Skip to content
KernelIndex
Search⌘K

submission 109320

JB Gage · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-109320?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
2.96ms
#674 of 678
2025-11-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:430e1e8675cd1f3e7b400d51a4e670f68b1d199968fd5d2386d464bdd04e0d87
license declaredunknown
license concludedunknown
authorsJB Gage
imported2026-08-15

Techniques

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

fp4return (float)float_e2m1_t::bitcast(bits);
fp8return (float)((__nv_fp8_e4m3*)ptr)[idx];
shared-memory__shared__ uint8_t smemB[TILE_K_BYTES];
vector-width = int4int4 vec = *((const int4*)(row_A + k_byte_start + i));

Kernel source

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

# ==============================================================================
# CONFIGURATION
# ==============================================================================
TARGET_B200 = True

# ==============================================================================
# 1. PATH FINDER
# ==============================================================================
def find_cutlass():
    import os
    if os.path.exists("./cutlass/include"):
        return [os.path.abspath("./cutlass/include"), os.path.abspath("./cutlass/tools/util/include")]
    return ["/opt/cutlass/4.3.0/include", "/opt/cutlass/4.3.0/tools/util/include"]

# ==============================================================================
# 2. CUDA SOURCE - SHARED MEMORY + PIPELINING
# ==============================================================================
cuda_source = r"""
#include <cuda_runtime.h>
#include <cstdint>
#include <cuda_fp16.h>

#ifdef TARGET_B200
#include <cuda_fp8.h>
#include <cutlass/numeric_types.h>
using namespace cutlass;
#endif

__device__ __forceinline__ float unpack_e2m1(uint8_t packed_byte, int which_nibble) {
    uint8_t bits = (which_nibble == 0) ? (packed_byte & 0x0F) : (packed_byte >> 4);
#ifdef TARGET_B200
    return (float)float_e2m1_t::bitcast(bits);
#else
    const float 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
    };
    return lut[bits];
#endif
}

__device__ __forceinline__ float load_scale(const void* ptr, int idx) {
#ifdef TARGET_B200
    return (float)((__nv_fp8_e4m3*)ptr)[idx];
#else
    return __half2float(((const half*)ptr)[idx]);
#endif
}

// Tile size: 64 elements = 32 bytes = 4 scale groups
#define TILE_K_BYTES 32
#define TILE_K_ELEM 64

extern "C" __global__ void __launch_bounds__(128) gemv_kernel_shared(
    const uint8_t* __restrict__ A,
    const uint8_t* __restrict__ B,
    const void* __restrict__ SFA,
    const void* __restrict__ SFB,
    half* __restrict__ C,
    int M, int K,
    long long stride_a_0, long long stride_a_2,
    long long stride_b_2,
    long long stride_sfa_0, long long stride_sfa_2,
    long long stride_sfb_2,
    long long stride_c_0, long long stride_c_2)
{
    int tid = threadIdx.x;
    int block_row_start = blockIdx.x * 128;
    int global_row = block_row_start + tid;
    int batch_idx = blockIdx.z;
    
    // Batch-offset pointers
    const uint8_t* pA = A + batch_idx * stride_a_2;
    const uint8_t* pB = B + batch_idx * stride_b_2;
    const void* pSFA = (const uint8_t*)SFA + batch_idx * stride_sfa_2;
    const void* pSFB = (const uint8_t*)SFB + batch_idx * stride_sfb_2;
    half* pC = C + batch_idx * stride_c_2;
    
    // Shared memory for B vector tile and SFB scales
    __shared__ uint8_t smemB[TILE_K_BYTES];
    __shared__ float smemSFB[4];  // 4 scale factors per tile
    
    int K_bytes = K / 2;
    int num_tiles = (K_bytes + TILE_K_BYTES - 1) / TILE_K_BYTES;
    
    float acc = 0.0f;
    bool row_valid = (global_row < M);
    
    const uint8_t* row_A = row_valid ? (pA + global_row * stride_a_0) : pA;
    
    for (int tile = 0; tile < num_tiles; ++tile) {
        int k_byte_start = tile * TILE_K_BYTES;
        int tile_bytes = min(TILE_K_BYTES, K_bytes - k_byte_start);
        
        // Cooperative load of B into shared memory
        if (tid < TILE_K_BYTES) {
            smemB[tid] = (tid < tile_bytes) ? pB[k_byte_start + tid] : 0;
        }
        
        // Load scale factors for B (4 per tile)
        if (tid < 4) {
            int scale_idx = (k_byte_start * 2) / 16 + tid;
            int max_scales = (K + 15) / 16;
            smemSFB[tid] = (scale_idx < max_scales) ? load_scale(pSFB, scale_idx) : 0.0f;
        }
        
        __syncthreads();
        
        // Each thread computes its row's contribution
        if (row_valid) {
            // Load A data for this tile - use vectorized load if aligned
            uint8_t localA[TILE_K_BYTES];
            
            #pragma unroll
            for (int i = 0; i < TILE_K_BYTES; i += 16) {
                if (i < tile_bytes) {
                    int4 vec = *((const int4*)(row_A + k_byte_start + i));
                    *((int4*)&localA[i]) = vec;
                }
            }
            
            // Process 4 scale groups
            #pragma unroll
            for (int sg = 0; sg < 4; ++sg) {
                int byte_start = sg * 8;
                if (byte_start >= tile_bytes) break;
                
                int scale_idx = (k_byte_start * 2) / 16 + sg;
                float sa = load_scale(pSFA, global_row * stride_sfa_0 + scale_idx);
                float sb = smemSFB[sg];
                float scale = sa * sb;
                
                int bytes_in_group = min(8, tile_bytes - byte_start);
                
                #pragma unroll
                for (int b = 0; b < 8; ++b) {
                    if (b < bytes_in_group) {
                        uint8_t raw_a = localA[byte_start + b];
                        uint8_t raw_b = smemB[byte_start + b];
                        
                        float va0 = unpack_e2m1(raw_a, 0);
                        float vb0 = unpack_e2m1(raw_b, 0);
                        float va1 = unpack_e2m1(raw_a, 1);
                        float vb1 = unpack_e2m1(raw_b, 1);
                        
                        acc += (va0 * vb0 + va1 * vb1) * scale;
                    }
                }
            }
        }
        
        __syncthreads();
    }
    
    if (row_valid) {
        pC[global_row * stride_c_0] = __float2half(acc);
    }
}

extern "C" void launch_gemv(
    void* a, void* b, void* sfa, void* sfb, void* c,
    int m, int k, int l,
    int stride_a_0, int stride_a_1, int stride_a_2,
    int stride_b_0, int stride_b_1, int stride_b_2,
    int stride_sfa_0, int stride_sfa_1, int stride_sfa_2,
    int stride_sfb_0, int stride_sfb_1, int stride_sfb_2,
    int stride_c_0, int stride_c_1, int stride_c_2)
{
    dim3 block(128);
    dim3 grid((m + 127) / 128, 1, l);
    
    gemv_kernel_shared<<<grid, block>>>(
        (const uint8_t*)a,
        (const uint8_t*)b,
        sfa, sfb,
        (half*)c,
        m, k,
        (long long)stride_a_0, (long long)stride_a_2,
        (long long)stride_b_2,
        (long long)stride_sfa_0, (long long)stride_sfa_2,
        (long long)stride_sfb_2,
        (long long)stride_c_0, (long long)stride_c_2
    );
}
"""

