Skip to content
KernelIndex
Search⌘K

submission 110187

revess · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission2_staging.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-110187?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
35.6µs
#194 of 678
2025-11-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:594ee22cb0b8af8e685bb7370fa5a9dede6509f3b301a3d8a94bcdbce8e78ac8
license declaredunknown
license concludedunknown
authorsrevess
imported2026-08-15

Techniques

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

shared-memory__global__ void __launch_bounds__(256) nvfp4_gemv_smem_lut(
vector-width = half2half2* smem_b = (half2*)(smem_buffer + 1024);

Kernel source

submission2_staging.py293 lines
import torch
import math
import struct
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

# -------------------------------------------------------------------------
# 1. Precompute Lookup Tables
# -------------------------------------------------------------------------

def float_to_half_bits(f):
    return struct.unpack('H', struct.pack('e', f))[0]

def create_fp4_lut_hex():
    values = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0]
    lut_vals = []
    for i in range(16):
        sign = -1.0 if i >= 8 else 1.0
        idx = i if i < 8 else i - 8
        lut_vals.append(sign * values[idx])
    
    cpp_array = []
    for byte_val in range(256):
        lo = lut_vals[byte_val & 0x0F]
        hi = lut_vals[(byte_val >> 4) & 0x0F]
        packed = (float_to_half_bits(hi) << 16) | float_to_half_bits(lo)
        cpp_array.append(f"0x{packed:08X}")
    return "{" + ",".join(cpp_array) + "}"

def create_fp8_lut():
    vals = []
    for i in range(256):
        sign = (i >> 7) & 1
        exp = (i >> 3) & 0xF
        mant = i & 0x7
        val = 0.0
        if exp == 0:
            val = (mant / 8.0) * (2 ** -6) if mant != 0 else 0.0
        elif exp == 15:
            val = float('nan')
        else:
            val = (1.0 + mant / 8.0) * (2 ** (exp - 7))
        if sign: val = -val
        if math.isnan(val):
            vals.append("NAN")
        else:
            vals.append(f"{val}f")
    return "{" + ",".join(vals) + "}"

FP4_LUT_HEX = create_fp4_lut_hex()
FP8_LUT_STR = create_fp8_lut()

# -------------------------------------------------------------------------
# 2. CUDA Kernel Source
# -------------------------------------------------------------------------

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

__constant__ float FP8_LUT_HOST[256] = {FP8_LUT_STR};
__constant__ uint32_t FP4_LUT_CONST[256] = {FP4_LUT_HEX};

#define WARP_SIZE 32
#define WARPS_PER_BLOCK 8
#define THREADS_PER_BLOCK 256

// Padding config to avoid bank conflicts for B
#define CHUNK_STRIDE 9 

__global__ void __launch_bounds__(256) nvfp4_gemv_smem_lut(
    const uint8_t* __restrict__ A,
    const uint8_t* __restrict__ B,
    const uint8_t* __restrict__ SFA,
    const uint8_t* __restrict__ SFB,
    half* __restrict__ C,
    const int M,
    const int K_packed,
    const int K_sf,
    const int stride_a_l,
    const int stride_b_l,
    const int stride_sfa_l,
    const int stride_sfb_l
) {{
    // Shared Memory Layout:
    // 1. FP4 LUT (256 * 4 bytes = 1024 bytes)
    // 2. B Matrix (Dynamic size, padded)
    
    extern __shared__ char smem_buffer[];
    uint32_t* smem_lut = (uint32_t*)smem_buffer;
    half2* smem_b = (half2*)(smem_buffer + 1024);
    
    const int tid = threadIdx.x;
    const int l_idx = blockIdx.z;
    const int warp_id = tid / WARP_SIZE;
    const int lane_id = tid % WARP_SIZE;
    const int row_idx = blockIdx.x * WARPS_PER_BLOCK + warp_id;
    
    const uint8_t* B_global = B + l_idx * stride_b_l;
    const uint8_t* SFB_global = SFB + l_idx * stride_sfb_l;

    // ----------------------------------------------------------------
    // 1. Init Shared Memory LUT & Decode B
    // ----------------------------------------------------------------
    
    // Copy LUT to SMEM (Coalesced, 256 threads exactly fill it)
    if (tid < 256) {{
        smem_lut[tid] = FP4_LUT_CONST[tid];
    }}
    // No sync needed yet if we rely on warp sync, but let's be safe.
    // Actually, we process B below, which needs LUT if we decode on the fly.
    // But we are decoding B using the *same* LUT.
    // Wait, B-decoding uses the LUT too.
    // So we must sync after loading LUT.
    
    __syncthreads();

    // Decode B
    for (int i = tid; i < K_sf; i += THREADS_PER_BLOCK) {{
        uint8_t sfb_raw = SFB_global[i];
        float scale = FP8_LUT_HOST[sfb_raw];
        half2 h_scale = __float2half2_rn(scale);
        
        uint2 b_pack = *reinterpret_cast<const uint2*>(B_global + i * 8);
        uint8_t* b_bytes = (uint8_t*)&b_pack;
        
        half2* dst_chunk = smem_b + (i * CHUNK_STRIDE);
        
        #pragma unroll
        for (int j = 0; j < 8; ++j) {{
            // Use SMEM LUT for B decoding too
            uint32_t lut = smem_lut[b_bytes[j]];
            half2 val = *reinterpret_cast<half2*>(&lut);
            dst_chunk[j] = __hmul2(val, h_scale);
        }}
    }}
    
    __syncthreads(); // B and LUT are ready
    
    // ----------------------------------------------------------------
    // 2. Compute Row
    // ----------------------------------------------------------------
    
    if (row_idx < M) {{
        const uint8_t* A_base = A + l_idx * stride_a_l + row_idx * K_packed;
        const uint8_t* SFA_base = SFA + l_idx * stride_sfa_l + row_idx * K_sf;
        
        float row_acc = 0.0f;
        
        int k_packed_end = K_packed;
        
        // Loop over A (32 bytes / 64 elems per stride, 16 bytes / 32 elems per thread)
        for (int k = lane_id * 16; k < k_packed_end; k += WARP_SIZE * 16) {{
            int sfa_idx = k >> 3; // k / 8
            
            // Load SFA
            uint16_t sfa_pack = *reinterpret_cast<const uint16_t*>(SFA_base + sfa_idx);
            float scale0 = FP8_LUT_HOST[sfa_pack & 0xFF];
            float scale1 = FP8_LUT_HOST[sfa_pack >> 8];
            
            // Load A
            int4 a_pack = *reinterpret_cast<const int4*>(A_base + k);
            uint8_t* a_bytes = (uint8_t*)&a_pack;
            
            // Pointers to B
            half2* b_ptr0 = smem_b + (sfa_idx * CHUNK_STRIDE);
            half2* b_ptr1 = smem_b + ((sfa_idx + 1) * CHUNK_STRIDE);
            
            half2 acc0 = __float2half2_rn(0.0f);
            half2 acc1 = __float2half2_rn(0.0f);
            
            // Inner loops using SMEM LUT
            #pragma unroll
            for (int j = 0; j < 8; ++j) {{
                // Random access to SMEM LUT
                uint32_t lut = smem_lut[a_bytes[j]];
                half2 va = *reinterpret_cast<half2*>(&lut);
                acc0 = __hfma2(va, b_ptr0[j], acc0);
            }}
            
            #pragma unroll
            for (int j = 8; j < 16; ++j) {{
                uint32_t lut = smem_lut[a_bytes[j]];
                half2 va = *reinterpret_cast<half2*>(&lut);
                acc1 = __hfma2(va, b_ptr1[j - 8], acc1);
            }}
            
            float2 f0 = __half22float2(acc0);
            float2 f1 = __half22float2(acc1);
            
            row_acc += (f0.x + f0.y) * scale0;
            row_acc += (f1.x + f1.y) * scale1;
        }}
        
        #pragma unroll
        for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {{
            row_acc += __shfl_down_sync(0xFFFFFFFF, row_acc, offset);
        }}
        
        if (lane_id == 0) {{
            C[row_idx + l_idx * M] = __float2half(row_acc);
        }}
    }}
}}

