Skip to content
KernelIndex
Search⌘K

submission 115988

Theta Sigma · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

nvfp4_dom6.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-115988?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
36.9µs
#201 of 678
2025-11-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:12b6cc029ec7dc38c9442b34e2fb818db9920eb80ce759553eb31ce8195daf7f
license declaredunknown
license concludedunknown
authorsTheta Sigma
imported2026-08-15

Techniques

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

fp4custom implementation of NVFP4 block-scaled GEMV.
fp8__nv_fp8_storage_t fp8_val = *reinterpret_cast<__nv_fp8_storage_t*>(&val);
shared-memoryextern __shared__ char smem_buffer[];
vector-width = half2half2* lut_smem = (half2*)smem_buffer;

Kernel source

nvfp4_dom6.py253 lines
import torch
from task import input_t, output_t
from utils import make_match_reference
from torch.utils.cpp_extension import load_inline
import os
import subprocess

cuda_source = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <cutlass/cutlass.h>

__device__ const float fp4_e2m1_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
};

__device__ inline half unpack_fp8_half(uint8_t val) {
    __nv_fp8_storage_t fp8_val = *reinterpret_cast<__nv_fp8_storage_t*>(&val);
    return __nv_cvt_fp8_to_halfraw(fp8_val, __NV_E4M3);
}

#define ROWS_PER_BLOCK 4
#define WARP_SIZE 32
#define THREADS_PER_BLOCK 128
#define K_TILE_VEC 128  // 128 vectors * 16 bytes = 2048 bytes

__global__ void __launch_bounds__(THREADS_PER_BLOCK) gemv_vectorized(
    const void* __restrict__ Aptr,    
    const void* __restrict__ Bptr,    
    const void* __restrict__ SAptr,   
    const void* __restrict__ SBptr,   
    half* __restrict__ Cptr,          
    int m, 
    int num_k_vecs,
    long stride_l_a, long stride_m_a,
    long stride_l_b,     
    long stride_l_sa, long stride_m_sa,
    long stride_l_sb,    
    long stride_l_c      
) {
    // Shared Memory: lookup table (1KB) + B_Tile (2KB) + SB_Tile (256B)
    extern __shared__ char smem_buffer[];
    half2* lut_smem = (half2*)smem_buffer;
    uint4* B_smem_vec = (uint4*)(smem_buffer + 1024);
    ushort* SB_smem_vec = (ushort*)((char*)B_smem_vec + (K_TILE_VEC * sizeof(uint4)));

    int tid = threadIdx.x;
    int warp_id = tid / WARP_SIZE;
    int lane_id = tid % WARP_SIZE;
    
    if (tid < 128) {
        uint8_t packed1 = tid;
        half low1 = __float2half(fp4_e2m1_lut[packed1 & 0xF]);
        half high1 = __float2half(fp4_e2m1_lut[packed1 >> 4]);
        lut_smem[tid] = __halves2half2(low1, high1);
        
        uint8_t packed2 = tid + 128;
        half low2 = __float2half(fp4_e2m1_lut[packed2 & 0xF]);
        half high2 = __float2half(fp4_e2m1_lut[packed2 >> 4]);
        lut_smem[packed2] = __halves2half2(low2, high2);
    }
    __syncthreads();

    int block_row_start = blockIdx.x * ROWS_PER_BLOCK;
    int batch_idx = blockIdx.y;
    int my_row = block_row_start + warp_id;
    bool active_row = (my_row < m);

    const char* A_batch_ptr = (const char*)Aptr + (batch_idx * stride_l_a);
    const char* B_batch_ptr = (const char*)Bptr + (batch_idx * stride_l_b);
    const char* SA_batch_ptr = (const char*)SAptr + (batch_idx * stride_l_sa);
    const char* SB_batch_ptr = (const char*)SBptr + (batch_idx * stride_l_sb);
    
    const uint4* A_row_ptr = nullptr;
    const ushort* SA_row_ptr = nullptr;

    if (active_row) {
        A_row_ptr = (const uint4*)(A_batch_ptr + (my_row * stride_m_a));
        SA_row_ptr = (const ushort*)(SA_batch_ptr + (my_row * stride_m_sa));
    }

    const uint4* B_g = (const uint4*)B_batch_ptr;
    const ushort* SB_g = (const ushort*)SB_batch_ptr;

    half2 acc0 = __float2half2_rn(0.0f);
    half2 acc1 = __float2half2_rn(0.0f);

    for (int k_base = 0; k_base < num_k_vecs; k_base += K_TILE_VEC) {
        int tiles_remaining = num_k_vecs - k_base;
        int current_tile_size = (tiles_remaining < K_TILE_VEC) ? tiles_remaining : K_TILE_VEC;

        if (tid < current_tile_size) {
            B_smem_vec[tid] = B_g[k_base + tid];
            SB_smem_vec[tid] = SB_g[k_base + tid];
        }
        __syncthreads();

        if (active_row) {
            for (int i = lane_id; i < current_tile_size; i += WARP_SIZE) {
                uint4 pack_a = A_row_ptr[k_base + i];
                ushort scale_pack_a = SA_row_ptr[k_base + i];

                uint4 pack_b = B_smem_vec[i];
                ushort scale_pack_b = SB_smem_vec[i];

                half s_a_0 = unpack_fp8_half((uint8_t)(scale_pack_a & 0xFF));
                half s_a_1 = unpack_fp8_half((uint8_t)(scale_pack_a >> 8));
                half s_b_0 = unpack_fp8_half((uint8_t)(scale_pack_b & 0xFF));
                half s_b_1 = unpack_fp8_half((uint8_t)(scale_pack_b >> 8));

                half2 common_0 = __halves2half2(__hmul(s_a_0, s_b_0), __hmul(s_a_0, s_b_0));
                half2 common_1 = __halves2half2(__hmul(s_a_1, s_b_1), __hmul(s_a_1, s_b_1));

                uint32_t a_words[4] = {pack_a.x, pack_a.y, pack_a.z, pack_a.w};
                uint32_t b_words[4] = {pack_b.x, pack_b.y, pack_b.z, pack_b.w};

                #pragma unroll
                for(int w = 0; w < 4; w++) { 
                    uint32_t wa = a_words[w];
                    uint32_t wb = b_words[w];
                    half2 scale = (w < 2) ? common_0 : common_1;
                    // byte 0
                    acc0 = __hfma2(lut_smem[(uint8_t)(wa)], __hmul2(lut_smem[(uint8_t)(wb)], scale), acc0);
                    // byte 1
                    acc0 = __hfma2(lut_smem[(uint8_t)(wa >> 8)], __hmul2(lut_smem[(uint8_t)(wb >> 8)], scale), acc0);
                    // byte 2
                    acc1 = __hfma2(lut_smem[(uint8_t)(wa >> 16)],__hmul2(lut_smem[(uint8_t)(wb >> 16)], scale),acc1);
                    // byte 3
                    acc1 = __hfma2(lut_smem[(uint8_t)(wa >> 24)],__hmul2(lut_smem[(uint8_t)(wb >> 24)], scale),acc1);
                }
            }
        }
        __syncthreads();
    }

    // reduction and store
    if (active_row) {
        half2 sum_h2 = __hadd2(acc0, acc1);
        float psum = __low2float(sum_h2) + __high2float(sum_h2);

        // warp Reduction
        for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {
            psum += __shfl_down_sync(0xffffffff, psum, offset);
        }

        if (lane_id == 0) {
            half* C_batch = Cptr + (batch_idx * stride_l_c);
            C_batch[my_row] = __float2half(psum);
        }
    }
}