# ==============================================================================
# 3. C++ WRAPPER
# ==============================================================================
cpp_source = r"""
#include <torch/extension.h>

extern "C" void launch_gemv(
    void* a, void* b, void* sfa, void* sfb, void* c,
    int m, int k, int l,
    int stride_a_0, int stride_a_1, int stride_a_2,
    int stride_b_0, int stride_b_1, int stride_b_2,
    int stride_sfa_0, int stride_sfa_1, int stride_sfa_2,
    int stride_sfb_0, int stride_sfb_1, int stride_sfb_2,
    int stride_c_0, int stride_c_1, int stride_c_2);

void run_kernel_proxy(
    torch::Tensor a, 
    torch::Tensor b, 
    torch::Tensor sfa, 
    torch::Tensor sfb, 
    torch::Tensor c) 
{
    int m = a.size(0);
    int k = a.size(1) * 2;
    int l = a.size(2);

    launch_gemv(
        a.data_ptr(), b.data_ptr(), sfa.data_ptr(), sfb.data_ptr(), c.data_ptr(),
        m, k, l,
        a.stride(0), a.stride(1), a.stride(2),
        b.stride(0), b.stride(1), b.stride(2),
        sfa.stride(0), sfa.stride(1), sfa.stride(2),
        sfb.stride(0), sfb.stride(1), sfb.stride(2),
        c.stride(0), c.stride(1), c.stride(2)
    );
}
"""

