Skip to content
KernelIndex
Search⌘K

submission 89542

amackenzie-jumptrading · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8e36e59f9ba5d8f69ca6b4be91667db1e0d363656dfe47eed006f4a788c7d737
license declaredunknown
license concludedunknown
authorsamackenzie-jumptrading
imported2026-08-15

Techniques

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

async-copy__device__ __forceinline__ void cp_async_16(void* smem_ptr, const void* glob_ptr) {
shared-memory__device__ __forceinline__ void cp_async_16(void* smem_ptr, const void* glob_ptr) {
vector-width = half2half2* s_lut_pair = (half2*)smem;

Kernel source

submission.py401 lines
import torch
from torch.utils.cpp_extension import load_inline

# -----------------------------------------------------------------------------
# CUDA Kernel Source
# -----------------------------------------------------------------------------

cuda_source = r"""
#include <cuda_runtime.h>
#include <cuda_fp16.h>

#define WARP_SIZE 32

// Small LUT for initialization
__device__ __constant__ float C_FP4_LUT_CONST[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 decode_fp8_fast(uint8_t x) {
    if (x == 0x80) return 0.0f;
    uint32_t val = x;
    uint32_t sign = (val & 0x80) << 24;
    uint32_t exp  = (val & 0x78) >> 3;
    uint32_t mant = (val & 0x07);
    uint32_t exp32 = exp + 120;
    return (val == 0) ? 0.0f : __int_as_float(sign | (exp32 << 23) | (mant << 20));
}

__device__ __forceinline__ void cp_async_16(void* smem_ptr, const void* glob_ptr) {
    uint32_t smem = static_cast<uint32_t>(__cvta_generic_to_shared(smem_ptr));
    asm volatile("cp.async.ca.shared.global.L2::128B [%0], [%1], 16;" :: "r"(smem), "l"(glob_ptr));
}

__device__ __forceinline__ void cp_async_commit() { asm volatile("cp.async.commit_group;"); }
__device__ __forceinline__ void cp_async_wait_all() { asm volatile("cp.async.wait_group 0;"); }

template <int WARPS_PER_BLOCK>
__global__ void __launch_bounds__(WARPS_PER_BLOCK * WARP_SIZE) fp4_gemv_b200_sol(
    const uint8_t* __restrict__ A,
    const uint8_t* __restrict__ B,
    const uint8_t* __restrict__ SFA,
    const uint8_t* __restrict__ SFB,
    half* __restrict__ C,
    int M, int K_bytes, int L,
    int stride_am, int stride_ak, int stride_al,
    int stride_bm, int stride_bk, int stride_bl,
    int stride_sfam, int stride_sfak, int stride_sfal,
    int stride_sfbm, int stride_sfbk, int stride_sfbl,
    int stride_cm, int stride_cl
) {
    extern __shared__ char smem[];
    
    // 1. LUT in Shared Memory (256 * 4 bytes = 1KB)
    // Maps byte (2x FP4) -> half2 (2x half decoded)
    // Optimization: Reduced size from float2 (8B) to half2 (4B)
    half2* s_lut_pair = (half2*)smem;

    // 2. Decoded B in Shared Memory
    // Offset: 2048 bytes (aligned safe margin)
    half* s_b_decoded = (half*)(smem + 2048); 
    
    int k_elements = K_bytes * 2;
    int num_chunks = (k_elements + 31) / 32;
    int decoded_size_bytes = num_chunks * 72;
    decoded_size_bytes = (decoded_size_bytes + 15) & ~15; 
    
    uint8_t* s_b_raw = (uint8_t*)((char*)s_b_decoded + decoded_size_bytes);
    uint8_t* s_sfb_raw = (uint8_t*)(s_b_raw + K_bytes);

    int tid = threadIdx.x + threadIdx.y * blockDim.x;
    
    // --- Init Huge LUT (first 256 threads) ---
    if (tid < 256) {
        int lo = tid & 0x0F;
        int hi = (tid >> 4) & 0x0F;
        // Convert to half2 directly during init
        s_lut_pair[tid] = __float22half2_rn(make_float2(C_FP4_LUT_CONST[lo], C_FP4_LUT_CONST[hi]));
    }

    int l_idx = blockIdx.y;
    const uint8_t* B_g = B + l_idx * stride_bl;
    const uint8_t* SFB_g = SFB + l_idx * stride_sfbl;

    int num_threads = WARPS_PER_BLOCK * WARP_SIZE;

    // --- Async Load B & SFB ---
    for (int i = tid * 16; i < K_bytes; i += num_threads * 16) {
        if (i + 16 <= K_bytes) cp_async_16(&s_b_raw[i], &B_g[i]);
        else for(int j=0; j<16 && i+j<K_bytes; ++j) s_b_raw[i+j] = B_g[i+j];
    }
    int sfb_bytes = K_bytes / 8;
    for (int i = tid * 16; i < sfb_bytes; i += num_threads * 16) {
        if (i + 16 <= sfb_bytes) cp_async_16(&s_sfb_raw[i], &SFB_g[i]);
        else for(int j=0; j<16 && i+j<sfb_bytes; ++j) s_sfb_raw[i+j] = SFB_g[i+j];
    }
    
    cp_async_commit();
    cp_async_wait_all();
    __syncthreads();

    // --- Decode B + SFB -> Padded SMEM ---
    for (int k = tid * 32; k < k_elements; k += num_threads * 32) {
        if (k >= k_elements) break;

        uint4 b_vec = *reinterpret_cast<uint4*>(&s_b_raw[k / 2]);
        uint8_t* b_bytes = (uint8_t*)&b_vec;

        int sfb_idx = k / 16;
        float s0 = decode_fp8_fast(s_sfb_raw[sfb_idx]);
        float s1 = decode_fp8_fast(s_sfb_raw[sfb_idx+1]);

        int write_base = (k >> 5) * 36;

        #pragma unroll
        for (int i = 0; i < 32; ++i) {
            uint8_t packed = b_bytes[i/2];
            int nib = (i & 1) ? (packed >> 4) : (packed & 0xF);
            float val = C_FP4_LUT_CONST[nib]; 
            float scale = (i < 16) ? s0 : s1;
            s_b_decoded[write_base + i] = __float2half(val * scale);
        }
    }
    __syncthreads();

    // --- Math Phase ---
    int warp_id = threadIdx.y;
    int lane_id = threadIdx.x;
    int m = blockIdx.x * WARPS_PER_BLOCK + warp_id;

    if (m < M) {
        const uint8_t* A_ptr = A + m * stride_am + l_idx * stride_al;
        const uint8_t* SFA_ptr = SFA + m * stride_sfam + l_idx * stride_sfal;
        
        float acc = 0.0f;
        
        uint4 a_reg;
        ushort sfa_reg;

        // Prologue
        int k_first = lane_id * 32;
        if (k_first < k_elements) {
             a_reg = *reinterpret_cast<const uint4*>(A_ptr + k_first / 2);
             sfa_reg = *reinterpret_cast<const ushort*>(SFA_ptr + k_first / 16);
        }

        // Main Loop
        for (int k_base = 0; k_base < k_elements; k_base += 1024) {
            int k = k_base + lane_id * 32;
            
            // 1. Prefetch
            int k_next = k + 1024;
            uint4 a_next;
            ushort sfa_next;
            bool active_next = (k_next < k_elements);
            
            if (active_next) {
                a_next = *reinterpret_cast<const uint4*>(A_ptr + k_next / 2);
                sfa_next = *reinterpret_cast<const ushort*>(SFA_ptr + k_next / 16);
            }

            // 2. Compute
            if (k < k_elements) {
                float sfa0 = decode_fp8_fast(sfa_reg & 0xFF);
                float sfa1 = decode_fp8_fast(sfa_reg >> 8);
                
                uint8_t* a_bytes = (uint8_t*)&a_reg;
                
                int read_base_idx = (k >> 5) * 36; 
                int2* b_vec_ptr = reinterpret_cast<int2*>(&s_b_decoded[read_base_idx]);
                
                int2 b_vals[8];
                #pragma unroll
                for(int v=0; v<8; ++v) b_vals[v] = b_vec_ptr[v];

                half2* b_h2_ptr = (half2*)b_vals;

                // Accumulators in half2
                half2 sum0_h2 = __float2half2_rn(0.0f);
                half2 sum1_h2 = __float2half2_rn(0.0f);

                // Loop 0-7
                #pragma unroll
                for (int i = 0; i < 8; ++i) {
                    // SMEM Load half2 (4 bytes)
                    half2 va = s_lut_pair[a_bytes[i]];
                    half2 vb = b_h2_ptr[i];
                    // Vectorized FMA
                    sum0_h2 = __hfma2(va, vb, sum0_h2);
                }

                // Loop 8-15
                #pragma unroll
                for (int i = 8; i < 16; ++i) {
                    half2 va = s_lut_pair[a_bytes[i]];
                    half2 vb = b_h2_ptr[i];
                    sum1_h2 = __hfma2(va, vb, sum1_h2);
                }
                
                // Reduction: sum .x and .y components and apply scale
                // We cast to float here to maintain precision for the accumulation across K
                float partial0 = __half2float(sum0_h2.x) + __half2float(sum0_h2.y);
                float partial1 = __half2float(sum1_h2.x) + __half2float(sum1_h2.y);

                acc += partial0 * sfa0 + partial1 * sfa1;
            }
            
            // 3. Update
            if (active_next) {
                a_reg = a_next;
                sfa_reg = sfa_next;
            }
        }

        for (int offset = 16; offset > 0; offset /= 2) 
            acc += __shfl_down_sync(0xffffffff, acc, offset);
        
        if (lane_id == 0) {
            C[m * stride_cm + l_idx * stride_cl] = __float2half(acc);
        }
    }
}

// Helper template to launch kernel with specific WARPS_PER_BLOCK
template<int WARPS_PER_BLOCK>
void launch_kernel_with_config(
    const uint8_t* A, const uint8_t* B, const uint8_t* SFA, const uint8_t* SFB, half* C,
    int M, int K_bytes, int L,
    int stride_am, int stride_ak, int stride_al,
    int stride_bm, int stride_bk, int stride_bl,
    int stride_sfam, int stride_sfak, int stride_sfal,
    int stride_sfbm, int stride_sfbk, int stride_sfbl,
    int stride_cm, int stride_cl,
    size_t smem_bytes
) {
    dim3 block(32, WARPS_PER_BLOCK, 1);
    dim3 grid((M + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK, L, 1);

    fp4_gemv_b200_sol<WARPS_PER_BLOCK><<<grid, block, smem_bytes>>>(
        A, B, SFA, SFB, C,
        M, K_bytes, L,
        stride_am, stride_ak, stride_al,
        stride_bm, stride_bk, stride_bl,
        stride_sfam, stride_sfak, stride_sfal,
        stride_sfbm, stride_sfbk, stride_sfbl,
        stride_cm, stride_cl
    );
}

void fp4_gemv_launch(
    torch::Tensor A, torch::Tensor B, torch::Tensor SFA, torch::Tensor SFB, torch::Tensor C,
    int stride_am, int stride_ak, int stride_al,
    int stride_bm, int stride_bk, int stride_bl,
    int stride_sfam, int stride_sfak, int stride_sfal,
    int stride_sfbm, int stride_sfbk, int stride_sfbl,
    int stride_cm, int stride_cl,
    int M, int K_bytes, int L
) {
    int K_elements = K_bytes * 2;
    int num_chunks = (K_elements + 31) / 32;
    size_t size_decoded = num_chunks * 72;
    size_decoded = (size_decoded + 15) & ~15;
    size_t size_raw = K_bytes + (K_bytes/8);

    // SMEM: 2048 (LUT) + Decoded B + Raw + Padding
    size_t smem_bytes = 2048 + size_decoded + size_raw + 256;

    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>();
    half* C_ptr = reinterpret_cast<half*>(C.data_ptr<at::Half>());

    // Dispatch based on shape characteristics
    // Shape 1: M=7168, K=16384, L=1 -> warps=18
    // Shape 2: M=4096, K=7168, L=8 -> warps=16
    // Shape 3: M=7168, K=2048, L=4 -> warps=17

    // Launch appropriate kernel
    if (L == 1) {
        launch_kernel_with_config<18>(
            A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr,
            M, K_bytes, L,
            stride_am, stride_ak, stride_al,
            stride_bm, stride_bk, stride_bl,
            stride_sfam, stride_sfak, stride_sfal,
            stride_sfbm, stride_sfbk, stride_sfbl,
            stride_cm, stride_cl,
            smem_bytes
        );
    } else {
        launch_kernel_with_config<16>(
            A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr,
            M, K_bytes, L,
            stride_am, stride_ak, stride_al,
            stride_bm, stride_bk, stride_bl,
            stride_sfam, stride_sfak, stride_sfal,
            stride_sfbm, stride_sfbk, stride_sfbl,
            stride_cm, stride_cl,
            smem_bytes
        ); 
    }
}
"""


cpp_source = r"""
void fp4_gemv_launch(
    torch::Tensor A, torch::Tensor B, torch::Tensor SFA, torch::Tensor SFB, torch::Tensor C,
    int stride_am, int stride_ak, int stride_al,
    int stride_bm, int stride_bk, int stride_bl,
    int stride_sfam, int stride_sfak, int stride_sfal,
    int stride_sfbm, int stride_sfbk, int stride_sfbl,
    int stride_cm, int stride_cl,
    int M, int K_bytes, int L
);

void fp4_gemv_opt(
    torch::Tensor A, torch::Tensor B, torch::Tensor SFA, torch::Tensor SFB, torch::Tensor C,
    int stride_am, int stride_ak, int stride_al,
    int stride_bm, int stride_bk, int stride_bl,
    int stride_sfam, int stride_sfak, int stride_sfal,
    int stride_sfbm, int stride_sfbk, int stride_sfbl,
    int stride_cm, int stride_cl,
    int M, int K_bytes, int L
) {
    fp4_gemv_launch(A, B, SFA, SFB, C, 
        stride_am, stride_ak, stride_al, 
        stride_bm, stride_bk, stride_bl, 
        stride_sfam, stride_sfak, stride_sfal, 
        stride_sfbm, stride_sfbk, stride_sfbl, 
        stride_cm, stride_cl, 
        M, K_bytes, L);
}
"""

try:
    from task import input_t, output_t
except ImportError:
    import collections
    input_t = collections.namedtuple("input_t", ["a", "b", "sfa", "sfb", "sfa_perm", "sfb_perm", "c"])
    output_t = torch.Tensor

_cuda_module = None

def get_cuda_module():
    global _cuda_module
    if _cuda_module is None:
        _cuda_module = load_inline(
            name="fp4_gemv_b200_sol_v204",
            cpp_sources=cpp_source,
            cuda_sources=cuda_source,
            functions=["fp4_gemv_opt"],
            extra_cuda_cflags=[
                "-O3", 
                "--use_fast_math", 
                "-std=c++17", 
                "-maxrregcount=255", 
                "--generate-line-info"
            ],
        )
    return _cuda_module

def _as_uint8(t: torch.Tensor) -> torch.Tensor:
    if t.dtype == torch.uint8: return t
    return t.view(torch.uint8)

def _check_device(t):
    if not t.is_cuda: raise RuntimeError("All inputs must be on CUDA device")

def custom_kernel(data: input_t) -> output_t:
    a, b, sfa, sfb, _, _, c = data
    _check_device(a); _check_device(b); _check_device(sfa); _check_device(sfb); _check_device(c)

    M, K_bytes, L = a.shape 
    
    A_u8 = _as_uint8(a)
    B_u8 = _as_uint8(b)
    SFA_u8 = _as_uint8(sfa)
    SFB_u8 = _as_uint8(sfb)
    C_view = c.view(M, L)

    sam, sak, sal = map(int, A_u8.stride())
    sbm, sbk, sbl = map(int, B_u8.stride())
    ssfam, ssfak, ssfal = map(int, SFA_u8.stride())
    ssfbm, ssfbk, ssfbl = map(int, SFB_u8.stride())
    scm, scl = map(int, C_view.stride())

    mod = get_cuda_module()
    mod.fp4_gemv_opt(
        A_u8, B_u8, SFA_u8, SFB_u8, C_view,
        sam, sak, sal,
        sbm, sbk, sbl,
        ssfam, ssfak, ssfal,
        ssfbm, ssfbk, ssfbl,
        scm, scl,
        M, K_bytes, L
    )
    
    return c
scrolls · 401 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