torch::Tensor gemv_bs_forward(torch::Tensor A, torch::Tensor B, torch::Tensor SA, torch::Tensor SB) {
    int m = A.size(0);
    int k_bytes = A.size(1) * A.element_size();
    int l = A.size(2);
    int num_k_vecs = k_bytes / 16;

    long stride_l_a = A.stride(2) * A.element_size();
    long stride_m_a = A.stride(0) * A.element_size();
    long stride_l_b = B.stride(2) * B.element_size();
    long stride_l_sa = SA.stride(2) * SA.element_size();
    long stride_m_sa = SA.stride(0) * SA.element_size();
    long stride_l_sb = SB.stride(2) * SB.element_size();

    auto options = torch::TensorOptions().dtype(torch::kHalf).device(A.device());
    auto C = torch::empty({l, m, 1}, options).permute({1, 2, 0});
    long stride_l_c = C.stride(2);

    int rows_per_block = 4;
    dim3 grid((m + rows_per_block - 1) / rows_per_block, l);
    dim3 block(128); 
    int smem_size = 4096; 

    gemv_vectorized<<<grid, block, smem_size>>>(
       (void*) A.data_ptr(), 
       (void*) B.data_ptr(), 
       (void*) SA.data_ptr(), 
       (void*) SB.data_ptr(), 
       (half*) C.data_ptr(), 
       m, num_k_vecs,
       stride_l_a, stride_m_a,
       stride_l_b, 
       stride_l_sa, stride_m_sa, 
       stride_l_sb, 
       stride_l_c
    );
    return C;
}
"""

cpp_source = "torch::Tensor gemv_bs_forward(torch::Tensor A, torch::Tensor B, torch::Tensor SA, torch::Tensor SB);"

gemv_ext = load_inline(
    name='gemv_bs_b200_v2', 
    cpp_sources=cpp_source,
    cuda_sources=cuda_source,
    functions=['gemv_bs_forward'],
    with_cuda=True,
    extra_cuda_cflags=["-O3", "-std=c++17", "--expt-relaxed-constexpr", "-arch=sm_90"],
    # extra_include_paths=[cutlass_include],
)

# Scaling factor vector size
sf_vec_size = 16

# Helper function for ceiling division
def ceil_div(a, b):
    return (a + b - 1) // b


# Helper function to convert scale factor tensor to blocked format
def to_blocked(input_matrix):
    rows, cols = input_matrix.shape

    # Please ensure rows and cols are multiples of 128 and 4 respectively
    n_row_blocks = ceil_div(rows, 128)
    n_col_blocks = ceil_div(cols, 4)

    padded = input_matrix
    blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
    rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)

    return rearranged.flatten()


def custom_kernel(
    data: input_t,
) -> output_t:
    """
    custom implementation of NVFP4 block-scaled GEMV.
    """
    a_ref, b_ref, sfa_ref, sfb_ref, _, _, c_ref = data
    return gemv_ext.gemv_bs_forward(a_ref, b_ref, sfa_ref, sfb_ref)

# if __name__ == "__main__":
#     M = 8192
#     K = 8192
#     L = 32
#     seed = 69

#     from nvfp4_gemv_b200 import generate_input

#     print(f"Generating inputs M={M}, K={K}, L={L}...")
#     _input = generate_input(m=M, k=K, l=L, seed=seed)
    
#     a_ref, b_ref, sfa_ref, sfb_ref, _, _, c_ref = _input
    
#     custom_kernel(_input)
scrolls · 253 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