Skip to content
KernelIndex
Search⌘K

submission 118340

poornaravuri · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_triton.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-118340?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
43.9µs
#244 of 369
2025-12-01

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e93b7ac982ef292d8820a4de3d0b63eb30660498b9d80f00a4d4d92867121d7f
license declaredunknown
license concludedunknown
authorspoornaravuri
imported2026-08-26

Techniques

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

fp8other=tl.zeros((BLOCK_M, num_vec), dtype=tl.float8e4nv),
num-warps = 8num_warps=8,
stages = 2num_stages=2,
tile-k = 256BLOCK_K = 256
tile-m = 128BLOCK_M = 128
tile-n = 128BLOCK_N = 128

Kernel source

submission_triton.py164 lines
import torch

try:
    import triton
    import triton.language as tl
except Exception:
    triton = None

from task import input_t, output_t
from utils import make_match_reference

sf_vec_size = 16


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


if triton is not None:

    @triton.jit
    def _bsmm_kernel(
        a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,
        M, N, K, sfK,
        stride_am, stride_ak,
        stride_bn, stride_bk,
        stride_sm, stride_sk,
        stride_tn, stride_tk,
        stride_cm, stride_cn,
        BLOCK_M: tl.constexpr,
        BLOCK_N: tl.constexpr,
        BLOCK_K: tl.constexpr,
        VEC_SIZE: tl.constexpr,
    ):
        pid = tl.program_id(0)
        num_pid_m = tl.cdiv(M, BLOCK_M)
        pid_m = pid % num_pid_m
        pid_n = pid // num_pid_m

        offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
        offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
        m_mask = offs_m < M
        n_mask = offs_n < N

        acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

        K_packed = K // 2
        BLOCK_K_PACKED: tl.constexpr = BLOCK_K // 2

        for k0 in range(0, K, BLOCK_K):
            offs_k_packed = (k0 // 2) + tl.arange(0, BLOCK_K_PACKED)
            k_mask = offs_k_packed < K_packed

            a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k_packed[None, :] * stride_ak
            b_ptrs = b_ptr + offs_n[:, None] * stride_bn + offs_k_packed[None, :] * stride_bk

            a_tile = tl.load(a_ptrs, mask=m_mask[:, None] & k_mask[None, :], other=0)
            b_tile = tl.load(b_ptrs, mask=n_mask[:, None] & k_mask[None, :], other=0)

            num_vec: tl.constexpr = BLOCK_K // VEC_SIZE
            offs_vec = (k0 // VEC_SIZE) + tl.arange(0, num_vec)
            vec_mask = offs_vec < sfK

            sfa_ptrs = sfa_ptr + offs_m[:, None] * stride_sm + offs_vec[None, :] * stride_sk
            sfb_ptrs = sfb_ptr + offs_n[:, None] * stride_tn + offs_vec[None, :] * stride_tk

            scale_a = tl.load(
                sfa_ptrs,
                mask=m_mask[:, None] & vec_mask[None, :],
                other=tl.zeros((BLOCK_M, num_vec), dtype=tl.float8e4nv),
            )
            scale_b = tl.load(
                sfb_ptrs,
                mask=n_mask[:, None] & vec_mask[None, :],
                other=tl.zeros((BLOCK_N, num_vec), dtype=tl.float8e4nv),
            )

            acc = tl.dot_scaled(a_tile, scale_a, "e2m1", b_tile.T, scale_b, "e2m1", acc)

        c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
        out = acc.to(c_ptr.type.element_ty)
        tl.store(c_ptrs, out, mask=m_mask[:, None] & n_mask[None, :])


def _fallback_scaled_mm(data: input_t) -> output_t:
    a_ref, b_ref, sfa_ref, sfb_ref, _, _, c_ref = data
    _, _, L = c_ref.shape
    device = a_ref.device

    def to_blocked(mat: torch.Tensor) -> torch.Tensor:
        rows, cols = mat.shape
        n_row_blocks = ceil_div(rows, 128)
        n_col_blocks = ceil_div(cols, 4)
        blocks = mat.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()

    for l_idx in range(L):
        scale_a = to_blocked(sfa_ref[:, :, l_idx]).to(device=device)
        scale_b = to_blocked(sfb_ref[:, :, l_idx]).to(device=device)
        res = torch._scaled_mm(
            a_ref[:, :, l_idx],
            b_ref[:, :, l_idx].transpose(0, 1),
            scale_a,
            scale_b,
            bias=None,
            out_dtype=torch.float16,
        )
        c_ref[:, :, l_idx] = res
    return c_ref


def custom_kernel(data: input_t) -> output_t:
    a_ref, b_ref, sfa_ref, sfb_ref, _, _, c_ref = data
    M, K_packed, L = a_ref.shape
    N, _, _ = b_ref.shape
    K = K_packed * 2
    sfK = sfa_ref.shape[1]

    use_triton = (
        triton is not None
        and torch.cuda.is_available()
        and torch.cuda.get_device_capability()[0] >= 10
    )

    if not use_triton:
        return _fallback_scaled_mm(data)

    BLOCK_M = 128
    BLOCK_N = 128
    BLOCK_K = 256
    VEC_SIZE = sf_vec_size

    for l_idx in range(L):
        # packed NVFP4 must be uint8 for tl.dot_scaled("e2m1")
        a_l = a_ref[:, :, l_idx].view(torch.uint8).contiguous()
        b_l = b_ref[:, :, l_idx].view(torch.uint8).contiguous()
        sfa_l = sfa_ref[:, :, l_idx].contiguous()
        sfb_l = sfb_ref[:, :, l_idx].contiguous()
        c_l = c_ref[:, :, l_idx]

        grid = (triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N),)

        _bsmm_kernel[grid](
            a_l, b_l, sfa_l, sfb_l, c_l,
            M, N, K, sfK,
            a_l.stride(0), a_l.stride(1),
            b_l.stride(0), b_l.stride(1),
            sfa_l.stride(0), sfa_l.stride(1),
            sfb_l.stride(0), sfb_l.stride(1),
            c_l.stride(0), c_l.stride(1),
            BLOCK_M=BLOCK_M,
            BLOCK_N=BLOCK_N,
            BLOCK_K=BLOCK_K,
            VEC_SIZE=VEC_SIZE,
            num_warps=8,
            num_stages=2,
        )

    return c_ref


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