Skip to content
KernelIndex
Search⌘K

submission 88730

msuiche · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2b2e21875636a4ade0240d38e4d1834c9db012cc0c5c9ccd7eb2cab892b4c8ae
license declaredunknown
license concludedunknown
authorsmsuiche
imported2026-08-26

Techniques

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

fp4Optimized CUDA FP4 GEMV - V3 with PyTorch inline compilation (no cupy dependency)
shared-memoryextern __shared__ float reduction_shared_mem[];

Kernel source

submission.py354 lines
"""
Optimized CUDA FP4 GEMV - V3 with PyTorch inline compilation (no cupy dependency)
"""
import torch
from task import input_t, output_t

cuda_source = """
#include <cuda_fp16.h>
#include <torch/extension.h>

#ifndef HUGE_VALF
#define HUGE_VALF __int_as_float(0x7f800000)
#endif

__device__ __constant__ float fp4_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__ __forceinline__ float fp8_to_float(unsigned char fp8_val) {
    unsigned int sign = (fp8_val >> 7) & 1;
    unsigned int exp = (fp8_val >> 3) & 0xF;
    unsigned int mant = fp8_val & 0x7;

    if (exp == 0) {
        if (mant == 0) return sign ? -0.0f : 0.0f;
        return (sign ? -1.0f : 1.0f) * ldexpf((float)mant / 8.0f, -6);
    }
    if (exp == 0xF) {
        return sign ? -HUGE_VALF : HUGE_VALF;
    }

    float mantissa = 1.0f + (float)mant / 8.0f;
    int exponent = (int)exp - 7;
    float result = ldexpf(mantissa, exponent);

    return sign ? -result : result;
}

// Decode 4 FP4 values from 2 bytes and accumulate
__device__ __forceinline__ void decode_and_accumulate_4fp4(
    unsigned char a_packed, unsigned char b_packed,
    float scale_a, float scale_b,
    float& acc)
{
    unsigned char a_low = a_packed & 0xF;
    unsigned char a_high = (a_packed >> 4) & 0xF;
    unsigned char b_low = b_packed & 0xF;
    unsigned char b_high = (b_packed >> 4) & 0xF;

    acc += fp4_lut[a_low] * scale_a * fp4_lut[b_low] * scale_b;
    acc += fp4_lut[a_high] * scale_a * fp4_lut[b_high] * scale_b;
}

// Decode 8 FP4 values from 4 bytes (vectorized)
__device__ __forceinline__ void decode_and_accumulate_8fp4(
    const unsigned int a_vec, const unsigned int b_vec,
    float scale_a, float scale_b,
    float& acc)
{
    // Extract 4 bytes from each uint32_t
    #pragma unroll
    for (int i = 0; i < 4; i++) {
        unsigned char a_byte = (a_vec >> (i * 8)) & 0xFF;
        unsigned char b_byte = (b_vec >> (i * 8)) & 0xFF;
        decode_and_accumulate_4fp4(a_byte, b_byte, scale_a, scale_b, acc);
    }
}

__device__ float warp_reduce_sum(float val) {
    #pragma unroll
    for (int offset = 16; offset > 0; offset /= 2) {
        val += __shfl_down_sync(0xffffffff, val, offset);
    }
    return val;
}

__device__ float block_reduce_sum(float val, float* shared) {
    int lane = threadIdx.x % 32;
    int wid = threadIdx.x / 32;

    val = warp_reduce_sum(val);

    if (lane == 0) {
        shared[wid] = val;
    }
    __syncthreads();

    if (wid == 0) {
        val = (threadIdx.x < (blockDim.x + 31) / 32) ? shared[lane] : 0.0f;
        val = warp_reduce_sum(val);
    }

    return val;
}

/**
 * Vectorized parallel K-reduction kernel
 * Uses uint32_t loads for 4-byte (8 FP4) vectorization
 */
__global__ void fp4_gemv_vectorized_kernel(
    const unsigned char* __restrict__ a_ptr,
    const unsigned char* __restrict__ b_ptr,
    const unsigned char* __restrict__ sfa_ptr,
    const unsigned char* __restrict__ sfb_ptr,
    half* __restrict__ c_ptr,
    int M, int K, int L)
{
    extern __shared__ float reduction_shared_mem[];

    int K_packed = K / 2;
    int K_scales = K / 16;

    int m = blockIdx.x;
    int batch_idx = blockIdx.y;

    if (m >= M || batch_idx >= L) return;

    float local_sum = 0.0f;

    // Vectorized loop: process 4 bytes (8 FP4) at a time
    // Each thread processes multiple 4-byte chunks
    for (int k_packed = threadIdx.x * 4; k_packed < K_packed; k_packed += blockDim.x * 4) {
        // Check if we can do a vectorized load (need 4 contiguous bytes)
        if (k_packed + 3 < K_packed) {
            // Load 4 bytes at once using uint32_t
            int a_base = m * K_packed * L + k_packed * L + batch_idx;
            int b_base = k_packed * L + batch_idx;

            // Manual 4-byte load (safer than reinterpret_cast with alignment issues)
            unsigned int a_vec = 0, b_vec = 0;
            #pragma unroll
            for (int i = 0; i < 4; i++) {
                a_vec |= ((unsigned int)a_ptr[a_base + i * L]) << (i * 8);
                b_vec |= ((unsigned int)b_ptr[b_base + i * L]) << (i * 8);
            }

            // Get scale factor (1 per 8 packed bytes = 16 FP4)
            int scale_idx = k_packed / 8;
            if (scale_idx < K_scales) {
                int sfa_idx = m * K_scales * L + scale_idx * L + batch_idx;
                int sfb_idx = scale_idx * L + batch_idx;

                float scale_a = fp8_to_float(sfa_ptr[sfa_idx]);
                float scale_b = fp8_to_float(sfb_ptr[sfb_idx]);

                decode_and_accumulate_8fp4(a_vec, b_vec, scale_a, scale_b, local_sum);
            }
        } else {
            // Handle remaining elements (scalar)
            #pragma unroll
            for (int i = 0; i < 4 && k_packed + i < K_packed; i++) {
                int a_idx = m * K_packed * L + (k_packed + i) * L + batch_idx;
                int b_idx = (k_packed + i) * L + batch_idx;

                unsigned char a_packed = a_ptr[a_idx];
                unsigned char b_packed = b_ptr[b_idx];

                int scale_idx = (k_packed + i) / 8;
                if (scale_idx < K_scales) {
                    int sfa_idx = m * K_scales * L + scale_idx * L + batch_idx;
                    int sfb_idx = scale_idx * L + batch_idx;

                    float scale_a = fp8_to_float(sfa_ptr[sfa_idx]);
                    float scale_b = fp8_to_float(sfb_ptr[sfb_idx]);

                    decode_and_accumulate_4fp4(a_packed, b_packed, scale_a, scale_b, local_sum);
                }
            }
        }
    }

    // Reduce across threads
    float result = block_reduce_sum(local_sum, reduction_shared_mem);

    if (threadIdx.x == 0) {
        int c_idx = m * 1 * L + 0 * L + batch_idx;
        c_ptr[c_idx] = __float2half(result);
    }
}

/**
 * Hybrid kernel with shared memory for small K
 */
__global__ void fp4_gemv_hybrid_vectorized_kernel(
    const unsigned char* __restrict__ a_ptr,
    const unsigned char* __restrict__ b_ptr,
    const unsigned char* __restrict__ sfa_ptr,
    const unsigned char* __restrict__ sfb_ptr,
    half* __restrict__ c_ptr,
    int M, int K, int L)
{
    extern __shared__ unsigned char shared_mem[];

    int K_packed = K / 2;
    int K_scales = K / 16;

    unsigned char* b_shared = shared_mem;
    unsigned char* sfb_shared = shared_mem + K_packed;

    int batch_idx = blockIdx.z;
    if (batch_idx >= L) return;

    // Load B cooperatively
    for (int i = threadIdx.x; i < K_packed; i += blockDim.x) {
        b_shared[i] = b_ptr[i * L + batch_idx];
    }
    for (int i = threadIdx.x; i < K_scales; i += blockDim.x) {
        sfb_shared[i] = sfb_ptr[i * L + batch_idx];
    }
    __syncthreads();

    // Each thread processes one row
    int m = blockIdx.x * blockDim.x + threadIdx.x;
    if (m >= M) return;

    float acc = 0.0f;

    // Vectorized inner loop with shared memory
    #pragma unroll 4
    for (int k_packed = 0; k_packed < K_packed; k_packed++) {
        int a_idx = m * K_packed * L + k_packed * L + batch_idx;
        unsigned char a_packed = a_ptr[a_idx];
        unsigned char b_packed = b_shared[k_packed];

        int scale_idx = k_packed / 8;
        if (scale_idx < K_scales) {
            int sfa_idx = m * K_scales * L + scale_idx * L + batch_idx;

            float scale_a = fp8_to_float(sfa_ptr[sfa_idx]);
            float scale_b = fp8_to_float(sfb_shared[scale_idx]);

            decode_and_accumulate_4fp4(a_packed, b_packed, scale_a, scale_b, acc);
        }
    }

    int c_idx = m * 1 * L + 0 * L + batch_idx;
    c_ptr[c_idx] = __float2half(acc);
}

torch::Tensor fp4_gemv_vectorized_wrapper(
    torch::Tensor a,
    torch::Tensor b,
    torch::Tensor sfa,
    torch::Tensor sfb,
    torch::Tensor c,
    int M, int K, int L)
{
    const unsigned char* a_ptr = a.data_ptr<unsigned char>();
    const unsigned char* b_ptr = b.data_ptr<unsigned char>();
    const unsigned char* sfa_ptr = sfa.data_ptr<unsigned char>();
    const unsigned char* sfb_ptr = sfb.data_ptr<unsigned char>();
    at::Half* c_ptr = c.data_ptr<at::Half>();

    int threads = 256;
    dim3 grid(M, L, 1);
    dim3 block(threads, 1, 1);
    int shared_mem = threads * sizeof(float);

    fp4_gemv_vectorized_kernel<<<grid, block, shared_mem>>>(
        a_ptr, b_ptr, sfa_ptr, sfb_ptr, reinterpret_cast<half*>(c_ptr), M, K, L
    );

    return c;
}

torch::Tensor fp4_gemv_hybrid_wrapper(
    torch::Tensor a,
    torch::Tensor b,
    torch::Tensor sfa,
    torch::Tensor sfb,
    torch::Tensor c,
    int M, int K, int L)
{
    const uint8_t* a_ptr = a.data_ptr<uint8_t>();
    const uint8_t* b_ptr = b.data_ptr<uint8_t>();
    const uint8_t* sfa_ptr = sfa.data_ptr<uint8_t>();
    const uint8_t* sfb_ptr = sfb.data_ptr<uint8_t>();
    at::Half* c_ptr = c.data_ptr<at::Half>();

    int K_packed = K / 2;
    int K_scales = K / 16;

    int threads = 256;
    int blocks_m = (M + threads - 1) / threads;
    int shared_mem = K_packed + K_scales;

    dim3 grid(blocks_m, 1, L);
    dim3 block(threads, 1, 1);

    fp4_gemv_hybrid_vectorized_kernel<<<grid, block, shared_mem>>>(
        a_ptr, b_ptr, sfa_ptr, sfb_ptr, reinterpret_cast<half*>(c_ptr), M, K, L
    );

    return c;
}
"""

