Skip to content
KernelIndex
Search⌘K

submission 260929

phuc9702 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 248 lines, June 9 Researcher Reciprocity License v1.0.

nvfp4_dual_gemm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-260929?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 dual GEMMsuite of 4 cases
NVIDIA B200
36.2µs
#269 of 420
2026-01-02

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4de8aa5c9c3f9f0922819ee8555ca64133e72d80a1efd061e9f67082b20c52a9
license declaredunknown
license concludedunknown
authorsphuc9702
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4PyTorch reference implementation of NVFP4 block-scaled dual GEMM with silu activation,
num-warps = 4num_warps = 4
stages = 3num_stages = 3
tile-k = 256BLOCK_K = 256
tile-m = 128BLOCK_M = 128
tile-n = 64BLOCK_N = 64

Kernel source

nvfp4_dual_gemm.py248 lines
import torch
import triton
import triton.language as tl

from task import input_t, output_t
from utils import make_match_reference

# Scaling factor vector size
sf_vec_size = 16


def ceil_div(a, b):
    return (a + b - 1) // b


def to_blocked(input_matrix: torch.Tensor) -> torch.Tensor:
    rows, cols = input_matrix.shape
    n_row_blocks = ceil_div(rows, 128)
    n_col_blocks = ceil_div(cols, 4)
    blocks = input_matrix.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
    rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
    return rearranged.flatten()


def ref_kernel(data: input_t) -> output_t:
    """
    PyTorch reference implementation of NVFP4 block-scaled dual GEMM with silu activation,
    C = silu(A @ B1) * (A @ B2).
    """
    a_ref, b1_ref, b2_ref, sfa_ref_cpu, sfb1_ref_cpu, sfb2_ref_cpu, _, _, _, c_ref = data
    m, n, l = c_ref.shape

    ref1 = torch.empty((l, m, n), dtype=torch.float32, device="cuda").permute(1, 2, 0)
    ref2 = torch.empty((l, m, n), dtype=torch.float32, device="cuda").permute(1, 2, 0)
    for l_idx in range(l):
        scale_a = to_blocked(sfa_ref_cpu[:, :, l_idx])
        scale_b1 = to_blocked(sfb1_ref_cpu[:, :, l_idx])
        scale_b2 = to_blocked(sfb2_ref_cpu[:, :, l_idx])

        res1 = torch._scaled_mm(
            a_ref[:, :, l_idx],
            b1_ref[:, :, l_idx].transpose(0, 1),
            scale_a.cuda(),
            scale_b1.cuda(),
            bias=None,
            out_dtype=torch.float32,
        )
        ref1[:, :, l_idx] = res1

        res2 = torch._scaled_mm(
            a_ref[:, :, l_idx],
            b2_ref[:, :, l_idx].transpose(0, 1),
            scale_a.cuda(),
            scale_b2.cuda(),
            bias=None,
            out_dtype=torch.float32,
        )
        ref2[:, :, l_idx] = res2

    return (torch.nn.functional.silu(ref1) * ref2).to(torch.float16)


