Skip to content
KernelIndex
Search⌘K

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
NVFP4 GEMMsuite of 3 cases
NVIDIA B200
4.58ms
#368 of 369
2025-12-19

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 = 32constexpr int BK = 32;
tile-m = 128constexpr int BM = 128;
tile-n = 128constexpr int BN = 128;
vector-width = uint4uint4* 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