cpp_source = """
torch::Tensor fp4_gemv_vectorized_wrapper(
    torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c, int M, int K, int L);
torch::Tensor fp4_gemv_hybrid_wrapper(
    torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c, int M, int K, int L);
"""

# Global kernel cache
_kernel_module = None

def custom_kernel(data: input_t) -> output_t:
    global _kernel_module

    # Compile kernel once
    if _kernel_module is None:
        import sys
        if sys.stdout is None:
            class DummyFile:
                def write(self, x): pass
                def flush(self): pass
            sys.stdout = DummyFile()
            
        from torch.utils.cpp_extension import load_inline
        _kernel_module = load_inline(
            name='fp4_gemv_kernels',
            cpp_sources=cpp_source,
            cuda_sources=cuda_source,
            functions=['fp4_gemv_vectorized_wrapper', 'fp4_gemv_hybrid_wrapper'],
            verbose=False,
            extra_cuda_cflags=['-O3', '--use_fast_math', '-std=c++17']
        )

    a_ref, b_ref, sfa_ref, sfb_ref, sfa_permuted, sfb_permuted, c_ref = data

    M, K_packed, L = a_ref.shape
    K = K_packed * 2

    # Convert tensors
    a_uint8 = a_ref.view(torch.uint8).contiguous()
    b_uint8 = b_ref.view(torch.uint8).contiguous()
    sfa_uint8 = sfa_ref.view(torch.uint8).contiguous()
    sfb_uint8 = sfb_ref.view(torch.uint8).contiguous()
    c_out = c_ref.contiguous()

    # Choose kernel based on problem size
    if K >= 4096:
        _kernel_module.fp4_gemv_vectorized_wrapper(
            a_uint8, b_uint8, sfa_uint8, sfb_uint8, c_out, M, K, L
        )
    else:
        _kernel_module.fp4_gemv_hybrid_wrapper(
            a_uint8, b_uint8, sfa_uint8, sfb_uint8, c_out, M, K, L
        )

    return c_out
scrolls · 354 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