@triton.jit
def _dual_gemm_silu_kernel_opt(
    a_ptr,  # uint8 packed fp4, [M, K_bytes, L]
    b1_ptr,  # uint8 packed fp4, [N, K_bytes, L]
    b2_ptr,  # uint8 packed fp4, [N, K_bytes, L]
    sfa_ptr,  # fp8 scale, [M, K//VEC_SIZE, L]
    sfb1_ptr,  # fp8 scale, [N, K//VEC_SIZE, L]
    sfb2_ptr,  # fp8 scale, [N, K//VEC_SIZE, L]
    c_ptr,  # fp16 output, [M, N, L]
    stride_am,
    stride_akb,
    stride_al,
    stride_b1n,
    stride_b1kb,
    stride_b1l,
    stride_b2n,
    stride_b2kb,
    stride_b2l,
    stride_sfam,
    stride_sfak,
    stride_sfal,
    stride_sfb1n,
    stride_sfb1k,
    stride_sfb1l,
    stride_sfb2n,
    stride_sfb2k,
    stride_sfb2l,
    stride_cm,
    stride_cn,
    stride_cl,
    M: tl.constexpr,
    N: tl.constexpr,
    K: tl.constexpr,
    L: tl.constexpr,
    ELEM_PER_BYTE: tl.constexpr,
    VEC_SIZE: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
    NUM_STAGES: tl.constexpr,
):
    pid_m = tl.program_id(axis=0)
    pid_n = tl.program_id(axis=1)
    pid_l = tl.program_id(axis=2)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)

    acc1 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    acc2 = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    k_bytes_total = K // ELEM_PER_BYTE
    k_tiles = tl.cdiv(K, BLOCK_K)
    kb_tile = BLOCK_K // ELEM_PER_BYTE
    ks_tile = BLOCK_K // VEC_SIZE

    base_a_l = a_ptr + pid_l * stride_al
    base_b1_l = b1_ptr + pid_l * stride_b1l
    base_b2_l = b2_ptr + pid_l * stride_b2l
    base_sfa_l = sfa_ptr + pid_l * stride_sfal
    base_sfb1_l = sfb1_ptr + pid_l * stride_sfb1l
    base_sfb2_l = sfb2_ptr + pid_l * stride_sfb2l

    for kt in tl.range(0, k_tiles, num_stages=NUM_STAGES):
        offs_kb = kt * kb_tile + tl.arange(0, BLOCK_K // ELEM_PER_BYTE)
        offs_ks = kt * ks_tile + tl.arange(0, BLOCK_K // VEC_SIZE)

        a_ptrs = base_a_l + offs_m[:, None] * stride_am + offs_kb[None, :] * stride_akb
        a = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & (offs_kb[None, :] < k_bytes_total), other=0).to(tl.uint8)

        sfa_ptrs = base_sfa_l + offs_m[:, None] * stride_sfam + offs_ks[None, :] * stride_sfak
        scale_a = tl.load(sfa_ptrs, mask=(offs_m[:, None] < M) & (offs_ks[None, :] < (K // VEC_SIZE)), other=0.0)

        b_ptrs = base_b1_l + offs_n[:, None] * stride_b1n + offs_kb[None, :] * stride_b1kb
        b = tl.load(b_ptrs, mask=(offs_n[:, None] < N) & (offs_kb[None, :] < k_bytes_total), other=0).to(tl.uint8)

        sfb_ptrs = base_sfb1_l + offs_n[:, None] * stride_sfb1n + offs_ks[None, :] * stride_sfb1k
        scale_b = tl.load(sfb_ptrs, mask=(offs_n[:, None] < N) & (offs_ks[None, :] < (K // VEC_SIZE)), other=0.0)

        acc1 = tl.dot_scaled(
            a, scale_a, "e2m1",
            b.T, scale_b, "e2m1",
            acc=acc1,
            fast_math=True,
            lhs_k_pack=True,
            rhs_k_pack=True,
            out_dtype=tl.float32,
        )
 
        b_ptrs = base_b2_l + offs_n[:, None] * stride_b2n + offs_kb[None, :] * stride_b2kb
        b = tl.load(b_ptrs, mask=(offs_n[:, None] < N) & (offs_kb[None, :] < k_bytes_total), other=0).to(tl.uint8)

        sfb_ptrs = base_sfb2_l + offs_n[:, None] * stride_sfb2n + offs_ks[None, :] * stride_sfb2k
        scale_b = tl.load(sfb_ptrs, mask=(offs_n[:, None] < N) & (offs_ks[None, :] < (K // VEC_SIZE)), other=0.0)

        acc2 = tl.dot_scaled(
            a, scale_a, "e2m1",
            b.T, scale_b, "e2m1",
            acc=acc2,
            fast_math=True,
            lhs_k_pack=True,
            rhs_k_pack=True,
            out_dtype=tl.float32,
        )
 
    sig = 1.0 / (1.0 + tl.exp(-acc1))
    acc1 = acc1 * sig
    out = acc1 * acc2
    out = out.to(tl.float16)

    c_ptrs = c_ptr + pid_l * stride_cl + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
    tl.store(c_ptrs, out, mask=mask)


def custom_kernel(data: input_t) -> output_t:        
    a, b1, b2, sfa, sfb1, sfb2, _, _, _, c = data

    m, k_bytes, l = a.shape
    n = b1.shape[0]
    k = k_bytes * 2

    a_u8 = a.view(torch.uint8)
    b1_u8 = b1.view(torch.uint8)
    b2_u8 = b2.view(torch.uint8)

    out = c

    BLOCK_M = 128
    BLOCK_N = 64
    BLOCK_K = 256
    VEC_SIZE = 16
    ELEM_PER_BYTE = 2

    num_warps = 4
    num_stages = 3

    grid = (triton.cdiv(m, BLOCK_M), triton.cdiv(n, BLOCK_N), l)

    _dual_gemm_silu_kernel_opt[grid](
        a_u8,
        b1_u8,
        b2_u8,
        sfa,
        sfb1,
        sfb2,
        out,
        a_u8.stride(0),
        a_u8.stride(1),
        a_u8.stride(2),
        b1_u8.stride(0),
        b1_u8.stride(1),
        b1_u8.stride(2),
        b2_u8.stride(0),
        b2_u8.stride(1),
        b2_u8.stride(2),
        sfa.stride(0),
        sfa.stride(1),
        sfa.stride(2),
        sfb1.stride(0),
        sfb1.stride(1),
        sfb1.stride(2),
        sfb2.stride(0),
        sfb2.stride(1),
        sfb2.stride(2),
        out.stride(0),
        out.stride(1),
        out.stride(2),
        M=m,
        N=n,
        K=k,
        L=l,
        ELEM_PER_BYTE=ELEM_PER_BYTE,
        VEC_SIZE=VEC_SIZE,
        BLOCK_M=BLOCK_M,
        BLOCK_N=BLOCK_N,
        BLOCK_K=BLOCK_K,
        NUM_STAGES=num_stages,
        num_warps=num_warps,
    )

    return out


check_implementation = make_match_reference(custom_kernel, rtol=1e-03, atol=1e-03)
scrolls · 248 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