Skip to content
KernelIndex
Search⌘K

submission 109296

JB Gage · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-109296?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
3.42ms
#676 of 678
2025-11-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:542c89de0cbc398f51ababf0ef012e24cfe63a6e4f0ddf6fbe4560cb1c896a29
license declaredunknown
license concludedunknown
authorsJB Gage
imported2026-08-15

Techniques

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

fp4return (float)float_e2m1_t::bitcast(bits);
fp8const __nv_fp8_e4m3* fp8_ptr = (const __nv_fp8_e4m3*)ptr;
shared-memory__shared__ uint8_t smemA[BLOCK_M * TILE_K_BYTES];
tile-m = 128constexpr int BLOCK_M = 128;

Kernel source

submission.py248 lines
import torch
from torch.utils.cpp_extension import load_inline

# ==============================================================================
# CONFIGURATION
# ==============================================================================
TARGET_B200 = True

# ==============================================================================
# 1. PATH FINDER
# ==============================================================================
def find_cutlass():
    import os
    if os.path.exists("./cutlass/include"):
        return [os.path.abspath("./cutlass/include"), os.path.abspath("./cutlass/tools/util/include")]
    return ["/opt/cutlass/4.3.0/include", "/opt/cutlass/4.3.0/tools/util/include"]

# ==============================================================================
# 2. CUDA SOURCE - OPTIMIZED VERSION
# ==============================================================================
cuda_source = r"""
#include <cuda_runtime.h>
#include <cstdint>
#include <cuda_fp16.h>

#ifdef TARGET_B200
#include <cuda_fp8.h>
#include <cutlass/numeric_types.h>
using namespace cutlass;
#endif

__device__ __forceinline__ float unpack_e2m1(uint8_t packed_byte, int which_nibble) {
    uint8_t bits = (which_nibble == 0) ? (packed_byte & 0x0F) : (packed_byte >> 4);
#ifdef TARGET_B200
    return (float)float_e2m1_t::bitcast(bits);
#else
    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[bits];
#endif
}

__device__ __forceinline__ float load_scale_fp8(const void* ptr, int idx) {
#ifdef TARGET_B200
    const __nv_fp8_e4m3* fp8_ptr = (const __nv_fp8_e4m3*)ptr;
    return (float)fp8_ptr[idx];
#else
    const half* h_ptr = (const half*)ptr;
    return __half2float(h_ptr[idx]);
#endif
}

extern "C" __global__ void __launch_bounds__(128) gemv_kernel(
    const uint8_t* __restrict__ A,
    const uint8_t* __restrict__ B,
    const void* __restrict__ SFA,
    const void* __restrict__ SFB,
    half* __restrict__ C,
    int M, int K, int L,
    int stride_a_0, int stride_a_1, int stride_a_2,
    int stride_b_0, int stride_b_1, int stride_b_2,
    int stride_sfa_0, int stride_sfa_1, int stride_sfa_2,
    int stride_sfb_0, int stride_sfb_1, int stride_sfb_2,
    int stride_c_0, int stride_c_1, int stride_c_2)
{
    int l_idx = blockIdx.z;
    int tid = threadIdx.x;

    constexpr int BLOCK_M = 128;
    constexpr int TILE_K_BYTES = 32;  // 64 elements per tile
    
    __shared__ uint8_t smemA[BLOCK_M * TILE_K_BYTES];
    __shared__ uint8_t smemB[TILE_K_BYTES];

    int block_row_start = blockIdx.x * BLOCK_M;
    int global_row = block_row_start + tid;
    
    float acc = 0.0f;
    
    int K_bytes = K / 2;
    int num_k_tiles = (K_bytes + TILE_K_BYTES - 1) / TILE_K_BYTES;
    
    bool row_valid = (global_row < M);

    const uint8_t* A_l = A + l_idx * stride_a_2;
    const uint8_t* B_l = B + l_idx * stride_b_2;
    const void* SFA_l = (const uint8_t*)SFA + l_idx * stride_sfa_2;
    const void* SFB_l = (const uint8_t*)SFB + l_idx * stride_sfb_2;

    for (int k_tile = 0; k_tile < num_k_tiles; ++k_tile) {
        int k_byte_start = k_tile * TILE_K_BYTES;
        int tile_size = min(TILE_K_BYTES, K_bytes - k_byte_start);
        
        // LOAD A - each thread loads its row's tile
        if (row_valid) {
            #pragma unroll
            for (int i = 0; i < TILE_K_BYTES; ++i) {
                if (i < tile_size) {
                    int k_byte = k_byte_start + i;
                    smemA[tid * TILE_K_BYTES + i] = A_l[global_row * stride_a_0 + k_byte * stride_a_1];
                }
            }
        }
        
        // LOAD B - first threads load the shared B tile
        if (tid < TILE_K_BYTES && tid < tile_size) {
            int k_byte = k_byte_start + tid;
            smemB[tid] = B_l[k_byte * stride_b_1];
        }
        
        __syncthreads();
        
        // COMPUTE - process in groups of 16 elements (8 bytes) per scale
        if (row_valid) {
            // Process 4 scale groups per tile (64 elements = 4 * 16)
            #pragma unroll
            for (int sg = 0; sg < 4; ++sg) {
                int local_byte_start = sg * 8;
                if (local_byte_start >= tile_size) break;
                
                int k_elem = (k_byte_start + local_byte_start) * 2;
                int scale_idx = k_elem / 16;
                
                float sa = load_scale_fp8(SFA_l, global_row * stride_sfa_0 + scale_idx * stride_sfa_1);
                float sb = load_scale_fp8(SFB_l, scale_idx * stride_sfb_1);
                float combined_scale = sa * sb;
                
                int bytes_in_group = min(8, tile_size - local_byte_start);
                
                #pragma unroll
                for (int b = 0; b < 8; ++b) {
                    if (b < bytes_in_group) {
                        int local_byte = local_byte_start + b;
                        uint8_t raw_a = smemA[tid * TILE_K_BYTES + local_byte];
                        uint8_t raw_b = smemB[local_byte];

                        float va0 = unpack_e2m1(raw_a, 0);
                        float vb0 = unpack_e2m1(raw_b, 0);
                        float va1 = unpack_e2m1(raw_a, 1);
                        float vb1 = unpack_e2m1(raw_b, 1);
                        
                        acc += (va0 * vb0 + va1 * vb1) * combined_scale;
                    }
                }
            }
        }
        
        __syncthreads();
    }

    if (row_valid) {
        C[global_row * stride_c_0 + l_idx * stride_c_2] = __float2half(acc);
    }
}

extern "C" void launch_gemv(
    void* a, void* b, void* sfa, void* sfb, void* c,
    int m, int k, int l,
    int stride_a_0, int stride_a_1, int stride_a_2,
    int stride_b_0, int stride_b_1, int stride_b_2,
    int stride_sfa_0, int stride_sfa_1, int stride_sfa_2,
    int stride_sfb_0, int stride_sfb_1, int stride_sfb_2,
    int stride_c_0, int stride_c_1, int stride_c_2)
{
    constexpr int BLOCK_M = 128;
    dim3 block(BLOCK_M);
    dim3 grid((m + BLOCK_M - 1) / BLOCK_M, 1, l);
    
    gemv_kernel<<<grid, block>>>(
        (const uint8_t*)a,
        (const uint8_t*)b,
        sfa,
        sfb,
        (half*)c,
        m, k, l,
        stride_a_0, stride_a_1, stride_a_2,
        stride_b_0, stride_b_1, stride_b_2,
        stride_sfa_0, stride_sfa_1, stride_sfa_2,
        stride_sfb_0, stride_sfb_1, stride_sfb_2,
        stride_c_0, stride_c_1, stride_c_2
    );
}
"""

