Skip to content
KernelIndex
Search⌘K

submission 123301

Batuhanaktas · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

kai_3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-123301?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
13.4µs
#138 of 369
2025-12-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:501caaa25c0df94cb94dd3785d617dc60947d996e57b9c179915e463a5134c56
license declaredunknown
license concludedunknown
authorsBatuhanaktas
imported2026-08-26

Techniques

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

fp4"""Optimized NVFP4 block-scaled GEMM using pre-permuted GPU scales."""

Kernel source

kai_3.py76 lines

import torch
from typing import Tuple

# Competition harness signature
input_t = Tuple[
    torch.Tensor,  # a
    torch.Tensor,  # b
    torch.Tensor,  # sfa
    torch.Tensor,  # sfb
    torch.Tensor,  # sfa_permuted
    torch.Tensor,  # sfb_permuted
    torch.Tensor,  # c
]
output_t = torch.Tensor

_scaled_mm = torch._scaled_mm


def custom_kernel(data: input_t) -> output_t:
    """Optimized NVFP4 block-scaled GEMM using pre-permuted GPU scales."""
    a, b, sfa, sfb, sfa_permuted, sfb_permuted, c_out = data

    M, N, L = c_out.shape

    sfa_flat = sfa_permuted.permute(5, 2, 4, 0, 1, 3).reshape(L, -1)
    sfb_flat = sfb_permuted.permute(5, 2, 4, 0, 1, 3).reshape(L, -1)

    if L == 1:
        a_slice = a.view(M, -1)
        b_slice = b.view(N, -1).t()
        out = _scaled_mm(
            a_slice,
            b_slice,
            sfa_flat[0],
            sfb_flat[0],
            bias=None,
            out_dtype=torch.float16,
        )
        return out.view(M, N, 1)

    result = torch.empty((M, N, L), dtype=torch.float16, device=a.device)

    if L == 2:
        result[:, :, 0] = _scaled_mm(a[:, :, 0], b[:, :, 0].t(), sfa_flat[0], sfb_flat[0], bias=None, out_dtype=torch.float16)
        result[:, :, 1] = _scaled_mm(a[:, :, 1], b[:, :, 1].t(), sfa_flat[1], sfb_flat[1], bias=None, out_dtype=torch.float16)
        return result

    if L == 3:
        result[:, :, 0] = _scaled_mm(a[:, :, 0], b[:, :, 0].t(), sfa_flat[0], sfb_flat[0], bias=None, out_dtype=torch.float16)
        result[:, :, 1] = _scaled_mm(a[:, :, 1], b[:, :, 1].t(), sfa_flat[1], sfb_flat[1], bias=None, out_dtype=torch.float16)
        result[:, :, 2] = _scaled_mm(a[:, :, 2], b[:, :, 2].t(), sfa_flat[2], sfb_flat[2], bias=None, out_dtype=torch.float16)
        return result

    if L == 4:
        result[:, :, 0] = _scaled_mm(a[:, :, 0], b[:, :, 0].t(), sfa_flat[0], sfb_flat[0], bias=None, out_dtype=torch.float16)
        result[:, :, 1] = _scaled_mm(a[:, :, 1], b[:, :, 1].t(), sfa_flat[1], sfb_flat[1], bias=None, out_dtype=torch.float16)
        result[:, :, 2] = _scaled_mm(a[:, :, 2], b[:, :, 2].t(), sfa_flat[2], sfb_flat[2], bias=None, out_dtype=torch.float16)
        result[:, :, 3] = _scaled_mm(a[:, :, 3], b[:, :, 3].t(), sfa_flat[3], sfb_flat[3], bias=None, out_dtype=torch.float16)
        return result

    a_slices = [a[:, :, i] for i in range(L)]
    b_slices = [b[:, :, i].t() for i in range(L)]

    for l_idx in range(L):
        result[:, :, l_idx] = _scaled_mm(
            a_slices[l_idx],
            b_slices[l_idx],
            sfa_flat[l_idx],
            sfb_flat[l_idx],
            bias=None,
            out_dtype=torch.float16,
        )

    return result
scrolls · 76 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