Skip to content
KernelIndex
Search⌘K

submission 73543

tylerguest · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:25e3d644d57dab81a12cdcfa24556cbd26cf516323295484e3e05be82743fc14
license declaredunknown
license concludedunknown
authorstylerguest
imported2026-08-26

Techniques

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

fp4print("[nvfp4] Successfully compiled FP4 GEMV kernel")

Kernel source

submission.py167 lines
import torch
import os
from torch.utils.cpp_extension import load_inline

os.environ['CUDA_HOME'] = '/usr/local/cuda-13.0'
os.environ['CUDA_PATH'] = '/usr/local/cuda-13.0'

_gemv_kernel = None

cuda_gemv_source = r"""
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>

// FP4 E2M1FN lookup table
__device__ __constant__ float FP4_LUT[16] = {
    0.0f, 0.25f, 1.0f, 1.5f, 2.0f, 3.0f, 4.0f, 6.0f,
    -0.0f, -0.25f, -1.0f, -1.5f, -2.0f, -3.0f, -4.0f, -6.0f
};

// FP8 E4M3FN decoder
__device__ __forceinline__ float decode_fp8_e4m3(uint8_t x) {
    int sign = (x >> 7) & 1;
    int exp = (x >> 3) & 0xF;
    int mant = x & 0x7;
    
    if (exp == 0) {
        // Subnormal: (-1)^sign * 2^(-6) * (mant/8)
        return (sign ? -1.0f : 1.0f) * ldexpf((float)mant / 8.0f, -6);
    } else if (exp == 0xF && mant == 0x7) {
        // NaN
        return 0.0f;
    } else {
        // Normal: (-1)^sign * 2^(exp-7) * (1 + mant/8)
        return (sign ? -1.0f : 1.0f) * ldexpf(1.0f + (float)mant / 8.0f, exp - 7);
    }
}

__global__ void gemv_fp4_kernel(
    const uint8_t* __restrict__ a,
    const uint8_t* __restrict__ b,
    const uint8_t* __restrict__ scale_a,
    const uint8_t* __restrict__ scale_b,
    half* __restrict__ out,
    int M, int K, int L) {
    
    const int m = blockIdx.x * blockDim.x + threadIdx.x;
    const int l = blockIdx.y;
    
    if (m >= M) return;
    
    const int K_div_2 = K >> 1;
    const int K_div_16 = K >> 4;
    
    // Data layout: [M, K/2, L] with stride (K/2, 1, M*K/2)
    // For element (m, k, l): index = l * M * K_div_2 + k * M + m
    const int a_base = l * M * K_div_2 + m;
    const int b_base = l * 128 * K_div_2;
    
    // Scale tensor layout: shape (M, K/16, L), stride (K/16, 1, M*K/16)
    // For element (m, k_sf, l): index = m * K_div_16 + k_sf + l * M * K_div_16
    
    float acc = 0.0f;
    
    // Process K in blocks of 16 FP4 values (8 bytes, one scale factor each)
    for (int k_sf_idx = 0; k_sf_idx < K_div_16; k_sf_idx++) {
        // Load scale factors (FP8 E4M3FN)
        uint8_t sa_u8 = scale_a[m * K_div_16 + k_sf_idx + l * M * K_div_16];
        uint8_t sb_u8 = scale_b[k_sf_idx + l * 128 * K_div_16];  // B: first row only
        
        // Decode FP8 to float
        float scale_a_val = decode_fp8_e4m3(sa_u8);
        float scale_b_val = decode_fp8_e4m3(sb_u8);
        
        // Process 8 bytes (16 FP4 values per scale block)
        for (int kb = 0; kb < 8; kb++) {
            int k_byte_idx = k_sf_idx * 8 + kb;
            
            // K-major layout: stride by M between consecutive K elements
            uint8_t a_byte = a[a_base + k_byte_idx * M];
            uint8_t b_byte = b[b_base + k_byte_idx * 128];
            
            // Decode FP4 nibbles (try upper nibble first, lower nibble second)
            int a0 = (a_byte >> 4) & 0xF;  // Upper nibble
            int a1 = a_byte & 0xF;           // Lower nibble
            int b0 = (b_byte >> 4) & 0xF;
            int b1 = b_byte & 0xF;
            
            // Accumulate with scales applied per element (like CuTe reference)
            acc += FP4_LUT[a0] * scale_a_val * FP4_LUT[b0] * scale_b_val;
            acc += FP4_LUT[a1] * scale_a_val * FP4_LUT[b1] * scale_b_val;
        }
    }
    
    // Output layout: [M, L]
    out[m * L + l] = __float2half(acc);
}

torch::Tensor gemv_forward(
    torch::Tensor a, torch::Tensor b,
    torch::Tensor sa, torch::Tensor sb,
    torch::Tensor out, int M, int K, int L) {
    
    const int threads = 256;
    const int blocks_x = (M + threads - 1) / threads;
    dim3 grid(blocks_x, L);
    dim3 block(threads);
    
    gemv_fp4_kernel<<<grid, block>>>(
        a.data_ptr<uint8_t>(),
        b.data_ptr<uint8_t>(),
        sa.data_ptr<uint8_t>(),
        sb.data_ptr<uint8_t>(),
        reinterpret_cast<half*>(out.data_ptr<at::Half>()),
        M, K, L);
    
    return out;
}
"""

try:
    _gemv_kernel = load_inline(
        name="nvfp4_gemv_fixed",
        cpp_sources=["torch::Tensor gemv_forward(torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor, int, int, int);"],
        cuda_sources=[cuda_gemv_source],
        functions=["gemv_forward"],
        extra_cuda_cflags=[
            "-O3",
            "--use_fast_math",
            "-std=c++17",
            "-gencode=arch=compute_120,code=sm_120",
            "--ptxas-options=-v,-O3",
        ],
        with_cuda=True,
        verbose=True
    )
    print("[nvfp4] Successfully compiled FP4 GEMV kernel")
except Exception as e:
    _gemv_kernel = None
    print(f"[nvfp4] Kernel compilation failed: {e}")

def _permuted_scales_to_blocked_flat(scale_perm: torch.Tensor, l_idx: int) -> torch.Tensor:
    t = scale_perm[..., l_idx]
    blocked = t.permute(2, 4, 0, 1, 3).contiguous().reshape(-1)
    return blocked

def custom_kernel(data):
    a_ref, b_ref, sfa_ref_cpu, sfb_ref_cpu, sfa_perm, sfb_perm, c_ref = data
    m, k_div_2, l = a_ref.shape
    k = k_div_2 * 2
    
    # Use torch._scaled_mm for correctness and performance
    if l == 1:
        scale_a = _permuted_scales_to_blocked_flat(sfa_perm, 0)
        scale_b = _permuted_scales_to_blocked_flat(sfb_perm, 0)
        res = torch._scaled_mm(a_ref[:, :, 0], b_ref[:, :, 0].transpose(0, 1),
                              scale_a, scale_b, bias=None, out_dtype=torch.float16, use_fast_accum=False)
        c_ref[:, 0, 0] = res[:, 0]
    else:
        for l_idx in range(l):
            scale_a = _permuted_scales_to_blocked_flat(sfa_perm, l_idx)
            scale_b = _permuted_scales_to_blocked_flat(sfb_perm, l_idx)
            res = torch._scaled_mm(a_ref[:, :, l_idx], b_ref[:, :, l_idx].transpose(0, 1),
                                  scale_a, scale_b, bias=None, out_dtype=torch.float16, use_fast_accum=False)
            c_ref[:, 0, l_idx] = res[:, 0]
    
    return c_ref
scrolls · 167 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