submission 110745
rcmalli · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 356 lines, June 9 Researcher Reciprocity License v1.0.
submission_v9.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-110745?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:f6cf2a1631537d4efc6beb45e33fbcdc323d166a6837193f7cd64e70aeff6c78
license declaredunknown
license concludedunknown
authorsrcmalli
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
NVFP4 GEMV Kernel - v9: Adaptive Dispatch (SMEM B vs Direct)fp8
const __nv_fp8_e4m3* __restrict__ SFA,shared-memory
gemv_kernel_smem_b(vector-width = uint4
uint4 a_vec, uint4 b_vec,Kernel source
submission_v9.py356 lines
"""
NVFP4 GEMV Kernel - v9: Adaptive Dispatch (SMEM B vs Direct)
Combines best of v8 and direct approach:
- SMEM B caching for large K + single batch (Case 1): 26.1µs
- Direct loads for batched cases (Cases 2 & 3): ~49µs, ~20µs
Expected geomean: ~30µs (vs v3's 32.0µs = 6% improvement)
"""
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
cuda_src = r"""
#include <cuda.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <torch/extension.h>
using at::Tensor;
// Process 32 FP4 elements with 2 scale factors using PTX
__device__ __forceinline__
float process_32_fp4_with_scale_ptx(
uint4 a_vec, uint4 b_vec,
uint16_t sfa_packed, uint16_t sfb_packed)
{
float result;
asm volatile(
"{\n"
".reg .b8 a0, a1, a2, a3, a4, a5, a6, a7;\n"
".reg .b8 a8, a9, a10, a11, a12, a13, a14, a15;\n"
".reg .b8 b0, b1, b2, b3, b4, b5, b6, b7;\n"
".reg .b8 b8, b9, b10, b11, b12, b13, b14, b15;\n"
".reg .f16x2 cvtA0, cvtA1, cvtA2, cvtA3, cvtA4, cvtA5, cvtA6, cvtA7;\n"
".reg .f16x2 cvtA8, cvtA9, cvtA10, cvtA11, cvtA12, cvtA13, cvtA14, cvtA15;\n"
".reg .f16x2 cvtB0, cvtB1, cvtB2, cvtB3, cvtB4, cvtB5, cvtB6, cvtB7;\n"
".reg .f16x2 cvtB8, cvtB9, cvtB10, cvtB11, cvtB12, cvtB13, cvtB14, cvtB15;\n"
".reg .f16x2 acc0_0, acc0_1, acc0_2, acc0_3;\n"
".reg .f16x2 acc1_0, acc1_1, acc1_2, acc1_3;\n"
".reg .f16x2 sfA_f16x2, sfB_f16x2, sf_f16x2;\n"
".reg .f16 sf0, sf1, lane0, lane1;\n"
".reg .f32 result_f32, tmp_f32;\n"
"mov.b32 acc0_0, 0;\n" "mov.b32 acc0_1, 0;\n"
"mov.b32 acc0_2, 0;\n" "mov.b32 acc0_3, 0;\n"
"mov.b32 acc1_0, 0;\n" "mov.b32 acc1_1, 0;\n"
"mov.b32 acc1_2, 0;\n" "mov.b32 acc1_3, 0;\n"
"mov.b32 {a0, a1, a2, a3}, %1;\n"
"mov.b32 {a4, a5, a6, a7}, %2;\n"
"mov.b32 {a8, a9, a10, a11}, %3;\n"
"mov.b32 {a12, a13, a14, a15}, %4;\n"
"mov.b32 {b0, b1, b2, b3}, %5;\n"
"mov.b32 {b4, b5, b6, b7}, %6;\n"
"mov.b32 {b8, b9, b10, b11}, %7;\n"
"mov.b32 {b12, b13, b14, b15}, %8;\n"
"cvt.rn.f16x2.e2m1x2 cvtA0, a0;\n" "cvt.rn.f16x2.e2m1x2 cvtA1, a1;\n"
"cvt.rn.f16x2.e2m1x2 cvtA2, a2;\n" "cvt.rn.f16x2.e2m1x2 cvtA3, a3;\n"
"cvt.rn.f16x2.e2m1x2 cvtA4, a4;\n" "cvt.rn.f16x2.e2m1x2 cvtA5, a5;\n"
"cvt.rn.f16x2.e2m1x2 cvtA6, a6;\n" "cvt.rn.f16x2.e2m1x2 cvtA7, a7;\n"
"cvt.rn.f16x2.e2m1x2 cvtA8, a8;\n" "cvt.rn.f16x2.e2m1x2 cvtA9, a9;\n"
"cvt.rn.f16x2.e2m1x2 cvtA10, a10;\n" "cvt.rn.f16x2.e2m1x2 cvtA11, a11;\n"
"cvt.rn.f16x2.e2m1x2 cvtA12, a12;\n" "cvt.rn.f16x2.e2m1x2 cvtA13, a13;\n"
"cvt.rn.f16x2.e2m1x2 cvtA14, a14;\n" "cvt.rn.f16x2.e2m1x2 cvtA15, a15;\n"
"cvt.rn.f16x2.e2m1x2 cvtB0, b0;\n" "cvt.rn.f16x2.e2m1x2 cvtB1, b1;\n"
"cvt.rn.f16x2.e2m1x2 cvtB2, b2;\n" "cvt.rn.f16x2.e2m1x2 cvtB3, b3;\n"
"cvt.rn.f16x2.e2m1x2 cvtB4, b4;\n" "cvt.rn.f16x2.e2m1x2 cvtB5, b5;\n"
"cvt.rn.f16x2.e2m1x2 cvtB6, b6;\n" "cvt.rn.f16x2.e2m1x2 cvtB7, b7;\n"
"cvt.rn.f16x2.e2m1x2 cvtB8, b8;\n" "cvt.rn.f16x2.e2m1x2 cvtB9, b9;\n"
"cvt.rn.f16x2.e2m1x2 cvtB10, b10;\n" "cvt.rn.f16x2.e2m1x2 cvtB11, b11;\n"
"cvt.rn.f16x2.e2m1x2 cvtB12, b12;\n" "cvt.rn.f16x2.e2m1x2 cvtB13, b13;\n"
"cvt.rn.f16x2.e2m1x2 cvtB14, b14;\n" "cvt.rn.f16x2.e2m1x2 cvtB15, b15;\n"
"fma.rn.f16x2 acc0_0, cvtA0, cvtB0, acc0_0;\n"
"fma.rn.f16x2 acc0_0, cvtA1, cvtB1, acc0_0;\n"
"fma.rn.f16x2 acc0_1, cvtA2, cvtB2, acc0_1;\n"
"fma.rn.f16x2 acc0_1, cvtA3, cvtB3, acc0_1;\n"
"fma.rn.f16x2 acc0_2, cvtA4, cvtB4, acc0_2;\n"
"fma.rn.f16x2 acc0_2, cvtA5, cvtB5, acc0_2;\n"
"fma.rn.f16x2 acc0_3, cvtA6, cvtB6, acc0_3;\n"
"fma.rn.f16x2 acc0_3, cvtA7, cvtB7, acc0_3;\n"
"fma.rn.f16x2 acc1_0, cvtA8, cvtB8, acc1_0;\n"
"fma.rn.f16x2 acc1_0, cvtA9, cvtB9, acc1_0;\n"
"fma.rn.f16x2 acc1_1, cvtA10, cvtB10, acc1_1;\n"
"fma.rn.f16x2 acc1_1, cvtA11, cvtB11, acc1_1;\n"
"fma.rn.f16x2 acc1_2, cvtA12, cvtB12, acc1_2;\n"
"fma.rn.f16x2 acc1_2, cvtA13, cvtB13, acc1_2;\n"
"fma.rn.f16x2 acc1_3, cvtA14, cvtB14, acc1_3;\n"
"fma.rn.f16x2 acc1_3, cvtA15, cvtB15, acc1_3;\n"
"add.rn.f16x2 acc0_0, acc0_0, acc0_1;\n"
"add.rn.f16x2 acc0_2, acc0_2, acc0_3;\n"
"add.rn.f16x2 acc0_0, acc0_0, acc0_2;\n"
"add.rn.f16x2 acc1_0, acc1_0, acc1_1;\n"
"add.rn.f16x2 acc1_2, acc1_2, acc1_3;\n"
"add.rn.f16x2 acc1_0, acc1_0, acc1_2;\n"
"cvt.rn.f16x2.e4m3x2 sfA_f16x2, %9;\n"
"cvt.rn.f16x2.e4m3x2 sfB_f16x2, %10;\n"
"mul.rn.f16x2 sf_f16x2, sfA_f16x2, sfB_f16x2;\n"
"mov.b32 {sf0, sf1}, sf_f16x2;\n"
"mov.b32 {lane0, lane1}, acc0_0;\n"
"mul.rn.f16 lane0, lane0, sf0;\n"
"mul.rn.f16 lane1, lane1, sf0;\n"
"add.rn.f16 lane0, lane0, lane1;\n"
"cvt.f32.f16 result_f32, lane0;\n"
"mov.b32 {lane0, lane1}, acc1_0;\n"
"mul.rn.f16 lane0, lane0, sf1;\n"
"mul.rn.f16 lane1, lane1, sf1;\n"
"add.rn.f16 lane0, lane0, lane1;\n"
"cvt.f32.f16 tmp_f32, lane0;\n"
"add.f32 result_f32, result_f32, tmp_f32;\n"
"mov.f32 %0, result_f32;\n"
"}\n"
: "=f"(result)
: "r"(a_vec.x), "r"(a_vec.y), "r"(a_vec.z), "r"(a_vec.w),
"r"(b_vec.x), "r"(b_vec.y), "r"(b_vec.z), "r"(b_vec.w),
"h"(sfa_packed), "h"(sfb_packed)
);
return result;
}
//=============================================================================
// Kernel 1: SMEM B caching (for large K + single batch)
//=============================================================================
template<int WARPS_PER_BLOCK>
__global__ void __launch_bounds__(32 * WARPS_PER_BLOCK, 4)
gemv_kernel_smem_b(
const uint8_t* __restrict__ A,
const uint8_t* __restrict__ B,
const __nv_fp8_e4m3* __restrict__ SFA,
const __nv_fp8_e4m3* __restrict__ SFB,
__half* __restrict__ C,
int M, int K, int L,
int64_t sAm, int64_t sAk, int64_t sAl,
int64_t sBm, int64_t sBk, int64_t sBl,
int64_t sSFAm, int64_t sSFAk, int64_t sSFAl,
int64_t sSFBm, int64_t sSFBk, int64_t sSFBl,
int64_t sCm, int64_t sCn, int64_t sCl
) {
extern __shared__ uint8_t smem[];
const int K_packed = K / 2;
const int num_sf = K / 16;
uint8_t* smem_b = smem;
uint8_t* smem_sfb = smem + K_packed;
const int lane = threadIdx.x;
const int warp_in_blk = threadIdx.y;
const int tid = warp_in_blk * 32 + lane;
const int block_threads = WARPS_PER_BLOCK * 32;
const int row = blockIdx.x * WARPS_PER_BLOCK + warp_in_blk;
const int batch = blockIdx.y;
const uint8_t* B_row = B + (int64_t)batch * sBl;
const __nv_fp8_e4m3* SFB_row = SFB + (int64_t)batch * sSFBl;
// Cooperative load B into SMEM
for (int i = tid * 16; i < K_packed; i += block_threads * 16) {
if (i + 16 <= K_packed) {
*reinterpret_cast<uint4*>(smem_b + i) =
*reinterpret_cast<const uint4*>(B_row + i);
} else {
for (int j = i; j < min(i + 16, K_packed); j++) {
smem_b[j] = B_row[j];
}
}
}
// Cooperative load SFB into SMEM
for (int i = tid * 16; i < num_sf; i += block_threads * 16) {
if (i + 16 <= num_sf) {
*reinterpret_cast<uint4*>(smem_sfb + i) =
*reinterpret_cast<const uint4*>(
reinterpret_cast<const uint8_t*>(SFB_row) + i * sSFBk);
} else {
for (int j = i; j < min(i + 16, num_sf); j++) {
smem_sfb[j] = reinterpret_cast<const uint8_t*>(SFB_row)[j * sSFBk];
}
}
}
__syncthreads();
if (row >= M || batch >= L) return;
const int num_sf_blocks = K / 16;
const int sf_pairs = num_sf_blocks / 2;
float acc = 0.0f;
const uint8_t* A_row = A + (int64_t)row * sAm + (int64_t)batch * sAl;
const __nv_fp8_e4m3* SFA_row = SFA + (int64_t)row * sSFAm + (int64_t)batch * sSFAl;
for (int pair_idx = lane; pair_idx < sf_pairs; pair_idx += 32) {
int k_packed_base = pair_idx * 16;
int sf_idx_base = pair_idx * 2;
uint4 a_vec = *reinterpret_cast<const uint4*>(A_row + k_packed_base);
uint4 b_vec = *reinterpret_cast<const uint4*>(smem_b + k_packed_base);
uint16_t sfa_packed = *reinterpret_cast<const uint16_t*>(
reinterpret_cast<const uint8_t*>(SFA_row) + sf_idx_base * sSFAk);
uint16_t sfb_packed = *reinterpret_cast<const uint16_t*>(smem_sfb + sf_idx_base);
acc += process_32_fp4_with_scale_ptx(a_vec, b_vec, sfa_packed, sfb_packed);
}
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
acc += __shfl_down_sync(0xffffffffu, acc, offset);
}
if (lane == 0) {
int64_t c_idx = (int64_t)row * sCm + (int64_t)0 * sCn + (int64_t)batch * sCl;
C[c_idx] = __float2half(acc);
}
}
//=============================================================================
// Kernel 2: Direct loads (for batched / small K cases)
//=============================================================================
template<int WARPS_PER_BLOCK>
__global__ void gemv_kernel_direct(
const uint8_t* __restrict__ A,
const uint8_t* __restrict__ B,
const __nv_fp8_e4m3* __restrict__ SFA,
const __nv_fp8_e4m3* __restrict__ SFB,
__half* __restrict__ C,
int M, int K, int L,
int64_t sAm, int64_t sAk, int64_t sAl,
int64_t sBm, int64_t sBk, int64_t sBl,
int64_t sSFAm, int64_t sSFAk, int64_t sSFAl,
int64_t sSFBm, int64_t sSFBk, int64_t sSFBl,
int64_t sCm, int64_t sCn, int64_t sCl
) {
const int lane = threadIdx.x;
const int warp_in_blk = threadIdx.y;
const int row = blockIdx.x * WARPS_PER_BLOCK + warp_in_blk;
const int batch = blockIdx.y;
if (row >= M || batch >= L) return;
const int num_sf_blocks = K / 16;
const int sf_pairs = num_sf_blocks / 2;
float acc = 0.0f;
const uint8_t* A_row = A + (int64_t)row * sAm + (int64_t)batch * sAl;
const uint8_t* B_row = B + (int64_t)batch * sBl;
const __nv_fp8_e4m3* SFA_row = SFA + (int64_t)row * sSFAm + (int64_t)batch * sSFAl;
const __nv_fp8_e4m3* SFB_row = SFB + (int64_t)batch * sSFBl;
for (int pair_idx = lane; pair_idx < sf_pairs; pair_idx += 32) {
int k_packed_base = pair_idx * 16;
int sf_idx_base = pair_idx * 2;
uint4 a_vec = *reinterpret_cast<const uint4*>(A_row + k_packed_base);
uint4 b_vec = *reinterpret_cast<const uint4*>(B_row + k_packed_base);
uint16_t sfa_packed = *reinterpret_cast<const uint16_t*>(
reinterpret_cast<const uint8_t*>(SFA_row) + sf_idx_base * sSFAk);
uint16_t sfb_packed = *reinterpret_cast<const uint16_t*>(
reinterpret_cast<const uint8_t*>(SFB_row) + sf_idx_base * sSFBk);
acc += process_32_fp4_with_scale_ptx(a_vec, b_vec, sfa_packed, sfb_packed);
}
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
acc += __shfl_down_sync(0xffffffffu, acc, offset);
}
if (lane == 0) {
int64_t c_idx = (int64_t)row * sCm + (int64_t)0 * sCn + (int64_t)batch * sCl;
C[c_idx] = __float2half(acc);
}
}
//=============================================================================
// Adaptive dispatch
//=============================================================================
void run_gemv(Tensor a, Tensor b, Tensor sfa, Tensor sfb, Tensor c) {
const int64_t M = a.size(0);
const int64_t K_packed = a.size(1);
const int64_t K = K_packed * 2;
const int64_t L = a.size(2);
const uint8_t* A_ptr = reinterpret_cast<const uint8_t*>(a.data_ptr());
const uint8_t* B_ptr = reinterpret_cast<const uint8_t*>(b.data_ptr());
const __nv_fp8_e4m3* SFA = reinterpret_cast<const __nv_fp8_e4m3*>(sfa.data_ptr());
const __nv_fp8_e4m3* SFB = reinterpret_cast<const __nv_fp8_e4m3*>(sfb.data_ptr());
__half* C = reinterpret_cast<__half*>(c.data_ptr());
auto a_str = a.strides();
auto b_str = b.strides();
auto sfa_str = sfa.strides();
auto sfb_str = sfb.strides();
auto c_str = c.strides();
constexpr int WARPS_PER_BLOCK = 4;
dim3 block(32, WARPS_PER_BLOCK, 1);
dim3 grid((M + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK, L, 1);
// Adaptive dispatch: SMEM B for large K + single batch, direct otherwise
// SMEM B benefits when B is reused across many rows in same block
// For batched (L > 1), each batch needs separate B, reducing SMEM benefit
bool use_smem_b = (K >= 12000) && (L == 1);
if (use_smem_b) {
size_t smem_size = K_packed + (K / 16);
gemv_kernel_smem_b<WARPS_PER_BLOCK>
<<<grid, block, smem_size>>>(
A_ptr, B_ptr, SFA, SFB, C,
(int)M, (int)K, (int)L,
a_str[0], a_str[1], a_str[2],
b_str[0], b_str[1], b_str[2],
sfa_str[0], sfa_str[1], sfa_str[2],
sfb_str[0], sfb_str[1], sfb_str[2],
c_str[0], c_str[1], c_str[2]);
} else {
gemv_kernel_direct<WARPS_PER_BLOCK>
<<<grid, block>>>(
A_ptr, B_ptr, SFA, SFB, C,
(int)M, (int)K, (int)L,
a_str[0], a_str[1], a_str[2],
b_str[0], b_str[1], b_str[2],
sfa_str[0], sfa_str[1], sfa_str[2],
sfb_str[0], sfb_str[1], sfb_str[2],
c_str[0], c_str[1], c_str[2]);
}
}
"""
cpp_src = r"""
#include <torch/extension.h>
void run_gemv(torch::Tensor a, torch::Tensor b, torch::Tensor sfa, torch::Tensor sfb, torch::Tensor c);
"""
gemv_lib = load_inline(
name="fp4_gemv_v9_adaptive",
cpp_sources=cpp_src,
cuda_sources=cuda_src,
functions=["run_gemv"],
extra_cuda_cflags=[
"-O3",
"-use_fast_math",
"-gencode=arch=compute_100a,code=sm_100a",
"--expt-relaxed-constexpr",
],
)
def custom_kernel(data: input_t) -> output_t:
a, b, sfa, sfb, sfa_permuted, sfb_permuted, c = data
gemv_lib.run_gemv(a, b, sfa, sfb, c)
return c
scrolls · 356 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