submission 114152
JB Gage · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 48 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-114152?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:18ed4596c52d248a62ceda0584a9f150582d95116b799983f4882a3b37296608
license declaredunknown
license concludedunknown
authorsJB Gage
imported2026-08-15
Kernel source
submission.py48 lines
import torch
from typing import TypeVar
input_t = TypeVar("input_t", bound=tuple)
output_t = TypeVar("output_t", bound=torch.Tensor)
def ceil_div(a, b):
return (a + b - 1) // b
def to_blocked(input_matrix):
"""Convert scale factor tensor to blocked format required by torch._scaled_mm"""
rows, cols = input_matrix.shape
n_row_blocks = ceil_div(rows, 128)
n_col_blocks = ceil_div(cols, 4)
padded = input_matrix
blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
return rearranged.flatten()
def custom_kernel(data: input_t) -> output_t:
a_ref, b_ref, sfa_ref_cpu, sfb_ref_cpu, _, _, c_ref = data
_, _, l = c_ref.shape
# Pre-convert all scales to blocked format (CPU)
# This minimizes overhead in the main compute loop
scales_a = [to_blocked(sfa_ref_cpu[:, :, l_idx]) for l_idx in range(l)]
scales_b = [to_blocked(sfb_ref_cpu[:, :, l_idx]) for l_idx in range(l)]
# Batch transfer to GPU
scales_a_gpu = [s.cuda() for s in scales_a]
scales_b_gpu = [s.cuda() for s in scales_b]
# Process each batch using cuBLAS (fastest available FP4 GEMV)
for l_idx in range(l):
res = torch._scaled_mm(
a_ref[:, :, l_idx],
b_ref[:, :, l_idx].transpose(0, 1),
scales_a_gpu[l_idx],
scales_b_gpu[l_idx],
bias=None,
out_dtype=torch.float16,
)
c_ref[:, 0, l_idx] = res[:, 0]
return c_refscrolls · 48 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 109320.
import torch- from torch.utils.cpp_extension import load_inline+ from typing import TypeVar- # ==============================================================================- # CONFIGURATION- # ==============================================================================- TARGET_B200 = True+ input_t = TypeVar("input_t", bound=tuple)+ output_t = TypeVar("output_t", bound=torch.Tensor)- # ==============================================================================- # 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"]+ def ceil_div(a, b):+ return (a + b - 1) // b- # ==============================================================================- # 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;+ def to_blocked(input_matrix):+ """Convert scale factor tensor to blocked format required by torch._scaled_mm"""+ rows, cols = input_matrix.shape+ n_row_blocks = ceil_div(rows, 128)+ n_col_blocks = ceil_div(cols, 4)- // 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;+ padded = input_matrix+ blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)+ rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)- // 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+ return rearranged.flatten()++ def custom_kernel(data: input_t) -> output_t:++ a_ref, b_ref, sfa_ref_cpu, sfb_ref_cpu, _, _, c_ref = data+ _, _, l = c_ref.shape- int K_bytes = K / 2;- int num_tiles = (K_bytes + TILE_K_BYTES - 1) / TILE_K_BYTES;+ # Pre-convert all scales to blocked format (CPU)+ # This minimizes overhead in the main compute loop+ scales_a = [to_blocked(sfa_ref_cpu[:, :, l_idx]) for l_idx in range(l)]+ scales_b = [to_blocked(sfb_ref_cpu[:, :, l_idx]) for l_idx in range(l)]- float acc = 0.0f;- bool row_valid = (global_row < M);+ # Batch transfer to GPU+ scales_a_gpu = [s.cuda() for s in scales_a]+ scales_b_gpu = [s.cuda() for s in scales_b]- const uint8_t* row_A = row_valid ? (pA + global_row * stride_a_0) : pA;+ # Process each batch using cuBLAS (fastest available FP4 GEMV)+ for l_idx in range(l):+ res = torch._scaled_mm(+ a_ref[:, :, l_idx],+ b_ref[:, :, l_idx].transpose(0, 1),+ scales_a_gpu[l_idx],+ scales_b_gpu[l_idx],+ bias=None,+ out_dtype=torch.float16,+ )+ c_ref[:, 0, l_idx] = res[:, 0]- 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 cNo newline at end of file+ return c_refNo newline at end of file
scrolls · 291 diff lines total
Best evidence level for this revision: reported
JSON