submission 110187
revess · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 293 lines, June 9 Researcher Reciprocity License v1.0.
submission2_staging.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-110187?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:594ee22cb0b8af8e685bb7370fa5a9dede6509f3b301a3d8a94bcdbce8e78ac8
license declaredunknown
license concludedunknown
authorsrevess
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__global__ void __launch_bounds__(256) nvfp4_gemv_smem_lut(vector-width = half2
half2* smem_b = (half2*)(smem_buffer + 1024);Kernel source
submission2_staging.py293 lines
import torch
import math
import struct
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
# -------------------------------------------------------------------------
# 1. Precompute Lookup Tables
# -------------------------------------------------------------------------
def float_to_half_bits(f):
return struct.unpack('H', struct.pack('e', f))[0]
def create_fp4_lut_hex():
values = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0]
lut_vals = []
for i in range(16):
sign = -1.0 if i >= 8 else 1.0
idx = i if i < 8 else i - 8
lut_vals.append(sign * values[idx])
cpp_array = []
for byte_val in range(256):
lo = lut_vals[byte_val & 0x0F]
hi = lut_vals[(byte_val >> 4) & 0x0F]
packed = (float_to_half_bits(hi) << 16) | float_to_half_bits(lo)
cpp_array.append(f"0x{packed:08X}")
return "{" + ",".join(cpp_array) + "}"
def create_fp8_lut():
vals = []
for i in range(256):
sign = (i >> 7) & 1
exp = (i >> 3) & 0xF
mant = i & 0x7
val = 0.0
if exp == 0:
val = (mant / 8.0) * (2 ** -6) if mant != 0 else 0.0
elif exp == 15:
val = float('nan')
else:
val = (1.0 + mant / 8.0) * (2 ** (exp - 7))
if sign: val = -val
if math.isnan(val):
vals.append("NAN")
else:
vals.append(f"{val}f")
return "{" + ",".join(vals) + "}"
FP4_LUT_HEX = create_fp4_lut_hex()
FP8_LUT_STR = create_fp8_lut()
# -------------------------------------------------------------------------
# 2. CUDA Kernel Source
# -------------------------------------------------------------------------
cuda_source = f'''
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cstdint>
#include <cmath>
__constant__ float FP8_LUT_HOST[256] = {FP8_LUT_STR};
__constant__ uint32_t FP4_LUT_CONST[256] = {FP4_LUT_HEX};
#define WARP_SIZE 32
#define WARPS_PER_BLOCK 8
#define THREADS_PER_BLOCK 256
// Padding config to avoid bank conflicts for B
#define CHUNK_STRIDE 9
__global__ void __launch_bounds__(256) nvfp4_gemv_smem_lut(
const uint8_t* __restrict__ A,
const uint8_t* __restrict__ B,
const uint8_t* __restrict__ SFA,
const uint8_t* __restrict__ SFB,
half* __restrict__ C,
const int M,
const int K_packed,
const int K_sf,
const int stride_a_l,
const int stride_b_l,
const int stride_sfa_l,
const int stride_sfb_l
) {{
// Shared Memory Layout:
// 1. FP4 LUT (256 * 4 bytes = 1024 bytes)
// 2. B Matrix (Dynamic size, padded)
extern __shared__ char smem_buffer[];
uint32_t* smem_lut = (uint32_t*)smem_buffer;
half2* smem_b = (half2*)(smem_buffer + 1024);
const int tid = threadIdx.x;
const int l_idx = blockIdx.z;
const int warp_id = tid / WARP_SIZE;
const int lane_id = tid % WARP_SIZE;
const int row_idx = blockIdx.x * WARPS_PER_BLOCK + warp_id;
const uint8_t* B_global = B + l_idx * stride_b_l;
const uint8_t* SFB_global = SFB + l_idx * stride_sfb_l;
// ----------------------------------------------------------------
// 1. Init Shared Memory LUT & Decode B
// ----------------------------------------------------------------
// Copy LUT to SMEM (Coalesced, 256 threads exactly fill it)
if (tid < 256) {{
smem_lut[tid] = FP4_LUT_CONST[tid];
}}
// No sync needed yet if we rely on warp sync, but let's be safe.
// Actually, we process B below, which needs LUT if we decode on the fly.
// But we are decoding B using the *same* LUT.
// Wait, B-decoding uses the LUT too.
// So we must sync after loading LUT.
__syncthreads();
// Decode B
for (int i = tid; i < K_sf; i += THREADS_PER_BLOCK) {{
uint8_t sfb_raw = SFB_global[i];
float scale = FP8_LUT_HOST[sfb_raw];
half2 h_scale = __float2half2_rn(scale);
uint2 b_pack = *reinterpret_cast<const uint2*>(B_global + i * 8);
uint8_t* b_bytes = (uint8_t*)&b_pack;
half2* dst_chunk = smem_b + (i * CHUNK_STRIDE);
#pragma unroll
for (int j = 0; j < 8; ++j) {{
// Use SMEM LUT for B decoding too
uint32_t lut = smem_lut[b_bytes[j]];
half2 val = *reinterpret_cast<half2*>(&lut);
dst_chunk[j] = __hmul2(val, h_scale);
}}
}}
__syncthreads(); // B and LUT are ready
// ----------------------------------------------------------------
// 2. Compute Row
// ----------------------------------------------------------------
if (row_idx < M) {{
const uint8_t* A_base = A + l_idx * stride_a_l + row_idx * K_packed;
const uint8_t* SFA_base = SFA + l_idx * stride_sfa_l + row_idx * K_sf;
float row_acc = 0.0f;
int k_packed_end = K_packed;
// Loop over A (32 bytes / 64 elems per stride, 16 bytes / 32 elems per thread)
for (int k = lane_id * 16; k < k_packed_end; k += WARP_SIZE * 16) {{
int sfa_idx = k >> 3; // k / 8
// Load SFA
uint16_t sfa_pack = *reinterpret_cast<const uint16_t*>(SFA_base + sfa_idx);
float scale0 = FP8_LUT_HOST[sfa_pack & 0xFF];
float scale1 = FP8_LUT_HOST[sfa_pack >> 8];
// Load A
int4 a_pack = *reinterpret_cast<const int4*>(A_base + k);
uint8_t* a_bytes = (uint8_t*)&a_pack;
// Pointers to B
half2* b_ptr0 = smem_b + (sfa_idx * CHUNK_STRIDE);
half2* b_ptr1 = smem_b + ((sfa_idx + 1) * CHUNK_STRIDE);
half2 acc0 = __float2half2_rn(0.0f);
half2 acc1 = __float2half2_rn(0.0f);
// Inner loops using SMEM LUT
#pragma unroll
for (int j = 0; j < 8; ++j) {{
// Random access to SMEM LUT
uint32_t lut = smem_lut[a_bytes[j]];
half2 va = *reinterpret_cast<half2*>(&lut);
acc0 = __hfma2(va, b_ptr0[j], acc0);
}}
#pragma unroll
for (int j = 8; j < 16; ++j) {{
uint32_t lut = smem_lut[a_bytes[j]];
half2 va = *reinterpret_cast<half2*>(&lut);
acc1 = __hfma2(va, b_ptr1[j - 8], acc1);
}}
float2 f0 = __half22float2(acc0);
float2 f1 = __half22float2(acc1);
row_acc += (f0.x + f0.y) * scale0;
row_acc += (f1.x + f1.y) * scale1;
}}
#pragma unroll
for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {{
row_acc += __shfl_down_sync(0xFFFFFFFF, row_acc, offset);
}}
if (lane_id == 0) {{
C[row_idx + l_idx * M] = __float2half(row_acc);
}}
}}
}}
void set_shared_mem_config() {{
cudaFuncSetAttribute(nvfp4_gemv_smem_lut, cudaFuncAttributeMaxDynamicSharedMemorySize, 98304);
}}
torch::Tensor nvfp4_gemv_cuda_launch(
torch::Tensor A,
torch::Tensor B,
torch::Tensor SFA,
torch::Tensor SFB,
torch::Tensor C,
int64_t M,
int64_t K,
int64_t L
) {{
const int K_packed = K / 2;
const int K_sf = K / 16;
int blocks_x = (M + 8 - 1) / 8;
dim3 grid(blocks_x, 1, L);
dim3 block(256);
// SMEM: 1024 bytes (LUT) + B_padded
size_t smem_size = 1024 + K_sf * 36;
nvfp4_gemv_smem_lut<<<grid, block, smem_size>>>(
A.data_ptr<uint8_t>(),
B.data_ptr<uint8_t>(),
SFA.data_ptr<uint8_t>(),
SFB.data_ptr<uint8_t>(),
reinterpret_cast<half*>(C.data_ptr<at::Half>()),
M, K_packed, K_sf,
M * K_packed,
B.size(0) * K_packed,
M * K_sf,
B.size(0) * K_sf
);
return C;
}}
'''
cpp_source = r'''
#include <torch/extension.h>
void set_shared_mem_config();
torch::Tensor nvfp4_gemv_cuda_launch(
torch::Tensor A,
torch::Tensor B,
torch::Tensor SFA,
torch::Tensor SFB,
torch::Tensor C,
int64_t M,
int64_t K,
int64_t L
);
'''
_nvfp4_module = None
def get_nvfp4_module():
global _nvfp4_module
if _nvfp4_module is None:
_nvfp4_module = load_inline(
name='nvfp4_gemv_smem_lut',
cpp_sources=[cpp_source],
cuda_sources=[cuda_source],
functions=['nvfp4_gemv_cuda_launch', 'set_shared_mem_config'],
verbose=False,
extra_cuda_cflags=['-O3', '--use_fast_math', '-lineinfo', '-std=c++17', '--maxrregcount=64']
)
_nvfp4_module.set_shared_mem_config()
return _nvfp4_module
def custom_kernel(data: input_t) -> output_t:
a, b, sfa, sfb, x, y, c = data
M, _, L = c.shape
K = a.shape[1] * 2
a_u8 = a.view(torch.uint8)
b_u8 = b.view(torch.uint8)
sfa_u8 = sfa.view(torch.uint8)
sfb_u8 = sfb.view(torch.uint8)
module = get_nvfp4_module()
module.nvfp4_gemv_cuda_launch(a_u8, b_u8, sfa_u8, sfb_u8, c, M, K, L)
return cscrolls · 293 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