# ==============================================================================
# 3. C++ WRAPPER
# ==============================================================================
cpp_source = r"""
#include <torch/extension.h>

extern "C" void launch_gemv(
    void* a, void* b, void* sfa, void* sfb, void* c,
    int m, int k, int l,
    int stride_a_0, int stride_a_1, int stride_a_2,
    int stride_b_0, int stride_b_1, int stride_b_2,
    int stride_sfa_0, int stride_sfa_1, int stride_sfa_2,
    int stride_sfb_0, int stride_sfb_1, int stride_sfb_2,
    int stride_c_0, int stride_c_1, int stride_c_2);

void run_kernel_proxy(
    torch::Tensor a, 
    torch::Tensor b, 
    torch::Tensor sfa, 
    torch::Tensor sfb, 
    torch::Tensor c) 
{
    int m = a.size(0);
    int k = a.size(1) * 2;
    int l = a.size(2);

    launch_gemv(
        a.data_ptr(), b.data_ptr(), sfa.data_ptr(), sfb.data_ptr(), c.data_ptr(),
        m, k, l,
        a.stride(0), a.stride(1), a.stride(2),
        b.stride(0), b.stride(1), b.stride(2),
        sfa.stride(0), sfa.stride(1), sfa.stride(2),
        sfb.stride(0), sfb.stride(1), sfb.stride(2),
        c.stride(0), c.stride(1), c.stride(2)
    );
}
"""

# ==============================================================================
# 4. COMPILE
# ==============================================================================
extra_flags = ['-O3', '-std=c++17', '--use_fast_math', '-lineinfo']
if TARGET_B200:
    extra_flags.append('-DTARGET_B200')

custom_gemv_inline = load_inline(
    name='custom_gemv_v21',
    cpp_sources=cpp_source,
    cuda_sources=cuda_source,
    functions=['run_kernel_proxy'],
    extra_include_paths=find_cutlass(),
    extra_cuda_cflags=extra_flags,
    with_cuda=True
)

# ==============================================================================
# 5. ENTRY POINT
# ==============================================================================
def custom_kernel(data):
    a, b, sfa, sfb, _, _, c = data
    custom_gemv_inline.run_kernel_proxy(a, b, sfa, sfb, c)
    return c
scrolls · 248 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