# ==============================================================================
# 4. COMPILE
# ==============================================================================
extra_flags = ['-O3', '-std=c++17', '--use_fast_math', '-lineinfo']
if TARGET_B200:
    extra_flags.append('-DTARGET_B200')

custom_gemv_inline = load_inline(
    name='custom_gemv_v24',
    cpp_sources=cpp_source,
    cuda_sources=cuda_source,
    functions=['run_kernel_proxy'],
    extra_include_paths=find_cutlass(),
    extra_cuda_cflags=extra_flags,
    with_cuda=True
)

# ==============================================================================
# 5. ENTRY POINT
# ==============================================================================
def custom_kernel(data):
    a, b, sfa, sfb, _, _, c = data
    custom_gemv_inline.run_kernel_proxy(a, b, sfa, sfb, c)
    return c
scrolls · 251 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 109296.

⋯ 15 unchanged lines
return ["/opt/cutlass/4.3.0/include", "/opt/cutlass/4.3.0/tools/util/include"]
# ==============================================================================
- # 2. CUDA SOURCE - OPTIMIZED VERSION
+ # 2. CUDA SOURCE - SHARED MEMORY + PIPELINING
# ==============================================================================
cuda_source = r"""
#include <cuda_runtime.h>
⋯ 19 unchanged lines
#endif
}
- __device__ __forceinline__ float load_scale_fp8(const void* ptr, int idx) {
+ __device__ __forceinline__ float load_scale(const void* ptr, int idx) {
#ifdef TARGET_B200
- const __nv_fp8_e4m3* fp8_ptr = (const __nv_fp8_e4m3*)ptr;
- return (float)fp8_ptr[idx];
+ return (float)((__nv_fp8_e4m3*)ptr)[idx];
#else
- const half* h_ptr = (const half*)ptr;
- return __half2float(h_ptr[idx]);
+ return __half2float(((const half*)ptr)[idx]);
#endif
}
- extern "C" __global__ void __launch_bounds__(128) gemv_kernel(
+ // Tile size: 64 elements = 32 bytes = 4 scale groups
+ #define TILE_K_BYTES 32
+ #define TILE_K_ELEM 64
+
+ extern "C" __global__ void __launch_bounds__(128) gemv_kernel_shared(
const uint8_t* __restrict__ A,
const uint8_t* __restrict__ B,
const void* __restrict__ SFA,
const void* __restrict__ SFB,
half* __restrict__ C,
- int M, int K, int L,
- int stride_a_0, int stride_a_1, int stride_a_2,
- int stride_b_0, int stride_b_1, int stride_b_2,
- int stride_sfa_0, int stride_sfa_1, int stride_sfa_2,
- int stride_sfb_0, int stride_sfb_1, int stride_sfb_2,
- int stride_c_0, int stride_c_1, int stride_c_2)
+ int M, int K,
+ long long stride_a_0, long long stride_a_2,
+ long long stride_b_2,
+ long long stride_sfa_0, long long stride_sfa_2,
+ long long stride_sfb_2,
+ long long stride_c_0, long long stride_c_2)
{
- int l_idx = blockIdx.z;
int tid = threadIdx.x;
-
- constexpr int BLOCK_M = 128;
- constexpr int TILE_K_BYTES = 32; // 64 elements per tile
-
- __shared__ uint8_t smemA[BLOCK_M * TILE_K_BYTES];
- __shared__ uint8_t smemB[TILE_K_BYTES];
-
- int block_row_start = blockIdx.x * BLOCK_M;
+ int block_row_start = blockIdx.x * 128;
int global_row = block_row_start + tid;
+ int batch_idx = blockIdx.z;
- float acc = 0.0f;
+ // Batch-offset pointers
+ const uint8_t* pA = A + batch_idx * stride_a_2;
+ const uint8_t* pB = B + batch_idx * stride_b_2;
+ const void* pSFA = (const uint8_t*)SFA + batch_idx * stride_sfa_2;
+ const void* pSFB = (const uint8_t*)SFB + batch_idx * stride_sfb_2;
+ half* pC = C + batch_idx * stride_c_2;
+ // Shared memory for B vector tile and SFB scales
+ __shared__ uint8_t smemB[TILE_K_BYTES];
+ __shared__ float smemSFB[4]; // 4 scale factors per tile
+
int K_bytes = K / 2;
- int num_k_tiles = (K_bytes + TILE_K_BYTES - 1) / TILE_K_BYTES;
+ int num_tiles = (K_bytes + TILE_K_BYTES - 1) / TILE_K_BYTES;
+ float acc = 0.0f;
bool row_valid = (global_row < M);
-
- const uint8_t* A_l = A + l_idx * stride_a_2;
- const uint8_t* B_l = B + l_idx * stride_b_2;
- const void* SFA_l = (const uint8_t*)SFA + l_idx * stride_sfa_2;
- const void* SFB_l = (const uint8_t*)SFB + l_idx * stride_sfb_2;
-
- for (int k_tile = 0; k_tile < num_k_tiles; ++k_tile) {
- int k_byte_start = k_tile * TILE_K_BYTES;
- int tile_size = min(TILE_K_BYTES, K_bytes - k_byte_start);
+
+ const uint8_t* row_A = row_valid ? (pA + global_row * stride_a_0) : pA;
+
+ for (int tile = 0; tile < num_tiles; ++tile) {
+ int k_byte_start = tile * TILE_K_BYTES;
+ int tile_bytes = min(TILE_K_BYTES, K_bytes - k_byte_start);
- // LOAD A - each thread loads its row's tile
- if (row_valid) {
- #pragma unroll
- for (int i = 0; i < TILE_K_BYTES; ++i) {
- if (i < tile_size) {
- int k_byte = k_byte_start + i;
- smemA[tid * TILE_K_BYTES + i] = A_l[global_row * stride_a_0 + k_byte * stride_a_1];
- }
- }
+ // Cooperative load of B into shared memory
+ if (tid < TILE_K_BYTES) {
+ smemB[tid] = (tid < tile_bytes) ? pB[k_byte_start + tid] : 0;
}
- // LOAD B - first threads load the shared B tile
- if (tid < TILE_K_BYTES && tid < tile_size) {
- int k_byte = k_byte_start + tid;
- smemB[tid] = B_l[k_byte * stride_b_1];
+ // Load scale factors for B (4 per tile)
+ if (tid < 4) {
+ int scale_idx = (k_byte_start * 2) / 16 + tid;
+ int max_scales = (K + 15) / 16;
+ smemSFB[tid] = (scale_idx < max_scales) ? load_scale(pSFB, scale_idx) : 0.0f;
}
__syncthreads();
- // COMPUTE - process in groups of 16 elements (8 bytes) per scale
+ // Each thread computes its row's contribution
if (row_valid) {
- // Process 4 scale groups per tile (64 elements = 4 * 16)
+ // Load A data for this tile - use vectorized load if aligned
+ uint8_t localA[TILE_K_BYTES];
+
#pragma unroll
+ for (int i = 0; i < TILE_K_BYTES; i += 16) {
+ if (i < tile_bytes) {
+ int4 vec = *((const int4*)(row_A + k_byte_start + i));
+ *((int4*)&localA[i]) = vec;
+ }
+ }
+
+ // Process 4 scale groups
+ #pragma unroll
for (int sg = 0; sg < 4; ++sg) {
- int local_byte_start = sg * 8;
- if (local_byte_start >= tile_size) break;
+ int byte_start = sg * 8;
+ if (byte_start >= tile_bytes) break;
- int k_elem = (k_byte_start + local_byte_start) * 2;
- int scale_idx = k_elem / 16;
+ int scale_idx = (k_byte_start * 2) / 16 + sg;
+ float sa = load_scale(pSFA, global_row * stride_sfa_0 + scale_idx);
+ float sb = smemSFB[sg];
+ float scale = sa * sb;
- float sa = load_scale_fp8(SFA_l, global_row * stride_sfa_0 + scale_idx * stride_sfa_1);
- float sb = load_scale_fp8(SFB_l, scale_idx * stride_sfb_1);
- float combined_scale = sa * sb;
+ int bytes_in_group = min(8, tile_bytes - byte_start);
- int bytes_in_group = min(8, tile_size - local_byte_start);
-
#pragma unroll
for (int b = 0; b < 8; ++b) {
if (b < bytes_in_group) {
- int local_byte = local_byte_start + b;
- uint8_t raw_a = smemA[tid * TILE_K_BYTES + local_byte];
- uint8_t raw_b = smemB[local_byte];
-
+ uint8_t raw_a = localA[byte_start + b];
+ uint8_t raw_b = smemB[byte_start + b];
+
float va0 = unpack_e2m1(raw_a, 0);
float vb0 = unpack_e2m1(raw_b, 0);
float va1 = unpack_e2m1(raw_a, 1);
float vb1 = unpack_e2m1(raw_b, 1);
- acc += (va0 * vb0 + va1 * vb1) * combined_scale;
+ acc += (va0 * vb0 + va1 * vb1) * scale;
}
}
}
⋯ 1 unchanged lines
__syncthreads();
}
-
+
if (row_valid) {
- C[global_row * stride_c_0 + l_idx * stride_c_2] = __float2half(acc);
+ pC[global_row * stride_c_0] = __float2half(acc);
}
}
⋯ 6 unchanged lines
int stride_sfb_0, int stride_sfb_1, int stride_sfb_2,
int stride_c_0, int stride_c_1, int stride_c_2)
{
- constexpr int BLOCK_M = 128;
- dim3 block(BLOCK_M);
- dim3 grid((m + BLOCK_M - 1) / BLOCK_M, 1, l);
+ dim3 block(128);
+ dim3 grid((m + 127) / 128, 1, l);
- gemv_kernel<<<grid, block>>>(
+ gemv_kernel_shared<<<grid, block>>>(
(const uint8_t*)a,
(const uint8_t*)b,
- sfa,
- sfb,
+ sfa, sfb,
(half*)c,
- m, k, l,
- stride_a_0, stride_a_1, stride_a_2,
- stride_b_0, stride_b_1, stride_b_2,
- stride_sfa_0, stride_sfa_1, stride_sfa_2,
- stride_sfb_0, stride_sfb_1, stride_sfb_2,
- stride_c_0, stride_c_1, stride_c_2
+ m, k,
+ (long long)stride_a_0, (long long)stride_a_2,
+ (long long)stride_b_2,
+ (long long)stride_sfa_0, (long long)stride_sfa_2,
+ (long long)stride_sfb_2,
+ (long long)stride_c_0, (long long)stride_c_2
);
}
"""
⋯ 44 unchanged lines
extra_flags.append('-DTARGET_B200')
custom_gemv_inline = load_inline(
- name='custom_gemv_v21',
+ name='custom_gemv_v24',
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=['run_kernel_proxy'],
scrolls · 236 diff lines total

Best evidence level for this revision: reported

JSON