submission 88730
msuiche · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 354 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-88730?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:2b2e21875636a4ade0240d38e4d1834c9db012cc0c5c9ccd7eb2cab892b4c8ae
license declaredunknown
license concludedunknown
authorsmsuiche
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Optimized CUDA FP4 GEMV - V3 with PyTorch inline compilation (no cupy dependency)shared-memory
extern __shared__ float reduction_shared_mem[];Kernel source
submission.py354 lines
"""
Optimized CUDA FP4 GEMV - V3 with PyTorch inline compilation (no cupy dependency)
"""
import torch
from task import input_t, output_t
cuda_source = """
#include <cuda_fp16.h>
#include <torch/extension.h>
#ifndef HUGE_VALF
#define HUGE_VALF __int_as_float(0x7f800000)
#endif
__device__ __constant__ float fp4_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
};
__device__ __forceinline__ float fp8_to_float(unsigned char fp8_val) {
unsigned int sign = (fp8_val >> 7) & 1;
unsigned int exp = (fp8_val >> 3) & 0xF;
unsigned int mant = fp8_val & 0x7;
if (exp == 0) {
if (mant == 0) return sign ? -0.0f : 0.0f;
return (sign ? -1.0f : 1.0f) * ldexpf((float)mant / 8.0f, -6);
}
if (exp == 0xF) {
return sign ? -HUGE_VALF : HUGE_VALF;
}
float mantissa = 1.0f + (float)mant / 8.0f;
int exponent = (int)exp - 7;
float result = ldexpf(mantissa, exponent);
return sign ? -result : result;
}
// Decode 4 FP4 values from 2 bytes and accumulate
__device__ __forceinline__ void decode_and_accumulate_4fp4(
unsigned char a_packed, unsigned char b_packed,
float scale_a, float scale_b,
float& acc)
{
unsigned char a_low = a_packed & 0xF;
unsigned char a_high = (a_packed >> 4) & 0xF;
unsigned char b_low = b_packed & 0xF;
unsigned char b_high = (b_packed >> 4) & 0xF;
acc += fp4_lut[a_low] * scale_a * fp4_lut[b_low] * scale_b;
acc += fp4_lut[a_high] * scale_a * fp4_lut[b_high] * scale_b;
}
// Decode 8 FP4 values from 4 bytes (vectorized)
__device__ __forceinline__ void decode_and_accumulate_8fp4(
const unsigned int a_vec, const unsigned int b_vec,
float scale_a, float scale_b,
float& acc)
{
// Extract 4 bytes from each uint32_t
#pragma unroll
for (int i = 0; i < 4; i++) {
unsigned char a_byte = (a_vec >> (i * 8)) & 0xFF;
unsigned char b_byte = (b_vec >> (i * 8)) & 0xFF;
decode_and_accumulate_4fp4(a_byte, b_byte, scale_a, scale_b, acc);
}
}
__device__ float warp_reduce_sum(float val) {
#pragma unroll
for (int offset = 16; offset > 0; offset /= 2) {
val += __shfl_down_sync(0xffffffff, val, offset);
}
return val;
}
__device__ float block_reduce_sum(float val, float* shared) {
int lane = threadIdx.x % 32;
int wid = threadIdx.x / 32;
val = warp_reduce_sum(val);
if (lane == 0) {
shared[wid] = val;
}
__syncthreads();
if (wid == 0) {
val = (threadIdx.x < (blockDim.x + 31) / 32) ? shared[lane] : 0.0f;
val = warp_reduce_sum(val);
}
return val;
}
/**
* Vectorized parallel K-reduction kernel
* Uses uint32_t loads for 4-byte (8 FP4) vectorization
*/
__global__ void fp4_gemv_vectorized_kernel(
const unsigned char* __restrict__ a_ptr,
const unsigned char* __restrict__ b_ptr,
const unsigned char* __restrict__ sfa_ptr,
const unsigned char* __restrict__ sfb_ptr,
half* __restrict__ c_ptr,
int M, int K, int L)
{
extern __shared__ float reduction_shared_mem[];
int K_packed = K / 2;
int K_scales = K / 16;
int m = blockIdx.x;
int batch_idx = blockIdx.y;
if (m >= M || batch_idx >= L) return;
float local_sum = 0.0f;
// Vectorized loop: process 4 bytes (8 FP4) at a time
// Each thread processes multiple 4-byte chunks
for (int k_packed = threadIdx.x * 4; k_packed < K_packed; k_packed += blockDim.x * 4) {
// Check if we can do a vectorized load (need 4 contiguous bytes)
if (k_packed + 3 < K_packed) {
// Load 4 bytes at once using uint32_t
int a_base = m * K_packed * L + k_packed * L + batch_idx;
int b_base = k_packed * L + batch_idx;
// Manual 4-byte load (safer than reinterpret_cast with alignment issues)
unsigned int a_vec = 0, b_vec = 0;
#pragma unroll
for (int i = 0; i < 4; i++) {
a_vec |= ((unsigned int)a_ptr[a_base + i * L]) << (i * 8);
b_vec |= ((unsigned int)b_ptr[b_base + i * L]) << (i * 8);
}
// Get scale factor (1 per 8 packed bytes = 16 FP4)
int scale_idx = k_packed / 8;
if (scale_idx < K_scales) {
int sfa_idx = m * K_scales * L + scale_idx * L + batch_idx;
int sfb_idx = scale_idx * L + batch_idx;
float scale_a = fp8_to_float(sfa_ptr[sfa_idx]);
float scale_b = fp8_to_float(sfb_ptr[sfb_idx]);
decode_and_accumulate_8fp4(a_vec, b_vec, scale_a, scale_b, local_sum);
}
} else {
// Handle remaining elements (scalar)
#pragma unroll
for (int i = 0; i < 4 && k_packed + i < K_packed; i++) {
int a_idx = m * K_packed * L + (k_packed + i) * L + batch_idx;
int b_idx = (k_packed + i) * L + batch_idx;
unsigned char a_packed = a_ptr[a_idx];
unsigned char b_packed = b_ptr[b_idx];
int scale_idx = (k_packed + i) / 8;
if (scale_idx < K_scales) {
int sfa_idx = m * K_scales * L + scale_idx * L + batch_idx;
int sfb_idx = scale_idx * L + batch_idx;
float scale_a = fp8_to_float(sfa_ptr[sfa_idx]);
float scale_b = fp8_to_float(sfb_ptr[sfb_idx]);
decode_and_accumulate_4fp4(a_packed, b_packed, scale_a, scale_b, local_sum);
}
}
}
}
// Reduce across threads
float result = block_reduce_sum(local_sum, reduction_shared_mem);
if (threadIdx.x == 0) {
int c_idx = m * 1 * L + 0 * L + batch_idx;
c_ptr[c_idx] = __float2half(result);
}
}
/**
* Hybrid kernel with shared memory for small K
*/
__global__ void fp4_gemv_hybrid_vectorized_kernel(
const unsigned char* __restrict__ a_ptr,
const unsigned char* __restrict__ b_ptr,
const unsigned char* __restrict__ sfa_ptr,
const unsigned char* __restrict__ sfb_ptr,
half* __restrict__ c_ptr,
int M, int K, int L)
{
extern __shared__ unsigned char shared_mem[];
int K_packed = K / 2;
int K_scales = K / 16;
unsigned char* b_shared = shared_mem;
unsigned char* sfb_shared = shared_mem + K_packed;
int batch_idx = blockIdx.z;
if (batch_idx >= L) return;
// Load B cooperatively
for (int i = threadIdx.x; i < K_packed; i += blockDim.x) {
b_shared[i] = b_ptr[i * L + batch_idx];
}
for (int i = threadIdx.x; i < K_scales; i += blockDim.x) {
sfb_shared[i] = sfb_ptr[i * L + batch_idx];
}
__syncthreads();
// Each thread processes one row
int m = blockIdx.x * blockDim.x + threadIdx.x;
if (m >= M) return;
float acc = 0.0f;
// Vectorized inner loop with shared memory
#pragma unroll 4
for (int k_packed = 0; k_packed < K_packed; k_packed++) {
int a_idx = m * K_packed * L + k_packed * L + batch_idx;
unsigned char a_packed = a_ptr[a_idx];
unsigned char b_packed = b_shared[k_packed];
int scale_idx = k_packed / 8;
if (scale_idx < K_scales) {
int sfa_idx = m * K_scales * L + scale_idx * L + batch_idx;
float scale_a = fp8_to_float(sfa_ptr[sfa_idx]);
float scale_b = fp8_to_float(sfb_shared[scale_idx]);
decode_and_accumulate_4fp4(a_packed, b_packed, scale_a, scale_b, acc);
}
}
int c_idx = m * 1 * L + 0 * L + batch_idx;
c_ptr[c_idx] = __float2half(acc);
}
torch::Tensor fp4_gemv_vectorized_wrapper(
torch::Tensor a,
torch::Tensor b,
torch::Tensor sfa,
torch::Tensor sfb,
torch::Tensor c,
int M, int K, int L)
{
const unsigned char* a_ptr = a.data_ptr<unsigned char>();
const unsigned char* b_ptr = b.data_ptr<unsigned char>();
const unsigned char* sfa_ptr = sfa.data_ptr<unsigned char>();
const unsigned char* sfb_ptr = sfb.data_ptr<unsigned char>();
at::Half* c_ptr = c.data_ptr<at::Half>();
int threads = 256;
dim3 grid(M, L, 1);
dim3 block(threads, 1, 1);
int shared_mem = threads * sizeof(float);
fp4_gemv_vectorized_kernel<<<grid, block, shared_mem>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, reinterpret_cast<half*>(c_ptr), M, K, L
);
return c;
}
torch::Tensor fp4_gemv_hybrid_wrapper(
torch::Tensor a,
torch::Tensor b,
torch::Tensor sfa,
torch::Tensor sfb,
torch::Tensor c,
int M, int K, int L)
{
const uint8_t* a_ptr = a.data_ptr<uint8_t>();
const uint8_t* b_ptr = b.data_ptr<uint8_t>();
const uint8_t* sfa_ptr = sfa.data_ptr<uint8_t>();
const uint8_t* sfb_ptr = sfb.data_ptr<uint8_t>();
at::Half* c_ptr = c.data_ptr<at::Half>();
int K_packed = K / 2;
int K_scales = K / 16;
int threads = 256;
int blocks_m = (M + threads - 1) / threads;
int shared_mem = K_packed + K_scales;
dim3 grid(blocks_m, 1, L);
dim3 block(threads, 1, 1);
fp4_gemv_hybrid_vectorized_kernel<<<grid, block, shared_mem>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, reinterpret_cast<half*>(c_ptr), M, K, L
);
return c;
}
"""
cpp_source = """
torch::Tensor fp4_gemv_vectorized_wrapper(
torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c, int M, int K, int L);
torch::Tensor fp4_gemv_hybrid_wrapper(
torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c, int M, int K, int L);
"""
# Global kernel cache
_kernel_module = None
def custom_kernel(data: input_t) -> output_t:
global _kernel_module
# Compile kernel once
if _kernel_module is None:
import sys
if sys.stdout is None:
class DummyFile:
def write(self, x): pass
def flush(self): pass
sys.stdout = DummyFile()
from torch.utils.cpp_extension import load_inline
_kernel_module = load_inline(
name='fp4_gemv_kernels',
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=['fp4_gemv_vectorized_wrapper', 'fp4_gemv_hybrid_wrapper'],
verbose=False,
extra_cuda_cflags=['-O3', '--use_fast_math', '-std=c++17']
)
a_ref, b_ref, sfa_ref, sfb_ref, sfa_permuted, sfb_permuted, c_ref = data
M, K_packed, L = a_ref.shape
K = K_packed * 2
# Convert tensors
a_uint8 = a_ref.view(torch.uint8).contiguous()
b_uint8 = b_ref.view(torch.uint8).contiguous()
sfa_uint8 = sfa_ref.view(torch.uint8).contiguous()
sfb_uint8 = sfb_ref.view(torch.uint8).contiguous()
c_out = c_ref.contiguous()
# Choose kernel based on problem size
if K >= 4096:
_kernel_module.fp4_gemv_vectorized_wrapper(
a_uint8, b_uint8, sfa_uint8, sfb_uint8, c_out, M, K, L
)
else:
_kernel_module.fp4_gemv_hybrid_wrapper(
a_uint8, b_uint8, sfa_uint8, sfb_uint8, c_out, M, K, L
)
return c_out
scrolls · 354 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