Skip to content
KernelIndex
Search⌘K

submission 114273

Rayleon · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

cuda.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-114273?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
86.6µs
#364 of 678
2025-11-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:1e99b1cc0b8ad667161cc3e8cb0b8dcf02e841d98744c56f96389bcbeb932a22
license declaredunknown
license concludedunknown
authorsRayleon
imported2026-08-26

Techniques

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

fp4int halfK = K / 2; // number of fp4-packed bytes
split-k__global__ void reduce_splitk(
vector-width = int4int4 a_packed = reinterpret_cast<const int4*>(a)[A_offset + i];

Kernel source

cuda.py286 lines
import torch
from torch.utils.cpp_extension import load_inline
from typing import List
from task import input_t, output_t

_compiled_kernel_cache = None

add_cpp_source = """
#include <torch/extension.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <cstdio>

void bgemv_cuda(int M, int N, int K, int L, torch::Tensor A, torch::Tensor B, torch::Tensor sfA, torch::Tensor sfB, torch::Tensor C);
"""

add_cuda_source = """
#define VALUES_PER_LOOP 32
//__constant__ half fp4_e2m1_lut[16];

__global__ void implemented_bgemv(int M, int K, int L, uint8_t* a, uint8_t* b, uint8_t* sfa, uint8_t* sfb, at::Half* c, float* d, int K_split_count) {
    const unsigned int x = blockIdx.x * blockDim.x + threadIdx.x;
    const unsigned int y = blockIdx.y * blockDim.y + threadIdx.y; // actually batch dim

    int xOffset = x - M * (x / M);
    int xCount = (x / M);

    int halfK = K / 2;       // number of fp4-packed bytes
    int scaleK = K / 16;     // number of fp8 scale factors

    // precompute leading dimension products
    int A_row_stride   = halfK / 16;
    int A_batch_stride = M * halfK / 16;

    int B_batch_stride = 128 * halfK / 16;   // B has shape (128, K/2, L)
    int sfA_row_stride = scaleK;
    int sfA_batch_stride = M * scaleK;
    int sfB_batch_stride = 128 * scaleK;
    int sfA_offset = xOffset * sfA_row_stride + y * sfA_batch_stride;
    int sfB_offset = + y * sfB_batch_stride;
    int A_offset = xOffset * A_row_stride + y * A_batch_stride;
    int B_offset = y * B_batch_stride;

    // load B into local memory

    //int batch = blockIdx.x;      // one block per batch
    //int row   = threadIdx.x;     // thread per row (assume TILE_M threads)

    //const __nv_fp8_e4m3* x_batch = b + batch * K;

    // --- 1. Load vector x into shared memory (warp-strided) ---
    //int tid = threadIdx.x;
    //int num_warps = blockDim.x / warpSize;
    //int warp_id = tid / warpSize;
    //int lane    = tid % warpSize;

    //int warp_chunk = (K + num_warps - 1) / num_warps;
    //int start = warp_id * warp_chunk;
    //int end   = min(start + warp_chunk, K);

    //for (int i = start + lane; i < end; i += warpSize) {
    //    sX[i] = __half(x_batch[i]); // FP8 -> half
    //}

    //if (y >= L) return;

    //if (x >= M) {
        // second half of K
        float tmp = 0.0f;

        for (int i = (K/VALUES_PER_LOOP) / K_split_count * xCount; i < (K/VALUES_PER_LOOP) / K_split_count * (xCount + 1); i++) {
            
            //bs[threadIdx.x & 31 + (threadIdx.x >> 5) * 32] = b[B_offset * 8 + i * 8 + (threadIdx.x & 31)];

            //__syncwarp();
            //int idx_sfA = x * sfA_row_stride + (i) + y * sfA_batch_stride;
            //int idx_sfB = (i)               + y * sfB_batch_stride; // dim0 always 0

            half scaleA = __nv_cvt_fp8_to_halfraw(sfa[sfA_offset + i * VALUES_PER_LOOP/16], __NV_E4M3);
            half scaleA2 = __nv_cvt_fp8_to_halfraw(sfa[sfA_offset + i * VALUES_PER_LOOP/16 + 1], __NV_E4M3);
            

            half scaleB = __nv_cvt_fp8_to_halfraw(sfb[sfB_offset + i * VALUES_PER_LOOP/16], __NV_E4M3);
            half scaleB2 = __nv_cvt_fp8_to_halfraw(sfb[sfB_offset + i * VALUES_PER_LOOP/16 + 1], __NV_E4M3);

            // correct A index
            //int idxA = x * A_row_stride + i + y * A_batch_stride;
            
            int4 a_packed = reinterpret_cast<const int4*>(a)[A_offset + i];
            const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a_packed);

            half a_half[VALUES_PER_LOOP];
            #pragma unroll
            for (int j = 0; j < VALUES_PER_LOOP/2; j++) {
                reinterpret_cast<half2*>(&a_half)[j] = __nv_cvt_fp4x2_to_halfraw2(a_bytes[j],  __NV_E2M1);
                //a_half[2*j] = __nv_cvt_fp4_to_halfraw(a_bytes[j] & 0xF, __NV_E2M1);
                //a_half[2*j + 1] = __nv_cvt_fp4_to_halfraw(a_bytes[j] >> 4, __NV_E2M1);
            }

            int4 b_packed = reinterpret_cast<const int4*>(b)[B_offset + i];
            const uint8_t* b_bytes = reinterpret_cast<const uint8_t*>(&b_packed);

            half b_half[VALUES_PER_LOOP];
            #pragma unroll
            for (int j = 0; j < VALUES_PER_LOOP/2; j++) {
                reinterpret_cast<half2*>(&b_half)[j] = __nv_cvt_fp4x2_to_halfraw2(b_bytes[j],  __NV_E2M1);
                //b_half[2*j] = __nv_cvt_fp4_to_halfraw(b_bytes[j] & 0xF, __NV_E2M1);
                //b_half[2*j + 1] = __nv_cvt_fp4_to_halfraw(b_bytes[j] >> 4, __NV_E2M1);
            }

            half res[VALUES_PER_LOOP];
            #pragma unroll
            for (int j = 0; j < VALUES_PER_LOOP/2; j++) {
                reinterpret_cast<half2*>(&res)[j] = __hmul2(reinterpret_cast<half2*>(&a_half)[j], reinterpret_cast<half2*>(&b_half)[j]);
                //total = __hadd(total, __hmul(a_half[j], b_half[j]));
            }

            half total = __float2half(0.0f);
            #pragma unroll
            half pairedsums[16];
            reinterpret_cast<half2*>(&pairedsums)[0] = __hadd2(reinterpret_cast<half2*>(&res)[0], reinterpret_cast<half2*>(&res)[1]);
            reinterpret_cast<half2*>(&pairedsums)[1] = __hadd2(reinterpret_cast<half2*>(&res)[2], reinterpret_cast<half2*>(&res)[3]);
            reinterpret_cast<half2*>(&pairedsums)[2] = __hadd2(reinterpret_cast<half2*>(&res)[4], reinterpret_cast<half2*>(&res)[5]);
            reinterpret_cast<half2*>(&pairedsums)[3] = __hadd2(reinterpret_cast<half2*>(&res)[6], reinterpret_cast<half2*>(&res)[7]);
            reinterpret_cast<half2*>(&pairedsums)[4] = __hadd2(reinterpret_cast<half2*>(&pairedsums)[0], reinterpret_cast<half2*>(&pairedsums)[1]);
            reinterpret_cast<half2*>(&pairedsums)[5] = __hadd2(reinterpret_cast<half2*>(&pairedsums)[2], reinterpret_cast<half2*>(&pairedsums)[3]);
            reinterpret_cast<half2*>(&pairedsums)[6] = __hadd2(reinterpret_cast<half2*>(&pairedsums)[4], reinterpret_cast<half2*>(&pairedsums)[5]);
            total = __hadd(pairedsums[12], pairedsums[13]);
            //for (int j = 0; j < VALUES_PER_LOOP/2; j++) {
            //    total = __hadd(total, res[j]);
            //}

            half scale = __hmul(scaleA, scaleB);
            
            tmp += __half2float(scale) * __half2float(total);

            reinterpret_cast<half2*>(&pairedsums)[0] = __hadd2(reinterpret_cast<half2*>(&res)[8], reinterpret_cast<half2*>(&res)[9]);
            reinterpret_cast<half2*>(&pairedsums)[1] = __hadd2(reinterpret_cast<half2*>(&res)[10], reinterpret_cast<half2*>(&res)[11]);
            reinterpret_cast<half2*>(&pairedsums)[2] = __hadd2(reinterpret_cast<half2*>(&res)[12], reinterpret_cast<half2*>(&res)[13]);
            reinterpret_cast<half2*>(&pairedsums)[3] = __hadd2(reinterpret_cast<half2*>(&res)[14], reinterpret_cast<half2*>(&res)[15]);
            reinterpret_cast<half2*>(&pairedsums)[4] = __hadd2(reinterpret_cast<half2*>(&pairedsums)[0], reinterpret_cast<half2*>(&pairedsums)[1]);
            reinterpret_cast<half2*>(&pairedsums)[5] = __hadd2(reinterpret_cast<half2*>(&pairedsums)[2], reinterpret_cast<half2*>(&pairedsums)[3]);
            reinterpret_cast<half2*>(&pairedsums)[6] = __hadd2(reinterpret_cast<half2*>(&pairedsums)[4], reinterpret_cast<half2*>(&pairedsums)[5]);
            total = __hadd(pairedsums[12], pairedsums[13]);

            //total = __float2half(0.0f);
            //#pragma unroll
            //for (int j = VALUES_PER_LOOP/2; j < VALUES_PER_LOOP; j++) {
            //    total = __hadd(total, res[j]);
            //}

            scale = __hmul(scaleA2, scaleB2);
            tmp += __half2float(scale) * __half2float(total);
        }
        d[xOffset + M*y + M*L*(xCount)] = tmp;
        //c[x + M * y] =  *reinterpret_cast<at::Half*>(&tmp);
        //}
    //} else {
        
}

__global__ void reduce_splitk(
    at::Half *C, const float *C_partial, int M, int K, int L, int K_splits)
{
    int y = blockIdx.y * blockDim.y + threadIdx.y;
    int x = blockIdx.x * blockDim.x + threadIdx.x;
    if (y >= L) return;
    float total = 0.0f;
    #pragma unroll
    for (int i = 0; i < K_splits; i++) {
        total += C_partial[x + M*y + L*M*i];
    }
    C[x + M*y] = __float2half(total);
}


void bgemv_cuda(int M, int N, int K, int L, torch::Tensor A, torch::Tensor B, torch::Tensor sfA, torch::Tensor sfB, torch::Tensor C) {
    int size = 32;    
    int K_splits = 8;
    
    dim3 blockDim(size, 1, 1);
    dim3 gridDim((M + size - 1) / size * K_splits, L, 1);
    


    // test for split k impact
    float *d_d;
    cudaMalloc(&d_d, M * sizeof(float) * (K_splits) * L); 

    //__half h_host[16];

    // Fill with some values (example: 0.0, 0.1, 0.2, ...)
    //half fp4_e2m1_lut_host[16] = {
    //    /* 0x0: 0000 */  __float2half(0.0f),
    //    /* 0x1: 0001 */  __float2half(0.5f),
    //    /* 0x2: 0010 */  __float2half(1.0f),
    //    /* 0x3: 0011 */  __float2half(1.5f),
    //    /* 0x4: 0100 */  __float2half(2.0f),
    //    /* 0x5: 0101 */  __float2half(3.0f),
    //    /* 0x6: 0110 */  __float2half(INFINITY),
    //    /* 0x7: 0111 */  __float2half(NAN),
    //    /* 0x8: 1000 */  __float2half(-0.0f),
    //    /* 0x9: 1001 */  __float2half(-0.5f),
    //    /* 0xA: 1010 */  __float2half(-1.0f),
    //    /* 0xB: 1011 */  __float2half(-1.5f),
    //    /* 0xC: 1100 */  __float2half(-2.0f),
    //    /* 0xD: 1101 */  __float2half(-3.0f),
    //    /* 0xE: 1110 */  __float2half(-INFINITY),
    //    /* 0xF: 1111 */  __float2half(NAN)
    //};

    // Copy to GPU constant memory
    //cudaMemcpyToSymbol(fp4_e2m1_lut, fp4_e2m1_lut_host, sizeof(h_host));

    implemented_bgemv<<<gridDim, blockDim>>>(M,K,L,A.data_ptr<uint8_t>(),B.data_ptr<uint8_t>(),sfA.data_ptr<uint8_t>(),sfB.data_ptr<uint8_t>(),C.data_ptr<at::Half>(), d_d, K_splits);

    cudaDeviceSynchronize();

    dim3 blockDimReduce(size, 32, 1);
    dim3 gridDimReduce((M + size - 1) / size, 1, 1);

    reduce_splitk<<<gridDimReduce, blockDimReduce>>>(C.data_ptr<at::Half>(), d_d, M, K, L, K_splits);
    cudaFree(d_d);

    //cudaError_t err = cudaGetLastError();
    //if (err != cudaSuccess) {
    //    throw std::runtime_error(cudaGetErrorString(err));
    //}
}
"""

def compile_kernel():
    """
    Compile the kernel once and cache it.
    This should be called before any timing measurements.

    Returns:
        The compiled kernel function
    """
    global _compiled_kernel_cache

    if _compiled_kernel_cache is not None:
        return _compiled_kernel_cache

    # Compile the kernel
    _compiled_kernel_cache = load_inline(
        name='bgemv_cuda',
        cpp_sources=add_cpp_source,
        cuda_sources=add_cuda_source,
        functions=['bgemv_cuda'],
        extra_cuda_cflags=["-gencode=arch=compute_100a,code=sm_100a"],
        verbose=True,
    ).bgemv_cuda

    return _compiled_kernel_cache

def custom_kernel(data: input_t) -> output_t:
    """
    Custom implementation of vector addition using CUDA.
    Args:
        inputs: List of pairs of tensors [A, B] to be added.
    Returns:
        Tensor containing element-wise sum.
    """
    
    compiled_func = compile_kernel()

    a, b, sfa_natural, sfb_natural, _, _, c = data

    m, k, l = a.shape
    # Torch use e2m1_x2 data type, thus k is halved
    k = k * 2
    n = 1

    a_uint8 = a.view(torch.uint8)
    b_uint8 = b.view(torch.uint8)

    # sfa_uint8 = sfa.view(torch.uint8)
    # sfb_uint8 = sfb.view(torch.uint8)

    sfa_uint8_natural = sfa_natural.view(torch.uint8)
    sfb_uint8_natural = sfb_natural.view(torch.uint8)

    compiled_func(m, n, k, l, a_uint8, b_uint8, sfa_uint8_natural, sfb_uint8_natural, c)

    return c
scrolls · 286 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