Skip to content
KernelIndex
Search⌘K

submission 70351

snowclipsed · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 251 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-70351?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVFP4 GEMVsuite of 3 cases
NVIDIA B200
459.9µs
#593 of 678
2025-11-11

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c830d9f96e20a8e2218ad4dc654f6bb0fafcea576a692ca3211e97b939ff4f9c
license declaredunknown
license concludedunknown
authorssnowclipsed
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4Vectorized NVFP4 GEMV with shared memory - 8x fewer K-loop iterations
shared-memoryextern __shared__ unsigned char smem[];

Kernel source

submission.py251 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

nvfp4_gemv_cuda = """
#include <cuda_fp16.h>
#include <cuda_runtime.h>

// === Helper Functions (must come before kernel) ===

__device__ __forceinline__ float dequant_fp4_e2m1(unsigned char packed_val, int idx) {
    unsigned char fp4_bits = (idx == 0) ? (packed_val & 0x0F) : ((packed_val >> 4) & 0x0F);
    const float lut[16] = {
        0.0f, 0.5f, 1.0f, 1.5f,
        2.0f, 3.0f, 4.0f, 6.0f,
        -0.0f, -0.5f, -1.0f, -1.5f,
        -2.0f, -3.0f, -4.0f, -6.0f
    };
    return lut[fp4_bits];
}

__device__ __forceinline__ float dequant_fp8_e4m3(unsigned char fp8_bits) {
    int sign = (fp8_bits >> 7) & 0x1;
    int exp = (fp8_bits >> 3) & 0xF;
    int mant = fp8_bits & 0x7;
    
    float val;
    if (exp == 0) {
        val = ldexpf(mant / 8.0f, -6);
    } else if (exp == 15) {
        val = 448.0f;
    } else {
        val = ldexpf(1.0f + mant / 8.0f, exp - 7);
    }
    return sign ? -val : val;
}

__device__ __forceinline__ int64_t blocked_scale_offset(
    int m, int k_block, int l,
    int rest_m_dim, int rest_k_dim,
    int64_t s0, int64_t s1, int64_t s2, int64_t s3, int64_t s4, int64_t s5
) {
    int mm = m / 128;
    int mm32 = m % 32;
    int mm4 = (m % 128) / 32;
    int kk = k_block / 4;
    int kk4 = k_block % 4;
    
    if (mm >= rest_m_dim || kk >= rest_k_dim) return -1;
    
    return mm32 * s0 + mm4 * s1 + mm * s2 + kk4 * s3 + kk * s4 + l * s5;
}

// === Vectorized Kernel ===

__global__ void nvfp4_gemv_vectorized(
    const unsigned char* __restrict__ a,
    const unsigned char* __restrict__ b,
    const unsigned char* __restrict__ sfa,
    const unsigned char* __restrict__ sfb,
    __half* __restrict__ c,
    int M, int K, int L, int B_rows,
    int sfa_rest_m, int sfa_rest_k, int sfb_rest_m, int sfb_rest_k,
    int64_t a_s0, int64_t a_s1, int64_t a_s2,
    int64_t b_s0, int64_t b_s1, int64_t b_s2,
    int64_t sfa_s0, int64_t sfa_s1, int64_t sfa_s2, int64_t sfa_s3, int64_t sfa_s4, int64_t sfa_s5,
    int64_t sfb_s0, int64_t sfb_s1, int64_t sfb_s2, int64_t sfb_s3, int64_t sfb_s4, int64_t sfb_s5,
    int64_t c_s0, int64_t c_s1, int64_t c_s2
) {
    extern __shared__ unsigned char smem[];
    unsigned char* b_shared = smem;
    unsigned char* sfb_shared = smem + (K / 2);
    
    const int warp_id = threadIdx.y;
    const int lane_id = threadIdx.x;
    const int m = blockIdx.x * blockDim.y + warp_id;
    const int l = blockIdx.y;
    const int tid = threadIdx.x + threadIdx.y * blockDim.x;
    
    if (m >= M || l >= L) return;
    
    const int K_bytes = K / 2;
    const int K_blocks = K / 16;
    const int b_m = 0;
    
    // Load b (coalesced)
    #pragma unroll 8
    for (int idx = tid; idx < K_bytes; idx += blockDim.x * blockDim.y) {
        b_shared[idx] = __ldg(&b[idx * b_s1 + l * b_s2]);
    }
    
    // Precompute scale offsets
    const int mm = m / 128;
    const int mm32 = m % 32;
    const int mm4 = (m % 128) / 32;
    const int64_t sfa_m_part = mm32 * sfa_s0 + mm4 * sfa_s1 + mm * sfa_s2 + l * sfa_s5;
    
    const int64_t sfb_l_offset = l * sfb_s5;
    #pragma unroll 8
    for (int idx = tid; idx < K_blocks; idx += blockDim.x * blockDim.y) {
        const int kk = idx / 4;
        const int kk4 = idx % 4;
        const int64_t sfb_offset = b_m * sfb_s0 + kk4 * sfb_s3 + kk * sfb_s4 + sfb_l_offset;
        sfb_shared[idx] = __ldg(&sfb[sfb_offset]);
    }
    __syncthreads();
    
    // === VECTORIZED: Process 8 bytes per iteration ===
    float thread_acc = 0.0f;
    #pragma unroll 1
    for (int k_byte = lane_id * 8; k_byte < K_bytes; k_byte += warpSize * 8) {
        const int64_t a_offset = m * a_s0 + k_byte * a_s1 + l * a_s2;
        
        // Load 8 bytes (64 bits) at once
        uint2 a_vec = make_uint2(0, 0);
        if (k_byte + 8 <= K_bytes) {
            a_vec = *reinterpret_cast<const uint2*>(&a[a_offset]);
        } else {
            a_vec.x = (k_byte < K_bytes) ? *reinterpret_cast<const uint*>(&a[a_offset]) : 0;
            a_vec.y = (k_byte + 4 < K_bytes) ? *reinterpret_cast<const uint*>(&a[a_offset + 4]) : 0;
        }
        
        // Process 16 FP4 values
        #pragma unroll 8
        for (int i = 0; i < 8; i++) {
            const int byte_idx = k_byte + i;
            if (byte_idx >= K_bytes) break;
            
            const int k = byte_idx * 2;
            const int k_block = byte_idx >> 3;
            
            const unsigned char a_val = reinterpret_cast<unsigned char*>(&a_vec)[i];
            const unsigned char b_val = b_shared[byte_idx];
            
            const float a_fp4_0 = dequant_fp4_e2m1(a_val, 0);
            const float a_fp4_1 = (k + 1 < K) ? dequant_fp4_e2m1(a_val, 1) : 0.0f;
            const float b_fp4_0 = dequant_fp4_e2m1(b_val, 0);
            const float b_fp4_1 = dequant_fp4_e2m1(b_val, 1);
            
            const int kk = k_block / 4;
            const int kk4 = k_block % 4;
            const int64_t sfa_offset = sfa_m_part + kk4 * sfa_s3 + kk * sfa_s4;
            
            // Bounds check (computational equivalence)
            if (sfa_offset >= 0) {
                const float scale_a = dequant_fp8_e4m3(__ldg(&sfa[sfa_offset]));
                const float scale_b = dequant_fp8_e4m3(sfb_shared[k_block]);
                
                const float a_scaled_0 = a_fp4_0 * scale_a;
                const float a_scaled_1 = a_fp4_1 * scale_a;
                const float b_scaled_0 = b_fp4_0 * scale_b;
                const float b_scaled_1 = b_fp4_1 * scale_b;
                
                thread_acc += a_scaled_0 * b_scaled_0;
                if (k + 1 < K) thread_acc += a_scaled_1 * b_scaled_1;
            }
        }
    }
    
    // Warp reduction
    #pragma unroll
    for (int offset = 16; offset > 0; offset >>= 1) {
        thread_acc += __shfl_down_sync(0xFFFFFFFF, thread_acc, offset);
    }
    
    if (lane_id == 0) {
        const int64_t c_offset = m * c_s0 + l * c_s2;
        c[c_offset] = __float2half(thread_acc);
    }
}

// === PyTorch wrapper (matches original signature) ===

torch::Tensor nvfp4_gemv(
    torch::Tensor a,
    torch::Tensor b,
    torch::Tensor sfa_permuted,
    torch::Tensor sfb_permuted,
    torch::Tensor c
) {
    TORCH_CHECK(a.device().is_cuda(), "tensors must be CUDA");
    TORCH_CHECK(sfa_permuted.dim() == 6, "sfa_permuted must be 6D");
    TORCH_CHECK(sfb_permuted.dim() == 6, "sfb_permuted must be 6D");
    
    unsigned char* a_ptr = reinterpret_cast<unsigned char*>(a.data_ptr());
    unsigned char* b_ptr = reinterpret_cast<unsigned char*>(b.data_ptr());
    unsigned char* sfa_ptr = reinterpret_cast<unsigned char*>(sfa_permuted.data_ptr());
    unsigned char* sfb_ptr = reinterpret_cast<unsigned char*>(sfb_permuted.data_ptr());
    
    int M = a.size(0);
    int K_bytes = a.size(1);
    int L = a.size(2);
    int K = K_bytes * 2;
    int B_rows = b.size(0);
    
    int sfa_dim2 = sfa_permuted.size(2);
    int sfa_dim4 = sfa_permuted.size(4);
    int sfb_dim2 = sfb_permuted.size(2);
    int sfb_dim4 = sfb_permuted.size(4);
    
    dim3 block(32, 16);
    dim3 grid((M + 15) / 16, L);
    size_t smem_size = K_bytes + K / 16;
    
    nvfp4_gemv_vectorized<<<grid, block, smem_size>>>(
        a_ptr, b_ptr, sfa_ptr, sfb_ptr,
        reinterpret_cast<__half*>(c.data_ptr<at::Half>()),
        M, K, L, B_rows,
        sfa_dim2, sfa_dim4, sfb_dim2, sfb_dim4,
        a.stride(0), a.stride(1), a.stride(2),
        b.stride(0), b.stride(1), b.stride(2),
        sfa_permuted.stride(0), sfa_permuted.stride(1), sfa_permuted.stride(2),
        sfa_permuted.stride(3), sfa_permuted.stride(4), sfa_permuted.stride(5),
        sfb_permuted.stride(0), sfb_permuted.stride(1), sfb_permuted.stride(2),
        sfb_permuted.stride(3), sfb_permuted.stride(4), sfb_permuted.stride(5),
        c.stride(0), c.stride(1), c.stride(2)
    );
    
    cudaError_t err = cudaGetLastError();
    TORCH_CHECK(err == cudaSuccess, "CUDA error: ", cudaGetErrorString(err));
    
    return c;
}
"""

