Skip to content
KernelIndex
Search⌘K

submission 103699

lucifer_0000007 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-103699?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
33.2µs
#179 of 678
2025-11-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9d1512cd651ffc5b7c1866050a793044290bb1751192258651b7b274e89106da
license declaredunknown
license concludedunknown
authorslucifer_0000007
imported2026-08-15

Techniques

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

shared-memory__shared__ half2 lut[256];
vector-width = uint4const uint4* __restrict__ a,

Kernel source

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

def generate_fp8_lut():
    lut = []
    for i in range(256):
        sign = (i >> 7) & 0x1
        exp = (i >> 3) & 0xF
        mant = i & 0x7
        val = 0.0
        if exp == 0:
            if mant != 0: val = (mant / 8.0) * (2 ** -6)
        else:
            if exp == 15 and mant == 7: val = 0.0
            else: val = (1.0 + mant / 8.0) * (2 ** (exp - 7))
        if sign: val = -val
        lut.append(f"{val:.8f}f")
    return "{" + ",".join(lut) + "}"

fp8_lut_str = generate_fp8_lut()

cuda_source = f'''
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cstdint>

__constant__ float c_FP4[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
}};

__constant__ float c_FP8[256] = {fp8_lut_str};

__global__ void __launch_bounds__(256, 8) nvfp4_gemv_kernel(
    const uint4* __restrict__ a,
    const uint4* __restrict__ b,
    const unsigned char* __restrict__ sfa,
    const unsigned char* __restrict__ sfb,
    half* __restrict__ c,
    int M, int num_vec,
    int a_stride, int a_batch_stride, int b_batch_stride,
    int sfa_s0, int sfa_s1, int sfa_s2,
    int sfb_s1, int sfb_s2,
    int c_s0, int c_s2
) {{
    __shared__ half2 lut[256];
    
    const int tid = threadIdx.x;
    const int lane = tid & 31;
    const int warp_id = tid >> 5;
    const int batch = blockIdx.z;
    const int m = blockIdx.x * 8 + warp_id;
    
    lut[tid] = __floats2half2_rn(c_FP4[tid & 0xF], c_FP4[tid >> 4]);
    __syncthreads();
    
    if (m >= M) return;
    
    const uint4* a_row = a + m * a_stride + batch * a_batch_stride;
    const uint4* b_row = b + batch * b_batch_stride;
    const unsigned char* sfa_row = sfa + (size_t)m * sfa_s0 + (size_t)batch * sfa_s2;
    const unsigned char* sfb_row = sfb + (size_t)batch * sfb_s2;
    
    float sum = 0.0f;
    
    // Process 2 vectors per iteration for better ILP
    int v = lane;
    for (; v + 32 < num_vec; v += 64) {{
        // Load 2 pairs of vectors
        uint4 av0 = __ldg(a_row + v);
        uint4 bv0 = __ldg(b_row + v);
        uint4 av1 = __ldg(a_row + v + 32);
        uint4 bv1 = __ldg(b_row + v + 32);
        
        int si0 = v * 2;
        int si1 = si0 + 1;
        int si2 = (v + 32) * 2;
        int si3 = si2 + 1;
        
        float sf0 = c_FP8[sfa_row[si0 * sfa_s1]] * c_FP8[sfb_row[si0 * sfb_s1]];
        float sf1 = c_FP8[sfa_row[si1 * sfa_s1]] * c_FP8[sfb_row[si1 * sfb_s1]];
        float sf2 = c_FP8[sfa_row[si2 * sfa_s1]] * c_FP8[sfb_row[si2 * sfb_s1]];
        float sf3 = c_FP8[sfa_row[si3 * sfa_s1]] * c_FP8[sfb_row[si3 * sfb_s1]];
        
        // First vector pair
        half2 p0 = __floats2half2_rn(0.0f, 0.0f);
        half2 p1 = __floats2half2_rn(0.0f, 0.0f);
        
        p0 = __hfma2(lut[av0.x & 0xFF], lut[bv0.x & 0xFF], p0);
        p0 = __hfma2(lut[(av0.x >> 8) & 0xFF], lut[(bv0.x >> 8) & 0xFF], p0);
        p0 = __hfma2(lut[(av0.x >> 16) & 0xFF], lut[(bv0.x >> 16) & 0xFF], p0);
        p0 = __hfma2(lut[av0.x >> 24], lut[bv0.x >> 24], p0);
        p0 = __hfma2(lut[av0.y & 0xFF], lut[bv0.y & 0xFF], p0);
        p0 = __hfma2(lut[(av0.y >> 8) & 0xFF], lut[(bv0.y >> 8) & 0xFF], p0);
        p0 = __hfma2(lut[(av0.y >> 16) & 0xFF], lut[(bv0.y >> 16) & 0xFF], p0);
        p0 = __hfma2(lut[av0.y >> 24], lut[bv0.y >> 24], p0);
        
        p1 = __hfma2(lut[av0.z & 0xFF], lut[bv0.z & 0xFF], p1);
        p1 = __hfma2(lut[(av0.z >> 8) & 0xFF], lut[(bv0.z >> 8) & 0xFF], p1);
        p1 = __hfma2(lut[(av0.z >> 16) & 0xFF], lut[(bv0.z >> 16) & 0xFF], p1);
        p1 = __hfma2(lut[av0.z >> 24], lut[bv0.z >> 24], p1);
        p1 = __hfma2(lut[av0.w & 0xFF], lut[bv0.w & 0xFF], p1);
        p1 = __hfma2(lut[(av0.w >> 8) & 0xFF], lut[(bv0.w >> 8) & 0xFF], p1);
        p1 = __hfma2(lut[(av0.w >> 16) & 0xFF], lut[(bv0.w >> 16) & 0xFF], p1);
        p1 = __hfma2(lut[av0.w >> 24], lut[bv0.w >> 24], p1);
        
        float f0 = __low2float(p0) + __high2float(p0);
        float f1 = __low2float(p1) + __high2float(p1);
        sum += f0 * sf0 + f1 * sf1;
        
        // Second vector pair
        half2 q0 = __floats2half2_rn(0.0f, 0.0f);
        half2 q1 = __floats2half2_rn(0.0f, 0.0f);
        
        q0 = __hfma2(lut[av1.x & 0xFF], lut[bv1.x & 0xFF], q0);
        q0 = __hfma2(lut[(av1.x >> 8) & 0xFF], lut[(bv1.x >> 8) & 0xFF], q0);
        q0 = __hfma2(lut[(av1.x >> 16) & 0xFF], lut[(bv1.x >> 16) & 0xFF], q0);
        q0 = __hfma2(lut[av1.x >> 24], lut[bv1.x >> 24], q0);
        q0 = __hfma2(lut[av1.y & 0xFF], lut[bv1.y & 0xFF], q0);
        q0 = __hfma2(lut[(av1.y >> 8) & 0xFF], lut[(bv1.y >> 8) & 0xFF], q0);
        q0 = __hfma2(lut[(av1.y >> 16) & 0xFF], lut[(bv1.y >> 16) & 0xFF], q0);
        q0 = __hfma2(lut[av1.y >> 24], lut[bv1.y >> 24], q0);
        
        q1 = __hfma2(lut[av1.z & 0xFF], lut[bv1.z & 0xFF], q1);
        q1 = __hfma2(lut[(av1.z >> 8) & 0xFF], lut[(bv1.z >> 8) & 0xFF], q1);
        q1 = __hfma2(lut[(av1.z >> 16) & 0xFF], lut[(bv1.z >> 16) & 0xFF], q1);
        q1 = __hfma2(lut[av1.z >> 24], lut[bv1.z >> 24], q1);
        q1 = __hfma2(lut[av1.w & 0xFF], lut[bv1.w & 0xFF], q1);
        q1 = __hfma2(lut[(av1.w >> 8) & 0xFF], lut[(bv1.w >> 8) & 0xFF], q1);
        q1 = __hfma2(lut[(av1.w >> 16) & 0xFF], lut[(bv1.w >> 16) & 0xFF], q1);
        q1 = __hfma2(lut[av1.w >> 24], lut[bv1.w >> 24], q1);
        
        float g0 = __low2float(q0) + __high2float(q0);
        float g1 = __low2float(q1) + __high2float(q1);
        sum += g0 * sf2 + g1 * sf3;
    }}
    
    // Handle remaining
    for (; v < num_vec; v += 32) {{
        uint4 av = __ldg(a_row + v);
        uint4 bv = __ldg(b_row + v);
        
        int si0 = v * 2;
        int si1 = si0 + 1;
        
        float sf0 = c_FP8[sfa_row[si0 * sfa_s1]] * c_FP8[sfb_row[si0 * sfb_s1]];
        float sf1 = c_FP8[sfa_row[si1 * sfa_s1]] * c_FP8[sfb_row[si1 * sfb_s1]];
        
        half2 p0 = __floats2half2_rn(0.0f, 0.0f);
        half2 p1 = __floats2half2_rn(0.0f, 0.0f);
        
        p0 = __hfma2(lut[av.x & 0xFF], lut[bv.x & 0xFF], p0);
        p0 = __hfma2(lut[(av.x >> 8) & 0xFF], lut[(bv.x >> 8) & 0xFF], p0);
        p0 = __hfma2(lut[(av.x >> 16) & 0xFF], lut[(bv.x >> 16) & 0xFF], p0);
        p0 = __hfma2(lut[av.x >> 24], lut[bv.x >> 24], p0);
        p0 = __hfma2(lut[av.y & 0xFF], lut[bv.y & 0xFF], p0);
        p0 = __hfma2(lut[(av.y >> 8) & 0xFF], lut[(bv.y >> 8) & 0xFF], p0);
        p0 = __hfma2(lut[(av.y >> 16) & 0xFF], lut[(bv.y >> 16) & 0xFF], p0);
        p0 = __hfma2(lut[av.y >> 24], lut[bv.y >> 24], p0);
        
        p1 = __hfma2(lut[av.z & 0xFF], lut[bv.z & 0xFF], p1);
        p1 = __hfma2(lut[(av.z >> 8) & 0xFF], lut[(bv.z >> 8) & 0xFF], p1);
        p1 = __hfma2(lut[(av.z >> 16) & 0xFF], lut[(bv.z >> 16) & 0xFF], p1);
        p1 = __hfma2(lut[av.z >> 24], lut[bv.z >> 24], p1);
        p1 = __hfma2(lut[av.w & 0xFF], lut[bv.w & 0xFF], p1);
        p1 = __hfma2(lut[(av.w >> 8) & 0xFF], lut[(bv.w >> 8) & 0xFF], p1);
        p1 = __hfma2(lut[(av.w >> 16) & 0xFF], lut[(bv.w >> 16) & 0xFF], p1);
        p1 = __hfma2(lut[av.w >> 24], lut[bv.w >> 24], p1);
        
        float f0 = __low2float(p0) + __high2float(p0);
        float f1 = __low2float(p1) + __high2float(p1);
        sum += f0 * sf0 + f1 * sf1;
    }}
    
    #pragma unroll
    for (int offset = 16; offset > 0; offset >>= 1) {{
        sum += __shfl_down_sync(0xffffffff, sum, offset);
    }}
    
    if (lane == 0) {{
        c[(size_t)m * c_s0 + (size_t)batch * c_s2] = __float2half(sum);
    }}
}}

void launch_gemv(
    int64_t a_ptr, int64_t b_ptr, int64_t sfa_ptr, int64_t sfb_ptr, int64_t c_ptr,
    int M, int K_packed, int L,
    int a_s0, int a_s2, int b_s2,
    int sfa_s0, int sfa_s1, int sfa_s2,
    int sfb_s1, int sfb_s2,
    int c_s0, int c_s2
) {{
    int num_vec = K_packed >> 4;
    dim3 blocks((M + 7) / 8, 1, L);
    dim3 threads(256);
    
    nvfp4_gemv_kernel<<<blocks, threads>>>(
        (uint4*)a_ptr, (uint4*)b_ptr,
        (unsigned char*)sfa_ptr, (unsigned char*)sfb_ptr,
        (half*)c_ptr, M, num_vec,
        a_s0 / 16, a_s2 / 16, b_s2 / 16,
        sfa_s0, sfa_s1, sfa_s2, sfb_s1, sfb_s2, c_s0, c_s2
    );
}}
'''

cpp_source = '''
void launch_gemv(int64_t, int64_t, int64_t, int64_t, int64_t,
    int, int, int, int, int, int, int, int, int, int, int, int, int);
'''

module = None

def get_module():
    global module
    if module is None:
        module = load_inline(
            name='nvfp4_gemv_v28',
            cpp_sources=cpp_source,
            cuda_sources=cuda_source,
            functions=['launch_gemv'],
            verbose=False,
            extra_cuda_cflags=['-O3', '--use_fast_math']
        )
    return module

def custom_kernel(data: input_t) -> output_t:
    a, b, sfa, sfb, _, _, c = data
    M, K_packed, L = a.shape

    get_module().launch_gemv(
        a.data_ptr(), b.data_ptr(), sfa.data_ptr(), sfb.data_ptr(), c.data_ptr(),
        M, K_packed, L,
        a.stride(0), a.stride(2),
        b.stride(2),
        sfa.stride(0), sfa.stride(1), sfa.stride(2),
        sfb.stride(1), sfb.stride(2),
        c.stride(0), c.stride(2)
    )
    return c
scrolls · 242 lines total

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

Best evidence level for this revision: reported

JSON