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
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.
fp4
return (float)float_e2m1_t::bitcast(bits);fp8
return (float)((__nv_fp8_e4m3*)ptr)[idx];shared-memory
__shared__ uint8_t smemB[TILE_K_BYTES];vector-width = int4
int4 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 cscrolls · 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 linesreturn ["/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 contributionif (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 unrollfor (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 unrollfor (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 linesint 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 linesextra_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