nvfp4_gemv_cpp = """
#include <torch/extension.h>
torch::Tensor nvfp4_gemv(
    torch::Tensor a,
    torch::Tensor b,
    torch::Tensor sfa,
    torch::Tensor sfb,
    torch::Tensor c
);
"""

nvfp4_module = load_inline(
    name='nvfp4_gemv',
    cpp_sources=nvfp4_gemv_cpp,
    cuda_sources=nvfp4_gemv_cuda,
    functions=['nvfp4_gemv'],
    verbose=False,
    extra_cuda_cflags=['-O3', '--use_fast_math', '-arch=sm_100']
)

def custom_kernel(data: input_t) -> output_t:  # type: ignore
    """
    Vectorized NVFP4 GEMV with shared memory - 8x fewer K-loop iterations
    """
    a, b, _, _, sfa_permuted, sfb_permuted, c = data
    return nvfp4_module.nvfp4_gemv(a, b, sfa_permuted, sfb_permuted, c)
scrolls · 251 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 69575.

- from task import input_t, output_t
import torch
- import triton
- import triton.language as tl
+ from torch.utils.cpp_extension import load_inline
+ from task import input_t, output_t
- # =============================================================================
- # OPTIMIZED SCALE FACTOR BLOCKING - From your original kernel
- # =============================================================================
+ nvfp4_gemv_cuda = """
+ #include <cuda_fp16.h>
+ #include <cuda_runtime.h>
- def ceil_div(a, b):
- return (a + b - 1) // b
+ // === Helper Functions (must come before kernel) ===
- @triton.jit
- def blocked_transform_kernel(
- inp, out, M, K, L,
- s_im, s_ik, s_il, s_ol, s_oe,
- BLK: tl.constexpr
- ):
- """Optimized blocking transformation for scale factors"""
- pid_l = tl.program_id(0)
- pid_b = tl.program_id(1)
+ __device__ __forceinline__ float dequant_fp4_e2m1(unsigned char packed_val, int idx) {
+ unsigned char fp4_bits = (idx == 0) ? (packed_val & 0x0F) : ((packed_val >> 4) & 0x0F);
+ const float lut[16] = {
+ 0.0f, 0.5f, 1.0f, 1.5f,
+ 2.0f, 3.0f, 4.0f, 6.0f,
+ -0.0f, -0.5f, -1.0f, -1.5f,
+ -2.0f, -3.0f, -4.0f, -6.0f
+ };
+ return lut[fp4_bits];
+ }
+
+ __device__ __forceinline__ float dequant_fp8_e4m3(unsigned char fp8_bits) {
+ int sign = (fp8_bits >> 7) & 0x1;
+ int exp = (fp8_bits >> 3) & 0xF;
+ int mant = fp8_bits & 0x7;
- mk = M * K
- offs = pid_b * BLK + tl.arange(0, BLK)
- mask = offs < mk
+ float val;
+ if (exp == 0) {
+ val = ldexpf(mant / 8.0f, -6);
+ } else if (exp == 15) {
+ val = 448.0f;
+ } else {
+ val = ldexpf(1.0f + mant / 8.0f, exp - 7);
+ }
+ return sign ? -val : val;
+ }
+
+ __device__ __forceinline__ int64_t blocked_scale_offset(
+ int m, int k_block, int l,
+ int rest_m_dim, int rest_k_dim,
+ int64_t s0, int64_t s1, int64_t s2, int64_t s3, int64_t s4, int64_t s5
+ ) {
+ int mm = m / 128;
+ int mm32 = m % 32;
+ int mm4 = (m % 128) / 32;
+ int kk = k_block / 4;
+ int kk4 = k_block % 4;
- i = offs // K
- j = offs % K
+ if (mm >= rest_m_dim || kk >= rest_k_dim) return -1;
- nrb = (M + 127) // 128
- ncb = (K + 3) // 4
+ return mm32 * s0 + mm4 * s1 + mm * s2 + kk4 * s3 + kk * s4 + l * s5;
+ }
+
+ // === Vectorized Kernel ===
+
+ __global__ void nvfp4_gemv_vectorized(
+ const unsigned char* __restrict__ a,
+ const unsigned char* __restrict__ b,
+ const unsigned char* __restrict__ sfa,
+ const unsigned char* __restrict__ sfb,
+ __half* __restrict__ c,
+ int M, int K, int L, int B_rows,
+ int sfa_rest_m, int sfa_rest_k, int sfb_rest_m, int sfb_rest_k,
+ int64_t a_s0, int64_t a_s1, int64_t a_s2,
+ int64_t b_s0, int64_t b_s1, int64_t b_s2,
+ int64_t sfa_s0, int64_t sfa_s1, int64_t sfa_s2, int64_t sfa_s3, int64_t sfa_s4, int64_t sfa_s5,
+ int64_t sfb_s0, int64_t sfb_s1, int64_t sfb_s2, int64_t sfb_s3, int64_t sfb_s4, int64_t sfb_s5,
+ int64_t c_s0, int64_t c_s1, int64_t c_s2
+ ) {
+ extern __shared__ unsigned char smem[];
+ unsigned char* b_shared = smem;
+ unsigned char* sfb_shared = smem + (K / 2);
- rb = i // 128
- ri = i % 128
- cb = j // 4
- ci = j % 4
+ const int warp_id = threadIdx.y;
+ const int lane_id = threadIdx.x;
+ const int m = blockIdx.x * blockDim.y + warp_id;
+ const int l = blockIdx.y;
+ const int tid = threadIdx.x + threadIdx.y * blockDim.x;
- # Blocking logic from reference
- perm = rb * ncb * 128 * 4 + cb * 128 * 4 + ri * 4 + ci
+ if (m >= M || l >= L) return;
- chunk = perm // 512
- in_chunk = perm % 512
- d1 = in_chunk // 128
- rest = in_chunk % 128
- d2 = rest // 4
- d3 = rest % 4
+ const int K_bytes = K / 2;
+ const int K_blocks = K / 16;
+ const int b_m = 0;
- out_idx = chunk * 512 + d2 * 16 + d1 * 4 + d3
+ // Load b (coalesced)
+ #pragma unroll 8
+ for (int idx = tid; idx < K_bytes; idx += blockDim.x * blockDim.y) {
+ b_shared[idx] = __ldg(&b[idx * b_s1 + l * b_s2]);
+ }
- # Load and store
- inp_idx = pid_l * s_il + i * s_im + j * s_ik
- out_idx_final = pid_l * s_ol + out_idx * s_oe
+ // Precompute scale offsets
+ const int mm = m / 128;
+ const int mm32 = m % 32;
+ const int mm4 = (m % 128) / 32;
+ const int64_t sfa_m_part = mm32 * sfa_s0 + mm4 * sfa_s1 + mm * sfa_s2 + l * sfa_s5;
- val = tl.load(inp + inp_idx, mask=mask)
- tl.store(out + out_idx_final, val, mask=mask)
-
- def transform_scales_gpu(tensor):
- """GPU-based scale transformation"""
- M, K, L = tensor.shape
- mk = M * K
- result = torch.empty((L, mk), dtype=tensor.dtype, device='cuda')
+ const int64_t sfb_l_offset = l * sfb_s5;
+ #pragma unroll 8
+ for (int idx = tid; idx < K_blocks; idx += blockDim.x * blockDim.y) {
+ const int kk = idx / 4;
+ const int kk4 = idx % 4;
+ const int64_t sfb_offset = b_m * sfb_s0 + kk4 * sfb_s3 + kk * sfb_s4 + sfb_l_offset;
+ sfb_shared[idx] = __ldg(&sfb[sfb_offset]);
+ }
+ __syncthreads();
- t = tensor.cuda() if not tensor.is_cuda else tensor
-
- BLK = 256
- grid = (L, (mk + BLK - 1) // BLK)
-
- blocked_transform_kernel[grid](
- t, result, M, K, L,
- t.stride(0), t.stride(1), t.stride(2),
- result.stride(0), result.stride(1),
- BLK=BLK
- )
-
- return [result[i] for i in range(L)]
-
- # =============================================================================
- # CUDA GRAPH OPTIMIZATION
- # =============================================================================
-
- _graph_cache = {}
-
- class CUDAGraphExecutor:
- """Captures CUDA graph to eliminate kernel launch overhead"""
- def __init__(self, M, K, L):
- self.M = M
- self.K = K
- self.L = L
- self.graph = None
- self.static_a = None
- self.static_b = None
- self.static_sfa_list = None
- self.static_sfb_list = None
- self.static_c = None
+ // === VECTORIZED: Process 8 bytes per iteration ===
+ float thread_acc = 0.0f;
+ #pragma unroll 1
+ for (int k_byte = lane_id * 8; k_byte < K_bytes; k_byte += warpSize * 8) {
+ const int64_t a_offset = m * a_s0 + k_byte * a_s1 + l * a_s2;
- def capture(self, a, b, sfa_list, sfb_list, c):
- """Capture CUDA graph"""
- # Warmup
- for _ in range(3):
- for i in range(self.L):
- res = torch._scaled_mm(
- a[:, :, i],
- b[:, :, i].transpose(0, 1),
- sfa_list[i],
- sfb_list[i],
- bias=None,
- out_dtype=torch.float16,
- )
- c[:, 0, i] = res[:, 0]
- torch.cuda.synchronize()
+ // Load 8 bytes (64 bits) at once
+ uint2 a_vec = make_uint2(0, 0);
+ if (k_byte + 8 <= K_bytes) {
+ a_vec = *reinterpret_cast<const uint2*>(&a[a_offset]);
+ } else {
+ a_vec.x = (k_byte < K_bytes) ? *reinterpret_cast<const uint*>(&a[a_offset]) : 0;
+ a_vec.y = (k_byte + 4 < K_bytes) ? *reinterpret_cast<const uint*>(&a[a_offset + 4]) : 0;
+ }
- # Create static tensors
- self.static_a = a.clone()
- self.static_b = b.clone()
- self.static_sfa_list = [s.clone() for s in sfa_list]
- self.static_sfb_list = [s.clone() for s in sfb_list]
- self.static_c = c.clone()
-
- # Capture
- self.graph = torch.cuda.CUDAGraph()
- with torch.cuda.graph(self.graph):
- for i in range(self.L):
- res = torch._scaled_mm(
- self.static_a[:, :, i],
- self.static_b[:, :, i].transpose(0, 1),
- self.static_sfa_list[i],
- self.static_sfb_list[i],
- bias=None,
- out_dtype=torch.float16,
- )
- self.static_c[:, 0, i] = res[:, 0]
-
- return self
+ // Process 16 FP4 values
+ #pragma unroll 8
+ for (int i = 0; i < 8; i++) {
+ const int byte_idx = k_byte + i;
+ if (byte_idx >= K_bytes) break;
+
+ const int k = byte_idx * 2;
+ const int k_block = byte_idx >> 3;
+
+ const unsigned char a_val = reinterpret_cast<unsigned char*>(&a_vec)[i];
+ const unsigned char b_val = b_shared[byte_idx];
+
+ const float a_fp4_0 = dequant_fp4_e2m1(a_val, 0);
+ const float a_fp4_1 = (k + 1 < K) ? dequant_fp4_e2m1(a_val, 1) : 0.0f;
+ const float b_fp4_0 = dequant_fp4_e2m1(b_val, 0);
+ const float b_fp4_1 = dequant_fp4_e2m1(b_val, 1);
+
+ const int kk = k_block / 4;
+ const int kk4 = k_block % 4;
+ const int64_t sfa_offset = sfa_m_part + kk4 * sfa_s3 + kk * sfa_s4;
+
+ // Bounds check (computational equivalence)
+ if (sfa_offset >= 0) {
+ const float scale_a = dequant_fp8_e4m3(__ldg(&sfa[sfa_offset]));
+ const float scale_b = dequant_fp8_e4m3(sfb_shared[k_block]);
+
+ const float a_scaled_0 = a_fp4_0 * scale_a;
+ const float a_scaled_1 = a_fp4_1 * scale_a;
+ const float b_scaled_0 = b_fp4_0 * scale_b;
+ const float b_scaled_1 = b_fp4_1 * scale_b;
+
+ thread_acc += a_scaled_0 * b_scaled_0;
+ if (k + 1 < K) thread_acc += a_scaled_1 * b_scaled_1;
+ }
+ }
+ }
- def execute(self, a, b, sfa_list, sfb_list, c):
- """Execute with graph"""
- self.static_a.copy_(a, non_blocking=True)
- self.static_b.copy_(b, non_blocking=True)
- for i in range(self.L):
- self.static_sfa_list[i].copy_(sfa_list[i], non_blocking=True)
- self.static_sfb_list[i].copy_(sfb_list[i], non_blocking=True)
-
- self.graph.replay()
-
- c.copy_(self.static_c, non_blocking=True)
- return c
+ // Warp reduction
+ #pragma unroll
+ for (int offset = 16; offset > 0; offset >>= 1) {
+ thread_acc += __shfl_down_sync(0xFFFFFFFF, thread_acc, offset);
+ }
+
+ if (lane_id == 0) {
+ const int64_t c_offset = m * c_s0 + l * c_s2;
+ c[c_offset] = __float2half(thread_acc);
+ }
+ }
- # =============================================================================
- # MAIN KERNEL
- # =============================================================================
+ // === PyTorch wrapper (matches original signature) ===
- def custom_kernel(data: input_t) -> output_t:
- """
- Optimized NVFP4 batched GEMV kernel.
+ torch::Tensor nvfp4_gemv(
+ torch::Tensor a,
+ torch::Tensor b,
+ torch::Tensor sfa_permuted,
+ torch::Tensor sfb_permuted,
+ torch::Tensor c
+ ) {
+ TORCH_CHECK(a.device().is_cuda(), "tensors must be CUDA");
+ TORCH_CHECK(sfa_permuted.dim() == 6, "sfa_permuted must be 6D");
+ TORCH_CHECK(sfb_permuted.dim() == 6, "sfb_permuted must be 6D");
- KEY OPTIMIZATIONS:
- 1. GPU-based scale transformation (parallel across L batches)
- 2. CUDA graphs (eliminates kernel launch overhead)
- 3. Optimized memory access patterns
+ unsigned char* a_ptr = reinterpret_cast<unsigned char*>(a.data_ptr());
+ unsigned char* b_ptr = reinterpret_cast<unsigned char*>(b.data_ptr());
+ unsigned char* sfa_ptr = reinterpret_cast<unsigned char*>(sfa_permuted.data_ptr());
+ unsigned char* sfb_ptr = reinterpret_cast<unsigned char*>(sfb_permuted.data_ptr());
- LIMITATIONS:
- - Still uses torch._scaled_mm which doesn't use Blackwell tensor cores
- - To go faster: Need CUTLASS with tcgen05 instructions (multi-file setup)
+ int M = a.size(0);
+ int K_bytes = a.size(1);
+ int L = a.size(2);
+ int K = K_bytes * 2;
+ int B_rows = b.size(0);
- Expected speedup: 1.5-2.5x over reference
- """
- a, b, sfa, sfb, sfa_permuted, sfb_permuted, c = data
+ int sfa_dim2 = sfa_permuted.size(2);
+ int sfa_dim4 = sfa_permuted.size(4);
+ int sfb_dim2 = sfb_permuted.size(2);
+ int sfb_dim4 = sfb_permuted.size(4);
- M = a.size(0)
- K = a.size(1)
- L = a.size(2)
+ dim3 block(32, 16);
+ dim3 grid((M + 15) / 16, L);
+ size_t smem_size = K_bytes + K / 16;
- # Transform scales using the CORRECT format for torch._scaled_mm
- # (Not the permuted format - that's for CUTLASS)
- sfa_list = transform_scales_gpu(sfa)
- sfb_list = transform_scales_gpu(sfb)
+ nvfp4_gemv_vectorized<<<grid, block, smem_size>>>(
+ a_ptr, b_ptr, sfa_ptr, sfb_ptr,
+ reinterpret_cast<__half*>(c.data_ptr<at::Half>()),
+ M, K, L, B_rows,
+ sfa_dim2, sfa_dim4, sfb_dim2, sfb_dim4,
+ a.stride(0), a.stride(1), a.stride(2),
+ b.stride(0), b.stride(1), b.stride(2),
+ sfa_permuted.stride(0), sfa_permuted.stride(1), sfa_permuted.stride(2),
+ sfa_permuted.stride(3), sfa_permuted.stride(4), sfa_permuted.stride(5),
+ sfb_permuted.stride(0), sfb_permuted.stride(1), sfb_permuted.stride(2),
+ sfb_permuted.stride(3), sfb_permuted.stride(4), sfb_permuted.stride(5),
+ c.stride(0), c.stride(1), c.stride(2)
+ );
- # Use CUDA graphs for repeated calls
- cache_key = (M, K, L)
+ cudaError_t err = cudaGetLastError();
+ TORCH_CHECK(err == cudaSuccess, "CUDA error: ", cudaGetErrorString(err));
- if cache_key not in _graph_cache:
- # First time - capture graph
- executor = CUDAGraphExecutor(M, K, L)
- executor.capture(a, b, sfa_list, sfb_list, c)
- _graph_cache[cache_key] = executor
- else:
- # Reuse cached graph
- executor = _graph_cache[cache_key]
-
- result = executor.execute(a, b, sfa_list, sfb_list, c)
-
- return result
No newline at end of file
+ return c;
+ }
+ """
+
+ nvfp4_gemv_cpp = """
+ #include <torch/extension.h>
+ torch::Tensor nvfp4_gemv(
+ torch::Tensor a,
+ torch::Tensor b,
+ torch::Tensor sfa,
+ torch::Tensor sfb,
+ torch::Tensor c
+ );
+ """
+
+ nvfp4_module = load_inline(
+ name='nvfp4_gemv',
+ cpp_sources=nvfp4_gemv_cpp,
+ cuda_sources=nvfp4_gemv_cuda,
+ functions=['nvfp4_gemv'],
+ verbose=False,
+ extra_cuda_cflags=['-O3', '--use_fast_math', '-arch=sm_100']
+ )
+
+ def custom_kernel(data: input_t) -> output_t: # type: ignore
+ """
+ Vectorized NVFP4 GEMV with shared memory - 8x fewer K-loop iterations
+ """
+ a, b, _, _, sfa_permuted, sfb_permuted, c = data
+ return nvfp4_module.nvfp4_gemv(a, b, sfa_permuted, sfb_permuted, c)
No newline at end of file
scrolls · 419 diff lines total

Best evidence level for this revision: reported

JSON