submission 180610
Kazim · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 599 lines, June 9 Researcher Reciprocity License v1.0.
nvfp4_gemm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-180610?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:4ff19783f2cfe7fcc20f878b6ba6343cc30059d41bcd1ed3dab9d8b8082b374a
license declaredunknown
license concludedunknown
authorsKazim
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ uint64_t s_a[BM * BK / 8];tile-k = 32
constexpr int BK = 32;tile-m = 128
constexpr int BM = 128;tile-n = 128
constexpr int BN = 128;vector-width = uint4
uint4* s_a_load_ptr = (uint4*)s_a;Kernel source
nvfp4_gemm.py599 lines
# nvfp4_gemm_onefile.py
import torch
from torch.utils.cpp_extension import load_inline
input_t = tuple[
torch.Tensor, # A: uint8 [M, K_bytes, L]
torch.Tensor, # B: uint8 [N, K_bytes]
torch.Tensor, # SFA: uint8/uint16? packed FP8 scales [M, K_bytes/8, L] (see your harness)
torch.Tensor, # SFB: uint8/uint16? packed FP8 scales [N, K_bytes/8]
torch.Tensor, # unused
torch.Tensor, # unused
torch.Tensor, # C: fp16 [M, N, L] or [M, N] with L==1 depending on harness
]
output_t = torch.Tensor
gemm_cpp = r"""
#include <torch/extension.h>
torch::Tensor cuda_nvfp4_gemm(torch::Tensor A,
torch::Tensor B,
torch::Tensor C,
torch::Tensor SFA,
torch::Tensor SFB,
long long stride_batch_a,
long long stride_batch_b,
long long stride_batch_sfa,
long long stride_batch_sfb,
long long stride_batch_c,
long long stride_row_a,
long long stride_row_b,
long long stride_row_sfa,
long long stride_row_sfb,
long long stride_row_c);
"""
gemm_cuda = r"""
#include <cuda.h>
#include <cuda_runtime.h>
#include <torch/extension.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>
#include <cuda_bf16.h>
// NOTE:
// - This version fixes two common correctness issues:
// (1) true FP32 accumulation (not half2 accumulation)
// (2) expects ALL strides passed in BYTES (Python wrapper below does that)
// 32 bytes of packed FP4 per operand => 32 bytes == 64 fp4 values (2 per byte)
// scales are fp8 e4m3x2 packed (uint16) for two blocks; you use 4 scalar scales total.
__device__ __forceinline__ float block_scaled_fma_32x2fp4_fp32acc(
const uint64_t (&a_regs)[4],
const uint64_t (&b_regs)[4],
const uint16_t (&sfa_regs)[2],
const uint16_t (&sfb_regs)[2])
{
const uint32_t* a_regs_packed = reinterpret_cast<const uint32_t*>(&a_regs);
const uint32_t* b_regs_packed = reinterpret_cast<const uint32_t*>(&b_regs);
float result_f32;
asm volatile(
"{\n"
// Registers
".reg .b8 a0, a1, a2, a3, a4, a5, a6, a7;\n"
".reg .b8 b0, b1, b2, b3, b4, b5, b6, b7;\n"
".reg .f16x2 sfa_f16x2_0, sfa_f16x2_1, sfb_f16x2_0, sfb_f16x2_1;\n"
".reg .f16x2 sf_f16x2_0, sf_f16x2_1;\n"
".reg .f16 lane0, lane1;\n"
// half2 converted from fp4x2 bytes
".reg .f16x2 cvt_a0, cvt_a1, cvt_a2, cvt_a3, cvt_a4, cvt_a5, cvt_a6, cvt_a7;\n"
".reg .f16x2 cvt_b0, cvt_b1, cvt_b2, cvt_b3, cvt_b4, cvt_b5, cvt_b6, cvt_b7;\n"
// FP32 scales
".reg .f32 sc0, sc1, sc2, sc3;\n"
// FP32 accumulators (two lanes)
".reg .f32 acc0, acc1;\n"
".reg .f32 tmp0, tmp1;\n"
// FP32 temps
".reg .f32 a_f0, a_f1, b_f0, b_f1;\n"
"mov.f32 acc0, 0.0;\n"
"mov.f32 acc1, 0.0;\n"
// --- Scales Setup ---
// %17=sfa[0], %18=sfa[1], %19=sfb[0], %20=sfb[1]
"cvt.rn.f16x2.e4m3x2 sfa_f16x2_0, %17;\n"
"cvt.rn.f16x2.e4m3x2 sfa_f16x2_1, %18;\n"
"cvt.rn.f16x2.e4m3x2 sfb_f16x2_0, %19;\n"
"cvt.rn.f16x2.e4m3x2 sfb_f16x2_1, %20;\n"
"mul.rn.f16x2 sf_f16x2_0, sfa_f16x2_0, sfb_f16x2_0;\n"
"mul.rn.f16x2 sf_f16x2_1, sfa_f16x2_1, sfb_f16x2_1;\n"
// sf_f16x2_0 lanes -> sc0, sc1
"mov.b32 {lane0, lane1}, sf_f16x2_0;\n"
"cvt.f32.f16 sc0, lane0;\n"
"cvt.f32.f16 sc1, lane1;\n"
// sf_f16x2_1 lanes -> sc2, sc3
"mov.b32 {lane0, lane1}, sf_f16x2_1;\n"
"cvt.f32.f16 sc2, lane0;\n"
"cvt.f32.f16 sc3, lane1;\n"
// Helper macro-ish: accumulate 8 bytes worth of fp4x2 pairs into tmp0/tmp1 (FP32),
// then scale by scX and add into acc0/acc1.
// =========================
// --- Block 0 (scale sc0) ---
// A: %1,%2 B: %9,%10
// =========================
"mov.f32 tmp0, 0.0;\n"
"mov.f32 tmp1, 0.0;\n"
"mov.b32 {a0,a1,a2,a3}, %1; mov.b32 {a4,a5,a6,a7}, %2;\n"
"mov.b32 {b0,b1,b2,b3}, %9; mov.b32 {b4,b5,b6,b7}, %10;\n"
"cvt.rn.f16x2.e2m1x2 cvt_a0,a0; cvt.rn.f16x2.e2m1x2 cvt_a1,a1; cvt.rn.f16x2.e2m1x2 cvt_a2,a2; cvt.rn.f16x2.e2m1x2 cvt_a3,a3; cvt.rn.f16x2.e2m1x2 cvt_a4,a4; cvt.rn.f16x2.e2m1x2 cvt_a5,a5; cvt.rn.f16x2.e2m1x2 cvt_a6,a6; cvt.rn.f16x2.e2m1x2 cvt_a7,a7;\n"
"cvt.rn.f16x2.e2m1x2 cvt_b0,b0; cvt.rn.f16x2.e2m1x2 cvt_b1,b1; cvt.rn.f16x2.e2m1x2 cvt_b2,b2; cvt.rn.f16x2.e2m1x2 cvt_b3,b3; cvt.rn.f16x2.e2m1x2 cvt_b4,b4; cvt.rn.f16x2.e2m1x2 cvt_b5,b5; cvt.rn.f16x2.e2m1x2 cvt_b6,b6; cvt.rn.f16x2.e2m1x2 cvt_b7,b7;\n"
// For each half2: extract lanes -> FP32 -> fma into tmp0/tmp1
// pair 0
"mov.b32 {lane0,lane1}, cvt_a0; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b0; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
// pair 1
"mov.b32 {lane0,lane1}, cvt_a1; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b1; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
// pair 2
"mov.b32 {lane0,lane1}, cvt_a2; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b2; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
// pair 3
"mov.b32 {lane0,lane1}, cvt_a3; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b3; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
// pair 4
"mov.b32 {lane0,lane1}, cvt_a4; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b4; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
// pair 5
"mov.b32 {lane0,lane1}, cvt_a5; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b5; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
// pair 6
"mov.b32 {lane0,lane1}, cvt_a6; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b6; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
// pair 7
"mov.b32 {lane0,lane1}, cvt_a7; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b7; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
// scale by sc0 and add to acc
"mul.f32 tmp0, tmp0, sc0;\n"
"mul.f32 tmp1, tmp1, sc0;\n"
"add.f32 acc0, acc0, tmp0;\n"
"add.f32 acc1, acc1, tmp1;\n"
// =========================
// --- Block 1 (scale sc1) ---
// A: %3,%4 B: %11,%12
// =========================
"mov.f32 tmp0, 0.0;\n"
"mov.f32 tmp1, 0.0;\n"
"mov.b32 {a0,a1,a2,a3}, %3; mov.b32 {a4,a5,a6,a7}, %4;\n"
"mov.b32 {b0,b1,b2,b3}, %11; mov.b32 {b4,b5,b6,b7}, %12;\n"
"cvt.rn.f16x2.e2m1x2 cvt_a0,a0; cvt.rn.f16x2.e2m1x2 cvt_a1,a1; cvt.rn.f16x2.e2m1x2 cvt_a2,a2; cvt.rn.f16x2.e2m1x2 cvt_a3,a3; cvt.rn.f16x2.e2m1x2 cvt_a4,a4; cvt.rn.f16x2.e2m1x2 cvt_a5,a5; cvt.rn.f16x2.e2m1x2 cvt_a6,a6; cvt.rn.f16x2.e2m1x2 cvt_a7,a7;\n"
"cvt.rn.f16x2.e2m1x2 cvt_b0,b0; cvt.rn.f16x2.e2m1x2 cvt_b1,b1; cvt.rn.f16x2.e2m1x2 cvt_b2,b2; cvt.rn.f16x2.e2m1x2 cvt_b3,b3; cvt.rn.f16x2.e2m1x2 cvt_b4,b4; cvt.rn.f16x2.e2m1x2 cvt_b5,b5; cvt.rn.f16x2.e2m1x2 cvt_b6,b6; cvt.rn.f16x2.e2m1x2 cvt_b7,b7;\n"
// same 8-pair FP32 accumulate
"mov.b32 {lane0,lane1}, cvt_a0; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b0; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a1; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b1; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a2; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b2; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a3; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b3; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a4; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b4; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a5; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b5; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a6; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b6; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a7; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b7; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mul.f32 tmp0, tmp0, sc1;\n"
"mul.f32 tmp1, tmp1, sc1;\n"
"add.f32 acc0, acc0, tmp0;\n"
"add.f32 acc1, acc1, tmp1;\n"
// =========================
// --- Block 2 (scale sc2) ---
// A: %5,%6 B: %13,%14
// =========================
"mov.f32 tmp0, 0.0;\n"
"mov.f32 tmp1, 0.0;\n"
"mov.b32 {a0,a1,a2,a3}, %5; mov.b32 {a4,a5,a6,a7}, %6;\n"
"mov.b32 {b0,b1,b2,b3}, %13; mov.b32 {b4,b5,b6,b7}, %14;\n"
"cvt.rn.f16x2.e2m1x2 cvt_a0,a0; cvt.rn.f16x2.e2m1x2 cvt_a1,a1; cvt.rn.f16x2.e2m1x2 cvt_a2,a2; cvt.rn.f16x2.e2m1x2 cvt_a3,a3; cvt.rn.f16x2.e2m1x2 cvt_a4,a4; cvt.rn.f16x2.e2m1x2 cvt_a5,a5; cvt.rn.f16x2.e2m1x2 cvt_a6,a6; cvt.rn.f16x2.e2m1x2 cvt_a7,a7;\n"
"cvt.rn.f16x2.e2m1x2 cvt_b0,b0; cvt.rn.f16x2.e2m1x2 cvt_b1,b1; cvt.rn.f16x2.e2m1x2 cvt_b2,b2; cvt.rn.f16x2.e2m1x2 cvt_b3,b3; cvt.rn.f16x2.e2m1x2 cvt_b4,b4; cvt.rn.f16x2.e2m1x2 cvt_b5,b5; cvt.rn.f16x2.e2m1x2 cvt_b6,b6; cvt.rn.f16x2.e2m1x2 cvt_b7,b7;\n"
"mov.b32 {lane0,lane1}, cvt_a0; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b0; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a1; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b1; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a2; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b2; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a3; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b3; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a4; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b4; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a5; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b5; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a6; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b6; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a7; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b7; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mul.f32 tmp0, tmp0, sc2;\n"
"mul.f32 tmp1, tmp1, sc2;\n"
"add.f32 acc0, acc0, tmp0;\n"
"add.f32 acc1, acc1, tmp1;\n"
// =========================
// --- Block 3 (scale sc3) ---
// A: %7,%8 B: %15,%16
// =========================
"mov.f32 tmp0, 0.0;\n"
"mov.f32 tmp1, 0.0;\n"
"mov.b32 {a0,a1,a2,a3}, %7; mov.b32 {a4,a5,a6,a7}, %8;\n"
"mov.b32 {b0,b1,b2,b3}, %15; mov.b32 {b4,b5,b6,b7}, %16;\n"
"cvt.rn.f16x2.e2m1x2 cvt_a0,a0; cvt.rn.f16x2.e2m1x2 cvt_a1,a1; cvt.rn.f16x2.e2m1x2 cvt_a2,a2; cvt.rn.f16x2.e2m1x2 cvt_a3,a3; cvt.rn.f16x2.e2m1x2 cvt_a4,a4; cvt.rn.f16x2.e2m1x2 cvt_a5,a5; cvt.rn.f16x2.e2m1x2 cvt_a6,a6; cvt.rn.f16x2.e2m1x2 cvt_a7,a7;\n"
"cvt.rn.f16x2.e2m1x2 cvt_b0,b0; cvt.rn.f16x2.e2m1x2 cvt_b1,b1; cvt.rn.f16x2.e2m1x2 cvt_b2,b2; cvt.rn.f16x2.e2m1x2 cvt_b3,b3; cvt.rn.f16x2.e2m1x2 cvt_b4,b4; cvt.rn.f16x2.e2m1x2 cvt_b5,b5; cvt.rn.f16x2.e2m1x2 cvt_b6,b6; cvt.rn.f16x2.e2m1x2 cvt_b7,b7;\n"
"mov.b32 {lane0,lane1}, cvt_a0; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b0; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a1; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b1; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a2; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b2; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a3; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b3; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a4; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b4; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a5; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b5; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a6; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b6; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mov.b32 {lane0,lane1}, cvt_a7; cvt.f32.f16 a_f0,lane0; cvt.f32.f16 a_f1,lane1;\n"
"mov.b32 {lane0,lane1}, cvt_b7; cvt.f32.f16 b_f0,lane0; cvt.f32.f16 b_f1,lane1;\n"
"fma.rn.f32 tmp0, a_f0, b_f0, tmp0;\n"
"fma.rn.f32 tmp1, a_f1, b_f1, tmp1;\n"
"mul.f32 tmp0, tmp0, sc3;\n"
"mul.f32 tmp1, tmp1, sc3;\n"
"add.f32 acc0, acc0, tmp0;\n"
"add.f32 acc1, acc1, tmp1;\n"
// final reduction of the two lanes
"add.f32 %0, acc0, acc1;\n"
"}\n"
: "=f"(result_f32)
: "r"(a_regs_packed[0]), "r"(a_regs_packed[1]), "r"(a_regs_packed[2]), "r"(a_regs_packed[3]),
"r"(a_regs_packed[4]), "r"(a_regs_packed[5]), "r"(a_regs_packed[6]), "r"(a_regs_packed[7]),
"r"(b_regs_packed[0]), "r"(b_regs_packed[1]), "r"(b_regs_packed[2]), "r"(b_regs_packed[3]),
"r"(b_regs_packed[4]), "r"(b_regs_packed[5]), "r"(b_regs_packed[6]), "r"(b_regs_packed[7]),
"h"(sfa_regs[0]), "h"(sfa_regs[1]), "h"(sfb_regs[0]), "h"(sfb_regs[1])
: "memory"
);
return result_f32;
}
// 128x128 Tile, BK=32 (bytes)
constexpr int BM = 128;
constexpr int BN = 128;
constexpr int BK = 32;
constexpr int TM = 8;
constexpr int TN = 8;
__global__ void __launch_bounds__(256)
gemm_kernel_opt(
const void* __restrict__ A,
const void* __restrict__ B,
const void* __restrict__ SFA,
const void* __restrict__ SFB,
void* __restrict__ C,
int M, int N, int K,
long long batch_stride_a,
long long batch_stride_b,
long long batch_stride_sfa,
long long batch_stride_sfb,
long long batch_stride_c,
long long row_stride_a,
long long row_stride_b,
long long row_stride_sfa,
long long row_stride_sfb,
long long row_stride_c
) {
const int bx = blockIdx.x;
const int by = blockIdx.y;
const int bz = blockIdx.z;
const int tid = threadIdx.x;
const int ty = tid / 16;
const int tx = tid % 16;
const char* A_ptr = (const char*)A + bz * batch_stride_a;
const char* B_ptr = (const char*)B + bz * batch_stride_b;
const char* SFA_ptr = (const char*)SFA + bz * batch_stride_sfa;
const char* SFB_ptr = (const char*)SFB + bz * batch_stride_sfb;
// Shared Memory
__shared__ uint64_t s_a[BM * BK / 8];
__shared__ uint64_t s_b[BN * BK / 8];
__shared__ uint32_t s_sfa[BM * BK / 8 / 4];
__shared__ uint32_t s_sfb[BN * BK / 8 / 4];
uint4* s_a_load_ptr = (uint4*)s_a;
uint4* s_b_load_ptr = (uint4*)s_b;
uint16_t* s_sfa_load_ptr = (uint16_t*)s_sfa;
uint16_t* s_sfb_load_ptr = (uint16_t*)s_sfb;
float accum[TM][TN];
#pragma unroll
for (int i=0;i<TM;i++) {
#pragma unroll
for (int j=0;j<TN;j++) accum[i][j] = 0.0f;
}
// Loop over K in units of BYTES
for (int k = 0; k < K; k += BK) {
int t_idx = tid;
// --- Load A (128x32 bytes) ---
int load_row_a = t_idx / 2;
int load_col_bytes_a = (t_idx % 2) * 16;
int global_row_a = by * BM + load_row_a;
if (load_row_a < BM && global_row_a < M) {
const char* src = A_ptr + (size_t)global_row_a * row_stride_a + k + load_col_bytes_a;
s_a_load_ptr[t_idx] = *((const uint4*)src);
} else {
s_a_load_ptr[t_idx] = make_uint4(0, 0, 0, 0);
}
// --- Load B (128x32 bytes) ---
int load_row_b = t_idx / 2;
int load_col_bytes_b = (t_idx % 2) * 16;
int global_row_b = bx * BN + load_row_b;
if (load_row_b < BN && global_row_b < N) {
const char* src = B_ptr + (size_t)global_row_b * row_stride_b + k + load_col_bytes_b;
s_b_load_ptr[t_idx] = *((const uint4*)src);
} else {
s_b_load_ptr[t_idx] = make_uint4(0, 0, 0, 0);
}
// --- Load SFA (128 x 4 bytes per BK=32 chunk, since k/8 advances 4) ---
int sfa_global_row = (by * BM) + (tid / 2);
int sfa_col_offset = (tid % 2) * 2;
if (sfa_global_row < M) {
const char* src_sfa = SFA_ptr + (size_t)sfa_global_row * row_stride_sfa + (k / 8) + sfa_col_offset;
s_sfa_load_ptr[tid] = *((const uint16_t*)src_sfa);
} else {
s_sfa_load_ptr[tid] = 0;
}
// --- Load SFB ---
int sfb_global_row = (bx * BN) + (tid / 2);
if (sfb_global_row < N) {
const char* src_sfb = SFB_ptr + (size_t)sfb_global_row * row_stride_sfb + (k / 8) + sfa_col_offset;
s_sfb_load_ptr[tid] = *((const uint16_t*)src_sfb);
} else {
s_sfb_load_ptr[tid] = 0;
}
__syncthreads();
// --- Compute ---
uint64_t reg_a[TM][4];
uint64_t reg_b[TN][4];
uint16_t reg_sfa[TM][2];
uint16_t reg_sfb[TN][2];
int smem_row_base_a = ty * TM;
#pragma unroll
for(int i=0; i<TM; ++i) {
int row = smem_row_base_a + i;
reg_a[i][0] = s_a[row*4 + 0];
reg_a[i][1] = s_a[row*4 + 1];
reg_a[i][2] = s_a[row*4 + 2];
reg_a[i][3] = s_a[row*4 + 3];
uint32_t scale_pack = s_sfa[row];
reg_sfa[i][0] = scale_pack & 0xFFFF;
reg_sfa[i][1] = scale_pack >> 16;
}
int smem_row_base_b = tx * TN;
#pragma unroll
for(int j=0; j<TN; ++j) {
int row = smem_row_base_b + j;
reg_b[j][0] = s_b[row*4 + 0];
reg_b[j][1] = s_b[row*4 + 1];
reg_b[j][2] = s_b[row*4 + 2];
reg_b[j][3] = s_b[row*4 + 3];
uint32_t scale_pack = s_sfb[row];
reg_sfb[j][0] = scale_pack & 0xFFFF;
reg_sfb[j][1] = scale_pack >> 16;
}
#pragma unroll
for(int i=0; i<TM; ++i) {
#pragma unroll
for(int j=0; j<TN; ++j) {
float res = block_scaled_fma_32x2fp4_fp32acc(
reg_a[i],
reg_b[j],
reg_sfa[i],
reg_sfb[j]
);
accum[i][j] += res;
}
}
__syncthreads();
}
// --- Store ---
__half* C_ptr = (__half*)((char*)C + bz * batch_stride_c);
int row_start = by * BM + ty * TM;
int col_start = bx * BN + tx * TN;
const long long c_row_stride_elems = row_stride_c / (long long)sizeof(__half);
#pragma unroll
for(int i=0; i<TM; ++i) {
int r = row_start + i;
if(r < M) {
#pragma unroll
for(int j=0; j<TN; ++j) {
int c = col_start + j;
if(c < N) {
size_t offset = (size_t)r * (size_t)c_row_stride_elems + (size_t)c;
C_ptr[offset] = __float2half(accum[i][j]);
}
}
}
}
}
torch::Tensor cuda_nvfp4_gemm(torch::Tensor A,
torch::Tensor B,
torch::Tensor C,
torch::Tensor SFA,
torch::Tensor SFB,
long long stride_batch_a,
long long stride_batch_b,
long long stride_batch_sfa,
long long stride_batch_sfb,
long long stride_batch_c,
long long stride_row_a,
long long stride_row_b,
long long stride_row_sfa,
long long stride_row_sfb,
long long stride_row_c)
{
const int M = (int)A.size(0);
const int K = (int)A.size(1); // K is BYTES
const int L = (int)A.size(2);
const int N = (int)B.size(0);
dim3 grid( (N + 127) / 128, (M + 127) / 128, L );
dim3 block(256);
gemm_kernel_opt<<<grid, block>>>(
A.data_ptr(), B.data_ptr(), SFA.data_ptr(), SFB.data_ptr(), C.data_ptr(),
M, N, K,
stride_batch_a, stride_batch_b, stride_batch_sfa, stride_batch_sfb, stride_batch_c,
stride_row_a, stride_row_b, stride_row_sfa, stride_row_sfb, stride_row_c
);
return C;
}
"""
gemm_module = load_inline(
name="nvfp4_gemm_onefile",
cpp_sources=[gemm_cpp],
cuda_sources=[gemm_cuda],
functions=["cuda_nvfp4_gemm"],
extra_cuda_cflags=[
"-std=c++17",
"-O3",
"-use_fast_math",
"-maxrregcount=255",
"-gencode=arch=compute_100a,code=sm_100a",
"--ptxas-options=--gpu-name=sm_100a",
],
verbose=False,
)
def custom_kernel(data: input_t) -> output_t:
a, b, sfa, sfb, _, _, c = data
# IMPORTANT: pass strides in BYTES (kernel uses char* arithmetic everywhere)
def stride_bytes(t: torch.Tensor, dim: int) -> int:
return int(t.stride(dim) * t.element_size())
return gemm_module.cuda_nvfp4_gemm(
a, b, c, sfa, sfb,
stride_bytes(a, 2),
stride_bytes(b, 2),
stride_bytes(sfa, 2),
stride_bytes(sfb, 2),
stride_bytes(c, 2),
stride_bytes(a, 0),
stride_bytes(b, 0),
stride_bytes(sfa, 0),
stride_bytes(sfb, 0),
stride_bytes(c, 0),
)scrolls · 599 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