submission 70464
Infatoshi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 261 lines, June 9 Researcher Reciprocity License v1.0.
submission_v1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-70464?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:ea31c4d24904509d780466d66acf18ae2c3b09a0422b24a1501b00a6d6801a11
license declaredunknown
license concludedunknown
authorsInfatoshi
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
m.def("nvfp4_gemv", &nvfp4_gemv, "NVFP4 batched GEMV");Kernel source
submission_v1.py261 lines
from functools import lru_cache
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CUDA_SRC = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <cuda_fp16.h>
#include <cstdint>
#include <cmath>
namespace {
constexpr int WARP_SIZE = 32;
constexpr int ROWS_PER_BLOCK = 4;
__device__ __forceinline__ float decode_fp4_e2m1(uint8_t v) {
int sign = (v >> 3) & 0x1;
int exp = (v >> 1) & 0x3;
int mant = v & 0x1;
float value;
if (exp == 0) {
if (mant == 0) {
value = 0.0f;
} else {
float frac = float(mant) * 0.5f;
value = ldexpf(frac, 0); // 2^(1-bias) with bias=1
}
} else {
float frac = 1.0f + float(mant) * 0.5f;
value = ldexpf(frac, exp - 1);
}
return sign ? -value : value;
}
__device__ __forceinline__ float decode_fp8_e4m3(uint8_t v) {
int sign = (v >> 7) & 0x1;
int exp = (v >> 3) & 0xF;
int mant = v & 0x7;
const int bias = 7;
float value;
if (exp == 0) {
if (mant == 0) {
value = 0.0f;
} else {
float frac = float(mant) / 8.0f;
value = ldexpf(frac, 1 - bias);
}
} else if (exp == 0xF) {
// Finite-number saturate: clamp to the largest finite value representable.
float frac = 1.0f + 7.0f / 8.0f;
value = ldexpf(frac, (0xE - bias));
value = ldexpf(value, 1); // multiply by 2 once more to match exp==14, mant==7
} else {
float frac = 1.0f + float(mant) / 8.0f;
value = ldexpf(frac, exp - bias);
}
return sign ? -value : value;
}
struct Strides3 {
int64_t d0;
int64_t d1;
int64_t d2;
};
__global__ void nvfp4_gemv_kernel(
const uint8_t* __restrict__ A,
Strides3 a_stride,
const uint8_t* __restrict__ B,
Strides3 b_stride,
const uint8_t* __restrict__ SFA,
Strides3 sfa_stride,
const uint8_t* __restrict__ SFB,
Strides3 sfb_stride,
half* __restrict__ Out,
Strides3 out_stride,
int M,
int K,
int L,
int packs_per_row,
int blocks_per_row) {
int warp_id = threadIdx.x / WARP_SIZE;
int lane_id = threadIdx.x % WARP_SIZE;
int row = blockIdx.x * ROWS_PER_BLOCK + warp_id;
int batch = blockIdx.z;
if (row >= M || batch >= L) {
return;
}
float accum = 0.0f;
for (int pack_idx = lane_id; pack_idx < packs_per_row; pack_idx += WARP_SIZE) {
// Load packed nvfp4 data for A and B (first N row only).
int64_t a_offset = row * a_stride.d0 + pack_idx * a_stride.d1 + batch * a_stride.d2;
uint8_t packed_a = A[a_offset];
int64_t b_offset = pack_idx * b_stride.d1 + batch * b_stride.d2; // first row -> d0 term omitted
uint8_t packed_b = B[b_offset];
int col0 = pack_idx * 2;
int block0 = col0 >> 4;
int64_t sfa_off0 = row * sfa_stride.d0 + block0 * sfa_stride.d1 + batch * sfa_stride.d2;
int64_t sfb_off0 = block0 * sfb_stride.d1 + batch * sfb_stride.d2;
float scale_a0 = decode_fp8_e4m3(SFA[sfa_off0]);
float scale_b0 = decode_fp8_e4m3(SFB[sfb_off0]);
float base_a0 = decode_fp4_e2m1(packed_a & 0xF);
float base_b0 = decode_fp4_e2m1(packed_b & 0xF);
accum += (base_a0 * scale_a0) * (base_b0 * scale_b0);
int col1 = col0 + 1;
if (col1 < K) {
int block1 = col1 >> 4;
float scale_a1 = scale_a0;
float scale_b1 = scale_b0;
if (block1 != block0) {
int64_t sfa_off1 = row * sfa_stride.d0 + block1 * sfa_stride.d1 + batch * sfa_stride.d2;
int64_t sfb_off1 = block1 * sfb_stride.d1 + batch * sfb_stride.d2;
scale_a1 = decode_fp8_e4m3(SFA[sfa_off1]);
scale_b1 = decode_fp8_e4m3(SFB[sfb_off1]);
}
float base_a1 = decode_fp4_e2m1(packed_a >> 4);
float base_b1 = decode_fp4_e2m1(packed_b >> 4);
accum += (base_a1 * scale_a1) * (base_b1 * scale_b1);
}
}
// Warp reduction
for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {
accum += __shfl_down_sync(0xffffffff, accum, offset);
}
if (lane_id == 0) {
int64_t out_offset = row * out_stride.d0 + batch * out_stride.d2;
Out[out_offset] = __float2half(accum);
}
}
} // namespace
void nvfp4_gemv(torch::Tensor a,
torch::Tensor b,
torch::Tensor sfa,
torch::Tensor sfb,
torch::Tensor out) {
TORCH_CHECK(a.is_cuda(), "A must be on CUDA");
TORCH_CHECK(b.is_cuda(), "B must be on CUDA");
TORCH_CHECK(sfa.is_cuda(), "SFA must be on CUDA");
TORCH_CHECK(sfb.is_cuda(), "SFB must be on CUDA");
TORCH_CHECK(out.is_cuda(), "Output must be on CUDA");
TORCH_CHECK(a.dim() == 3, "A must be [M, K_pack, L]");
TORCH_CHECK(b.dim() == 3, "B must be [N, K_pack, L]");
TORCH_CHECK(out.dim() == 3, "Output must be [M, 1, L]");
TORCH_CHECK(sfa.dim() == 3, "SFA must be [M, K_blk, L]");
TORCH_CHECK(sfb.dim() == 3, "SFB must be [1, K_blk, L]");
int64_t M = a.size(0);
int64_t K = b.size(1) * 2; // packs -> actual columns
int64_t L = a.size(2);
TORCH_CHECK(b.size(1) == a.size(1), "Packed K dimension mismatch");
TORCH_CHECK(b.size(2) == L, "B batch mismatch");
TORCH_CHECK(out.size(0) == M, "Output M mismatch");
TORCH_CHECK(out.size(1) == 1, "Output N must be 1");
TORCH_CHECK(out.size(2) == L, "Output batch mismatch");
TORCH_CHECK(K % 16 == 0, "K must be divisible by 16");
TORCH_CHECK(sfa.size(0) == M, "SFA M mismatch");
TORCH_CHECK(sfa.size(1) * 16 >= K, "SFA K coverage mismatch");
TORCH_CHECK(sfa.size(2) == L, "SFA batch mismatch");
TORCH_CHECK(sfb.size(0) >= 1, "SFB requires at least one row");
TORCH_CHECK(sfb.size(1) * 16 >= K, "SFB K coverage mismatch");
TORCH_CHECK(sfb.size(2) == L, "SFB batch mismatch");
auto stream = at::cuda::getCurrentCUDAStream();
Strides3 a_stride{a.stride(0), a.stride(1), a.stride(2)};
Strides3 b_stride{b.stride(0), b.stride(1), b.stride(2)};
Strides3 sfa_stride{sfa.stride(0), sfa.stride(1), sfa.stride(2)};
Strides3 sfb_stride{sfb.stride(0), sfb.stride(1), sfb.stride(2)};
Strides3 out_stride{out.stride(0), out.stride(1), out.stride(2)};
const uint8_t* ptr_a = reinterpret_cast<const uint8_t*>(a.data_ptr());
const uint8_t* ptr_b = reinterpret_cast<const uint8_t*>(b.data_ptr());
const uint8_t* ptr_sfa = reinterpret_cast<const uint8_t*>(sfa.data_ptr());
const uint8_t* ptr_sfb = reinterpret_cast<const uint8_t*>(sfb.data_ptr());
half* ptr_out = reinterpret_cast<half*>(out.data_ptr<at::Half>());
int packs_per_row = static_cast<int>((K + 1) >> 1);
int blocks_per_row = static_cast<int>((K + 15) >> 4);
dim3 block_dim(WARP_SIZE * ROWS_PER_BLOCK);
dim3 grid_dim((static_cast<int>(M) + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK, 1, static_cast<unsigned>(L));
nvfp4_gemv_kernel<<<grid_dim, block_dim, 0, stream>>>(
ptr_a, a_stride,
ptr_b, b_stride,
ptr_sfa, sfa_stride,
ptr_sfb, sfb_stride,
ptr_out, out_stride,
static_cast<int>(M),
static_cast<int>(K),
static_cast<int>(L),
packs_per_row,
blocks_per_row);
TORCH_CHECK(cudaGetLastError() == cudaSuccess, "nvfp4_gemv kernel launch failed");
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("nvfp4_gemv", &nvfp4_gemv, "NVFP4 batched GEMV");
}
"""
@lru_cache(maxsize=1)
def _load_module():
extra_cuda_cflags = ["-O3", "--use_fast_math"]
return load_inline(
name="nvfp4_gemv_ext",
cpp_sources="",
cuda_sources=CUDA_SRC,
extra_cuda_cflags=extra_cuda_cflags,
)
def _ensure_device_tensor(tensor: torch.Tensor, device: torch.device) -> torch.Tensor:
if tensor.device == device:
return tensor
return tensor.to(device=device, non_blocking=True)
def custom_kernel(data: input_t) -> output_t:
a, b, sfa, sfb, _sfa_perm, _sfb_perm, c = data
device = a.device
if device.type != "cuda":
raise RuntimeError("nvfp4 kernel requires CUDA tensors")
module = _load_module()
a_dev = a
b_dev = b
sfa_dev = _ensure_device_tensor(sfa, device)
sfb_dev = _ensure_device_tensor(sfb, device)
c_dev = c
module.nvfp4_gemv(a_dev, b_dev, sfa_dev, sfb_dev, c_dev)
return c_dev
scrolls · 261 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