void set_shared_mem_config() {{
    cudaFuncSetAttribute(nvfp4_gemv_smem_lut, cudaFuncAttributeMaxDynamicSharedMemorySize, 98304);
}}

torch::Tensor nvfp4_gemv_cuda_launch(
    torch::Tensor A,
    torch::Tensor B, 
    torch::Tensor SFA,
    torch::Tensor SFB,
    torch::Tensor C,
    int64_t M,
    int64_t K,
    int64_t L
) {{
    const int K_packed = K / 2;
    const int K_sf = K / 16;
    
    int blocks_x = (M + 8 - 1) / 8;
    dim3 grid(blocks_x, 1, L);
    dim3 block(256);
    
    // SMEM: 1024 bytes (LUT) + B_padded
    size_t smem_size = 1024 + K_sf * 36; 

    nvfp4_gemv_smem_lut<<<grid, block, smem_size>>>(
        A.data_ptr<uint8_t>(),
        B.data_ptr<uint8_t>(),
        SFA.data_ptr<uint8_t>(),
        SFB.data_ptr<uint8_t>(),
        reinterpret_cast<half*>(C.data_ptr<at::Half>()),
        M, K_packed, K_sf,
        M * K_packed,
        B.size(0) * K_packed,
        M * K_sf,
        B.size(0) * K_sf
    );
    
    return C;
}}
'''

cpp_source = r'''
#include <torch/extension.h>

void set_shared_mem_config();

torch::Tensor nvfp4_gemv_cuda_launch(
    torch::Tensor A,
    torch::Tensor B,
    torch::Tensor SFA,
    torch::Tensor SFB,
    torch::Tensor C,
    int64_t M,
    int64_t K,
    int64_t L
);
'''

_nvfp4_module = None

def get_nvfp4_module():
    global _nvfp4_module
    if _nvfp4_module is None:
        _nvfp4_module = load_inline(
            name='nvfp4_gemv_smem_lut',
            cpp_sources=[cpp_source],
            cuda_sources=[cuda_source],
            functions=['nvfp4_gemv_cuda_launch', 'set_shared_mem_config'],
            verbose=False,
            extra_cuda_cflags=['-O3', '--use_fast_math', '-lineinfo', '-std=c++17', '--maxrregcount=64']
        )
        _nvfp4_module.set_shared_mem_config()
    return _nvfp4_module

def custom_kernel(data: input_t) -> output_t:
    a, b, sfa, sfb, x, y, c = data
    M, _, L = c.shape
    K = a.shape[1] * 2
    a_u8 = a.view(torch.uint8)
    b_u8 = b.view(torch.uint8)
    sfa_u8 = sfa.view(torch.uint8)
    sfb_u8 = sfb.view(torch.uint8)
    module = get_nvfp4_module()
    module.nvfp4_gemv_cuda_launch(a_u8, b_u8, sfa_u8, sfb_u8, c, M, K, L)
    return c
scrolls · 293 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