submission 92661
rt11 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 313 lines, June 9 Researcher Reciprocity License v1.0.
v1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-92661?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:3a90ee38db2f86117048b1fa10df9f36cdd83211e4a6578e898690c191e48437
license declaredunknown
license concludedunknown
authorsrt11
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
m.def("run_nvfp4_gemv", &run_nvfp4_gemv, "Warp-coalesced FP4 GEMV");fp8
const cutlass::float_e4m3_t* sfa_ptr,shared-memory
__shared__ __align__(16) __half2 sh_b[kPackedPerChunk];vector-width = uint4
__device__ __forceinline__ uint4 ld_uint4(const uint8_t* ptr) {Kernel source
v1.py313 lines
# Warp-coalesced FP4 GEMV tuned for SM100 (Blackwell)
# v1: map one warp -> one row to fix uncoalesced global loads, keep shared B/SF
# staging; vector-friendly strides and strict alignment for packed FP4.
import os
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
_module = None
# C++ glue
cpp_source = r"""
#include <cstdint>
#include <torch/extension.h>
#include "cutlass/float8.h"
#include "cutlass/half.h"
extern "C" void nvfp4_gemv_kernel_launch(
const uint8_t* a_ptr,
int64_t stride_a_m, int64_t stride_a_k, int64_t stride_a_l,
const uint8_t* b_ptr,
int64_t stride_b_k, int64_t stride_b_l,
const cutlass::float_e4m3_t* sfa_ptr,
int64_t stride_sfa_m, int64_t stride_sfa_k, int64_t stride_sfa_l,
const cutlass::float_e4m3_t* sfb_ptr,
int64_t stride_sfb_k, int64_t stride_sfb_l,
cutlass::half_t* c_ptr,
int64_t stride_c_m, int64_t stride_c_l,
int32_t m, int32_t l, int32_t k_actual, int32_t k_blocks
);
torch::Tensor run_nvfp4_gemv(
torch::Tensor a, torch::Tensor b,
torch::Tensor sfa, torch::Tensor sfb,
torch::Tensor c
) {
constexpr int kSfVecSize = 16;
TORCH_CHECK(a.is_cuda() && b.is_cuda() && sfa.is_cuda() && sfb.is_cuda() && c.is_cuda(),
"All tensors must be on CUDA");
TORCH_CHECK(a.dim() == 3 && b.dim() == 3 && sfa.dim() == 3 && sfb.dim() == 3 && c.dim() == 3,
"All tensors must be 3-dimensional");
const int64_t m = a.size(0);
const int64_t packed_k = a.size(1);
const int64_t l = a.size(2);
TORCH_CHECK(m > 0 && packed_k > 0 && l > 0, "Dimensions must be positive");
TORCH_CHECK(b.size(0) == 1 && b.size(1) == packed_k && b.size(2) == l,
"B must be compacted to N=1 dimension");
TORCH_CHECK(c.size(0) == m && c.size(1) == 1 && c.size(2) == l, "C shape mismatch");
const int64_t k_actual = packed_k * 2;
const int64_t k_blocks = (k_actual + kSfVecSize - 1) / kSfVecSize;
TORCH_CHECK(sfa.size(0) == m && sfa.size(1) == k_blocks && sfa.size(2) == l, "SFA shape mismatch");
TORCH_CHECK(sfb.size(0) == 1 && sfb.size(1) == k_blocks && sfb.size(2) == l,
"SFB must be compacted to N=1 dimension");
const int32_t m32 = static_cast<int32_t>(m);
const int32_t l32 = static_cast<int32_t>(l);
const int32_t k_actual32 = static_cast<int32_t>(k_actual);
const int32_t k_blocks32 = static_cast<int32_t>(k_blocks);
const int64_t a_bytes = a.element_size();
const int64_t b_bytes = b.element_size();
const int64_t sfa_bytes = sfa.element_size();
const int64_t sfb_bytes = sfb.element_size();
const int64_t c_bytes = c.element_size();
nvfp4_gemv_kernel_launch(
reinterpret_cast<const uint8_t*>(a.data_ptr()),
a.stride(0) * a_bytes, a.stride(1) * a_bytes, a.stride(2) * a_bytes,
reinterpret_cast<const uint8_t*>(b.data_ptr()),
b.stride(1) * b_bytes, b.stride(2) * b_bytes,
reinterpret_cast<const cutlass::float_e4m3_t*>(sfa.data_ptr()),
sfa.stride(0) * sfa_bytes, sfa.stride(1) * sfa_bytes, sfa.stride(2) * sfa_bytes,
reinterpret_cast<const cutlass::float_e4m3_t*>(sfb.data_ptr()),
sfb.stride(1) * sfb_bytes, sfb.stride(2) * sfb_bytes,
reinterpret_cast<cutlass::half_t*>(c.data_ptr()),
c.stride(0) * c_bytes, c.stride(2) * c_bytes,
m32, l32, k_actual32, k_blocks32
);
return c;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("run_nvfp4_gemv", &run_nvfp4_gemv, "Warp-coalesced FP4 GEMV");
}
"""
# CUDA kernel
cuda_source = r"""
#include <algorithm>
#include <cstdint>
#include <cuda.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include "cutlass/float8.h"
#include "cutlass/half.h"
#include "cutlass/numeric_conversion.h"
// Tunables
static constexpr int kLanesPerRow = 32; // full warp -> one row (coalesced)
static constexpr int kRowsPerWarp = 1;
static constexpr int kWarpsPerBlock = 8; // 8 rows per block
static constexpr int kThreads = kWarpsPerBlock * 32;
static constexpr int kRowsPerBlock = kRowsPerWarp * kWarpsPerBlock; // 8
static constexpr int kChunkElems = 256; // elements per K tile
static constexpr int kPackedPerChunk= kChunkElems / 2; // bytes
static constexpr int kSfVec = 16;
static constexpr int kSfPerChunk = kChunkElems / kSfVec; // 16
// Packed FP4x2 -> __half2 via CUTLASS converters (safe with ptxas)
__device__ __forceinline__ __half2 fp4x2_to_half2(const uint8_t v) {
cutlass::float_e2m1_t lo = cutlass::float_e2m1_t::bitcast(v & 0xF);
cutlass::float_e2m1_t hi = cutlass::float_e2m1_t::bitcast((v >> 4) & 0xF);
cutlass::NumericConverter<cutlass::half_t, cutlass::float_e2m1_t> conv;
cutlass::half_t h0 = conv(lo);
cutlass::half_t h1 = conv(hi);
return __halves2half2(reinterpret_cast<__half&>(h0), reinterpret_cast<__half&>(h1));
}
__device__ __forceinline__ __half fp8_to_half(const cutlass::float_e4m3_t v) {
cutlass::NumericConverter<cutlass::half_t, cutlass::float_e4m3_t> conv;
cutlass::half_t h = conv(v);
return reinterpret_cast<__half&>(h);
}
// Vectorized byte load for A (16B) to guarantee 128B transactions across warp
__device__ __forceinline__ uint4 ld_uint4(const uint8_t* ptr) {
return *reinterpret_cast<const uint4*>(__builtin_assume_aligned(ptr, 16));
}
extern "C" __global__ void __launch_bounds__(kThreads)
nvfp4_gemv_kernel_fast(
const uint8_t* __restrict__ a_ptr,
int64_t stride_a_m, int64_t stride_a_k, int64_t stride_a_l,
const uint8_t* __restrict__ b_ptr,
int64_t stride_b_k, int64_t stride_b_l,
const cutlass::float_e4m3_t* __restrict__ sfa_ptr,
int64_t stride_sfa_m, int64_t stride_sfa_k, int64_t stride_sfa_l,
const cutlass::float_e4m3_t* __restrict__ sfb_ptr,
int64_t stride_sfb_k, int64_t stride_sfb_l,
cutlass::half_t* __restrict__ c_ptr,
int64_t stride_c_m, int64_t stride_c_l,
int32_t m, int32_t l, int32_t k_actual, int32_t /*k_blocks*/
) {
// Batch selection
const int batch = blockIdx.z;
const uint8_t* a_base = a_ptr + batch * stride_a_l;
const uint8_t* b_base = b_ptr + batch * stride_b_l;
const uint8_t* sfa_base_bytes = reinterpret_cast<const uint8_t*>(sfa_ptr) + batch * stride_sfa_l;
const uint8_t* sfb_base_bytes = reinterpret_cast<const uint8_t*>(sfb_ptr) + batch * stride_sfb_l;
uint8_t* c_base_bytes = reinterpret_cast<uint8_t*>(c_ptr) + batch * stride_c_l;
// Thread layout: one warp per row
const int warp_id = threadIdx.x >> 5; // 0..7
const int lane_id = threadIdx.x & 31; // 0..31
const int block_row_start = blockIdx.x * kRowsPerBlock;
const int global_row = block_row_start + warp_id;
const bool active_row = (global_row < m);
// Shared staging (align to 16B for vector ld/st)
__shared__ __align__(16) __half2 sh_b[kPackedPerChunk];
__shared__ __align__(16) cutlass::float_e4m3_t sh_sfb[kSfPerChunk];
__shared__ __align__(16) cutlass::float_e4m3_t sh_sfa[kRowsPerBlock][kSfPerChunk];
float acc = 0.f;
for (int k_off = 0; k_off < k_actual; k_off += kChunkElems) {
const int remaining = k_actual - k_off;
const int chunk_elems = remaining > kChunkElems ? kChunkElems : remaining;
const int packed_count = (chunk_elems + 1) / 2;
const int sf_count = (chunk_elems + kSfVec - 1) / kSfVec;
// Stage SFB for this chunk (first kThreads threads)
for (int t = threadIdx.x; t < sf_count; t += kThreads) {
const int sf_global = (k_off / kSfVec) + t;
const auto* sf_ptr = reinterpret_cast<const cutlass::float_e4m3_t*>(
sfb_base_bytes + sf_global * stride_sfb_k);
sh_sfb[t] = *sf_ptr;
}
__syncthreads();
// Stage B (packed FP4x2 -> half2, scaled)
for (int t = threadIdx.x; t < packed_count; t += kThreads) {
const int pb_global = (k_off / 2) + t;
const uint8_t packed = *(b_base + pb_global * stride_b_k);
__half2 hb = fp4x2_to_half2(packed);
const int sf_idx = (t * 2) / kSfVec;
const __half sf_h = fp8_to_half(sh_sfb[sf_idx]);
const __half2 sf_h2 = __halves2half2(sf_h, sf_h);
hb = __hmul2(hb, sf_h2);
sh_b[t] = hb;
}
// Stage SFA for rows in this block (coalesced over M)
const int rows_this_block = min(kRowsPerBlock, m - block_row_start);
for (int t = threadIdx.x; t < rows_this_block * sf_count; t += kThreads) {
const int local_row = t / sf_count;
const int sf_local = t - local_row * sf_count;
const int sf_global = (k_off / kSfVec) + sf_local;
const auto* sfa_ptr = reinterpret_cast<const cutlass::float_e4m3_t*>(
sfa_base_bytes + (block_row_start + local_row) * stride_sfa_m + sf_global * stride_sfa_k);
sh_sfa[local_row][sf_local] = *sfa_ptr;
}
__syncthreads();
// Compute for active rows: full-warp contiguous loads -> coalesced
if (active_row) {
const int sf_base = k_off / kSfVec;
for (int pb = lane_id; pb < packed_count; pb += kLanesPerRow) {
const int sf_idx = (pb * 2) / kSfVec;
const __half sfa_h = fp8_to_half(sh_sfa[warp_id][sf_idx]);
const __half2 sfa_h2 = __halves2half2(sfa_h, sfa_h);
const int pb_global = (k_off / 2) + pb;
const uint8_t packed_a = *(a_base + global_row * stride_a_m + pb_global * stride_a_k);
__half2 ha = fp4x2_to_half2(packed_a);
ha = __hmul2(ha, sfa_h2);
const __half2 hb = sh_b[pb];
float2 af = __half22float2(ha);
float2 bf = __half22float2(hb);
acc += af.x * bf.x + af.y * bf.y;
}
}
__syncthreads();
}
// Reduce within warp
for (int offset = 16; offset > 0; offset >>= 1) {
acc += __shfl_down_sync(0xffffffff, acc, offset);
}
if (active_row && lane_id == 0) {
auto* out_ptr = reinterpret_cast<cutlass::half_t*>(c_base_bytes + global_row * stride_c_m);
*out_ptr = cutlass::half_t(acc);
}
}
extern "C" void nvfp4_gemv_kernel_launch(
const uint8_t* a_ptr,
int64_t stride_a_m, int64_t stride_a_k, int64_t stride_a_l,
const uint8_t* b_ptr,
int64_t stride_b_k, int64_t stride_b_l,
const cutlass::float_e4m3_t* sfa_ptr,
int64_t stride_sfa_m, int64_t stride_sfa_k, int64_t stride_sfa_l,
const cutlass::float_e4m3_t* sfb_ptr,
int64_t stride_sfb_k, int64_t stride_sfb_l,
cutlass::half_t* c_ptr,
int64_t stride_c_m, int64_t stride_c_l,
int32_t m, int32_t l, int32_t k_actual, int32_t k_blocks
) {
dim3 grid((m + kRowsPerBlock - 1) / kRowsPerBlock, 1, l);
dim3 block(kThreads);
nvfp4_gemv_kernel_fast<<<grid, block>>>(
a_ptr, stride_a_m, stride_a_k, stride_a_l,
b_ptr, stride_b_k, stride_b_l,
sfa_ptr, stride_sfa_m, stride_sfa_k, stride_sfa_l,
sfb_ptr, stride_sfb_k, stride_sfb_l,
c_ptr, stride_c_m, stride_c_l,
m, l, k_actual, k_blocks
);
}
"""
def _get_module():
"""Compile and cache the CUDA extension."""
global _module
if _module is None:
repo_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", ".."))
cutlass_include = os.path.join(repo_root, "cutlass", "include")
_module = load_inline(
name="nvfp4_gemv_v1_rowwarp",
cpp_sources=[cpp_source],
cuda_sources=[cuda_source],
extra_include_paths=[cutlass_include],
extra_cflags=["-std=c++20", "-O3"],
extra_cuda_cflags=[
"-std=c++20",
"--use_fast_math",
"-O3",
"--expt-relaxed-constexpr",
"-Xptxas=-v",
"-arch=sm_100",
],
verbose=False,
with_cuda=True,
)
return _module
def custom_kernel(data: input_t) -> output_t:
"""
Coalesced FP4 GEMV:
- full warp owns a row (no cross-row mixing) to maximize global coalescing
- shared decode of packed B and per-chunk scales
- FP4x2 -> half2 conversion via cvt.rn.f16x2.e2m1x2
"""
a, b_padded, sfa, sfb_padded, *_rest, c = data
b = b_padded.narrow(0, 0, 1)
sfb = sfb_padded.narrow(0, 0, 1)
module = _get_module()
module.run_nvfp4_gemv(a, b, sfa, sfb, c)
return c
scrolls · 313 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