submission 109296
JB Gage · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 248 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-109296?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
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:542c89de0cbc398f51ababf0ef012e24cfe63a6e4f0ddf6fbe4560cb1c896a29
license declaredunknown
license concludedunknown
authorsJB Gage
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
return (float)float_e2m1_t::bitcast(bits);fp8
const __nv_fp8_e4m3* fp8_ptr = (const __nv_fp8_e4m3*)ptr;shared-memory
__shared__ uint8_t smemA[BLOCK_M * TILE_K_BYTES];tile-m = 128
constexpr int BLOCK_M = 128;Kernel source
submission.py248 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 - OPTIMIZED VERSION
# ==============================================================================
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_fp8(const void* ptr, int idx) {
#ifdef TARGET_B200
const __nv_fp8_e4m3* fp8_ptr = (const __nv_fp8_e4m3*)ptr;
return (float)fp8_ptr[idx];
#else
const half* h_ptr = (const half*)ptr;
return __half2float(h_ptr[idx]);
#endif
}
extern "C" __global__ void __launch_bounds__(128) gemv_kernel(
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 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 global_row = block_row_start + tid;
float acc = 0.0f;
int K_bytes = K / 2;
int num_k_tiles = (K_bytes + TILE_K_BYTES - 1) / TILE_K_BYTES;
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);
// 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];
}
}
}
// 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];
}
__syncthreads();
// COMPUTE - process in groups of 16 elements (8 bytes) per scale
if (row_valid) {
// Process 4 scale groups per tile (64 elements = 4 * 16)
#pragma unroll
for (int sg = 0; sg < 4; ++sg) {
int local_byte_start = sg * 8;
if (local_byte_start >= tile_size) break;
int k_elem = (k_byte_start + local_byte_start) * 2;
int scale_idx = k_elem / 16;
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_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];
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;
}
}
}
}
__syncthreads();
}
if (row_valid) {
C[global_row * stride_c_0 + l_idx * stride_c_2] = __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)
{
constexpr int BLOCK_M = 128;
dim3 block(BLOCK_M);
dim3 grid((m + BLOCK_M - 1) / BLOCK_M, 1, l);
gemv_kernel<<<grid, block>>>(
(const uint8_t*)a,
(const uint8_t*)b,
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
);
}
"""
# ==============================================================================
# 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_v21',
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 cscrolls · 248 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