submission 79473
muxfd · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 426 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-79473?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:c2a2a418a1c480c65bc943eb782a21e80ac9cac9ded17fa1b8029308a9e3047e
license declaredunknown
license concludedunknown
authorsmuxfd
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
ab_dtype = cutlass.Float4E2M1FN # FP4 data type for A and Bfp8
__device__ inline float fp8_to_float(__nv_fp8_e4m3 val) {shared-memory
__shared__ float warp_sums[WARPS_PER_ROW];Kernel source
submission.py426 lines
import torch
from task import input_t, output_t
import cutlass
import cutlass.cute as cute
from cutlass.cute.runtime import make_ptr
import cutlass.utils.blockscaled_layout as blockscaled_utils
from torch.utils.cpp_extension import load_inline
# Kernel configuration parameters
mma_tiler_mnk = (128, 1, 128) # Tile sizes for M, N, K dimensions (larger K tile for big-K)
ab_dtype = cutlass.Float4E2M1FN # FP4 data type for A and B
sf_dtype = cutlass.Float8E4M3FN # FP8 data type for scale factors
c_dtype = cutlass.Float16 # FP16 output type
sf_vec_size = 16 # Scale factor block size (16 elements share one scale)
threads_per_cta = 128 # Number of threads per CUDA thread block
# Helper function for ceiling division
def ceil_div(a, b):
return (a + b - 1) // b
# Lightweight CUDA kernel for large-K, single-batch use (avoids CuTe overhead)
nvfp4_gemv_cuda_source = r"""
#include <cuda_fp16.h>
#include <cuda_fp8.h>
__device__ __constant__ float FP4_E2M1_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__ inline float unpack_fp4_low(uint8_t val) {
return FP4_E2M1_LUT[val & 0x0F];
}
__device__ inline float unpack_fp4_high(uint8_t val) {
return FP4_E2M1_LUT[(val >> 4) & 0x0F];
}
__device__ inline float fp8_to_float(__nv_fp8_e4m3 val) {
return float(val);
}
template<int WARPS_PER_ROW>
__global__ void nvfp4_gemv_kernel(
const uint8_t* __restrict__ A,
const uint8_t* __restrict__ B,
const __nv_fp8_e4m3* __restrict__ scale_a,
const __nv_fp8_e4m3* __restrict__ scale_b,
__half* __restrict__ C,
int M,
int K,
int L) {
const int K_packed = K / 2;
const int K_scale = K / 16;
const int WARP_SIZE = 32;
const int THREADS_PER_ROW = WARPS_PER_ROW * WARP_SIZE;
int row = blockIdx.x;
int batch = blockIdx.y;
int tid_in_row = threadIdx.x;
if (row >= M || batch >= L) return;
float sum = 0.0f;
const uint2* A_vec = (const uint2*)(A + (batch * M + row) * K_packed);
const uint2* B_vec = (const uint2*)(B + batch * K_packed);
int num_scale_blocks = K_scale;
const __nv_fp8_e4m3* scale_a_row = scale_a + (batch * M + row) * K_scale;
const __nv_fp8_e4m3* scale_b_row = scale_b + batch * K_scale;
for (int scale_idx = tid_in_row; scale_idx < num_scale_blocks; scale_idx += THREADS_PER_ROW) {
float sa = fp8_to_float(scale_a_row[scale_idx]);
float sb = fp8_to_float(scale_b_row[scale_idx]);
float scale = sa * sb;
int vec_idx = scale_idx;
uint2 a_data = A_vec[vec_idx];
uint2 b_data = B_vec[vec_idx];
uint8_t* a_bytes = (uint8_t*)&a_data;
uint8_t* b_bytes = (uint8_t*)&b_data;
float local_sum = 0.0f;
#pragma unroll
for (int i = 0; i < 8; i++) {
float a0 = unpack_fp4_low(a_bytes[i]);
float b0 = unpack_fp4_low(b_bytes[i]);
float a1 = unpack_fp4_high(a_bytes[i]);
float b1 = unpack_fp4_high(b_bytes[i]);
local_sum += a0 * b0 + a1 * b1;
}
sum += scale * local_sum;
}
#pragma unroll
for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {
sum += __shfl_down_sync(0xffffffff, sum, offset);
}
__shared__ float warp_sums[WARPS_PER_ROW];
int warp_id = tid_in_row / WARP_SIZE;
int lane_id = tid_in_row % WARP_SIZE;
if (lane_id == 0) {
warp_sums[warp_id] = sum;
}
__syncthreads();
if (warp_id == 0 && lane_id < WARPS_PER_ROW) {
sum = warp_sums[lane_id];
#pragma unroll
for (int offset = WARPS_PER_ROW / 2; offset > 0; offset /= 2) {
sum += __shfl_down_sync(0xffffffff, sum, offset);
}
if (lane_id == 0) {
C[batch * M + row] = __float2half(sum);
}
}
}
void nvfp4_gemv_cuda(
torch::Tensor A,
torch::Tensor B,
torch::Tensor scale_a,
torch::Tensor scale_b,
torch::Tensor C,
int M,
int K,
int L) {
int warps_per_row;
if (K >= 8192) {
warps_per_row = 8;
} else if (K >= 4096) {
warps_per_row = 4;
} else if (K >= 1024) {
warps_per_row = 2;
} else {
warps_per_row = 1;
}
int threads_per_block = warps_per_row * 32;
dim3 threads(threads_per_block);
dim3 blocks(M, L);
if (warps_per_row == 8) {
nvfp4_gemv_kernel<8><<<blocks, threads>>>(
reinterpret_cast<const uint8_t*>(A.data_ptr<uint8_t>()),
reinterpret_cast<const uint8_t*>(B.data_ptr<uint8_t>()),
reinterpret_cast<const __nv_fp8_e4m3*>(scale_a.data_ptr<at::Float8_e4m3fn>()),
reinterpret_cast<const __nv_fp8_e4m3*>(scale_b.data_ptr<at::Float8_e4m3fn>()),
reinterpret_cast<__half*>(C.data_ptr<at::Half>()),
M, K, L);
} else if (warps_per_row == 4) {
nvfp4_gemv_kernel<4><<<blocks, threads>>>(
reinterpret_cast<const uint8_t*>(A.data_ptr<uint8_t>()),
reinterpret_cast<const uint8_t*>(B.data_ptr<uint8_t>()),
reinterpret_cast<const __nv_fp8_e4m3*>(scale_a.data_ptr<at::Float8_e4m3fn>()),
reinterpret_cast<const __nv_fp8_e4m3*>(scale_b.data_ptr<at::Float8_e4m3fn>()),
reinterpret_cast<__half*>(C.data_ptr<at::Half>()),
M, K, L);
} else if (warps_per_row == 2) {
nvfp4_gemv_kernel<2><<<blocks, threads>>>(
reinterpret_cast<const uint8_t*>(A.data_ptr<uint8_t>()),
reinterpret_cast<const uint8_t*>(B.data_ptr<uint8_t>()),
reinterpret_cast<const __nv_fp8_e4m3*>(scale_a.data_ptr<at::Float8_e4m3fn>()),
reinterpret_cast<const __nv_fp8_e4m3*>(scale_b.data_ptr<at::Float8_e4m3fn>()),
reinterpret_cast<__half*>(C.data_ptr<at::Half>()),
M, K, L);
} else {
nvfp4_gemv_kernel<1><<<blocks, threads>>>(
reinterpret_cast<const uint8_t*>(A.data_ptr<uint8_t>()),
reinterpret_cast<const uint8_t*>(B.data_ptr<uint8_t>()),
reinterpret_cast<const __nv_fp8_e4m3*>(scale_a.data_ptr<at::Float8_e4m3fn>()),
reinterpret_cast<const __nv_fp8_e4m3*>(scale_b.data_ptr<at::Float8_e4m3fn>()),
reinterpret_cast<__half*>(C.data_ptr<at::Half>()),
M, K, L);
}
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(err));
}
}
"""
nvfp4_gemv_cpp_source = r"""
void nvfp4_gemv_cuda(
torch::Tensor A,
torch::Tensor B,
torch::Tensor scale_a,
torch::Tensor scale_b,
torch::Tensor C,
int M,
int K,
int L);
"""
nvfp4_cuda_module = load_inline(
name='nvfp4_gemv_hybrid_cuda',
cpp_sources=nvfp4_gemv_cpp_source,
cuda_sources=nvfp4_gemv_cuda_source,
functions=['nvfp4_gemv_cuda'],
verbose=False,
extra_cuda_cflags=['-O3', '--use_fast_math'],
)
# The CuTe implementation for NVFP4 block-scaled GEMV
@cute.kernel
def kernel(
mA_mkl: cute.Tensor,
mB_nkl: cute.Tensor,
mSFA_mkl: cute.Tensor,
mSFB_nkl: cute.Tensor,
mC_mnl: cute.Tensor,
):
# Get CUDA block and thread indices
bidx, bidy, bidz = cute.arch.block_idx()
tidx, _, _ = cute.arch.thread_idx()
# Extract the local tile for input matrix A (shape: [block_M, block_K, rest_M, rest_K, rest_L])
gA_mkl = cute.local_tile(
mA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
)
# Extract the local tile for scale factor tensor for A (same shape as gA_mkl)
gSFA_mkl = cute.local_tile(
mSFA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
)
# Extract the local tile for input matrix B (shape: [block_N, block_K, rest_N, rest_K, rest_L])
gB_nkl = cute.local_tile(
mB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
)
# Extract the local tile for scale factor tensor for B (same shape as gB_nkl)
gSFB_nkl = cute.local_tile(
mSFB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
)
# Extract the local tile for output matrix C (shape: [block_M, block_N, rest_M, rest_N, rest_L])
gC_mnl = cute.local_tile(
mC_mnl, cute.slice_(mma_tiler_mnk, (None, None, 0)), (None, None, None)
)
# Select output element corresponding to this thread and block indices
tCgC = gC_mnl[tidx, None, bidx, bidy, bidz]
tCgC = cute.make_tensor(tCgC.iterator, 1)
res = cute.zeros_like(tCgC, cutlass.Float32)
# Get the number of k tiles (depth dimension) for the reduction loop
k_tile_cnt = gA_mkl.layout[3].shape
for k_tile in range(k_tile_cnt):
tAgA = gA_mkl[tidx, None, bidx, k_tile, bidz]
tBgB = gB_nkl[0, None, bidy, k_tile, bidz]
tAgSFA = gSFA_mkl[tidx, None, bidx, k_tile, bidz]
tBgSFB = gSFB_nkl[0, None, bidy, k_tile, bidz]
tArA = cute.make_rmem_tensor_like(tAgA, cutlass.Float32)
tBrB = cute.make_rmem_tensor_like(tBgB, cutlass.Float32)
tArSFA = cute.make_rmem_tensor_like(tAgSFA, cutlass.Float32)
tBrSFB = cute.make_rmem_tensor_like(tBgSFB, cutlass.Float32)
# Load NVFP4 or FP8 values from global memory
a_val_nvfp4 = tAgA.load()
b_val_nvfp4 = tBgB.load()
sfa_val_fp8 = tAgSFA.load()
sfb_val_fp8 = tBgSFB.load()
# Convert loaded values to float32 for computation (FFMA)
a_val = a_val_nvfp4.to(cutlass.Float32)
b_val = b_val_nvfp4.to(cutlass.Float32)
sfa_val = sfa_val_fp8.to(cutlass.Float32)
sfb_val = sfb_val_fp8.to(cutlass.Float32)
# Store the converted values to RMEM CuTe tensors
tArA.store(a_val)
tBrB.store(b_val)
tArSFA.store(sfa_val)
tBrSFB.store(sfb_val)
# Iterate over SF vector tiles and compute the scale&matmul accumulation
for i in cutlass.range_constexpr(mma_tiler_mnk[2]):
res += tArA[i] * tArSFA[i] * tBrB[i] * tBrSFB[i]
# Store the final float16 result back to global memory
tCgC.store(res.to(cutlass.Float16))
return
@cute.jit
def my_kernel(
a_ptr: cute.Pointer,
b_ptr: cute.Pointer,
sfa_ptr: cute.Pointer,
sfb_ptr: cute.Pointer,
c_ptr: cute.Pointer,
problem_size: tuple,
):
"""
Host-side JIT function to prepare tensors and launch GPU kernel.
"""
m, _, k, l = problem_size
# Create CuTe Tensor via pointer and problem size.
a_tensor = cute.make_tensor(
a_ptr,
cute.make_layout(
(m, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m * k, 32)),
),
)
# We use n=128 to create the torch tensor to do fp4 computation via torch._scaled_mm
# then copy torch tensor to cute tensor for cute customize kernel computation
# therefore we need to ensure b_tensor has the right stride with this 128 padded size on n.
n_padded_128 = 128
b_tensor = cute.make_tensor(
b_ptr,
cute.make_layout(
(n_padded_128, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n_padded_128 * k, 32)),
),
)
c_tensor = cute.make_tensor(
c_ptr, cute.make_layout((cute.assume(m, 32), 1, l), stride=(1, 1, m))
)
# Convert scale factor tensors to MMA layout
# The layout matches Tensor Core requirements: (((32, 4), REST_M), ((SF_K, 4), REST_K), (1, REST_L))
sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a_tensor.shape, sf_vec_size)
sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)
sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(b_tensor.shape, sf_vec_size)
sfb_tensor = cute.make_tensor(sfb_ptr, sfb_layout)
# Grid is (M_blocks, 1, L)
grid = (
cute.ceil_div(c_tensor.shape[0], 128),
1,
c_tensor.shape[2],
)
# Launch the CUDA kernel
kernel(a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor).launch(
grid=grid,
block=[threads_per_cta, 1, 1],
cluster=(1, 1, 1),
)
return
# Global cache for compiled kernel
_compiled_kernel_cache = None
# Compile the kernel once and cache it.
def compile_kernel():
global _compiled_kernel_cache
if _compiled_kernel_cache is not None:
return _compiled_kernel_cache
# Create CuTe pointers for A/B/C/SFA/SFB via torch tensor data pointer
a_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
c_ptr = make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
sfa_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
sfb_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
# Compile the kernel
_compiled_kernel_cache = cute.compile(
my_kernel, a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (0, 0, 0, 0)
)
return _compiled_kernel_cache
def custom_kernel(data: input_t) -> output_t:
"""
Hybrid: use CUDA path for large-K single batch; CuTe path otherwise.
"""
a, b, sfa_ref, sfb_ref, sfa_permuted, sfb_permuted, c = data
# Get dimensions from MxKxL layout
m, k_packed, l = a.shape
k = k_packed * 2 # torch uses e2m1_x2 so K is packed by 2
# Large-K single batch: use CUDA kernel (faster for K>=8192, l==1)
if l == 1 and k >= 8192:
a_uint8 = a.contiguous().view(torch.uint8)
b_uint8 = b[0:1, :, :].contiguous().view(torch.uint8)
sfa = sfa_ref.contiguous()
sfb = sfb_ref[0:1, :, :].contiguous()
c_slice = c[:, 0, 0].contiguous()
nvfp4_cuda_module.nvfp4_gemv_cuda(
a_uint8,
b_uint8,
sfa,
sfb,
c_slice,
m,
k,
l,
)
c[:, 0, 0] = c_slice
return c
# Default: CuTe kernel
compiled_func = compile_kernel()
a_ptr = make_ptr(ab_dtype, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
b_ptr = make_ptr(ab_dtype, b.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
c_ptr = make_ptr(c_dtype, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
sfa_ptr = make_ptr(
sf_dtype, sfa_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32
)
sfb_ptr = make_ptr(
sf_dtype, sfb_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32
)
compiled_func(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, 1, k, l))
return c
scrolls · 426 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