Skip to content
KernelIndex
Search⌘K

submission 152699

dandanaka_hitman · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-152699?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
#133 of 369
2025-12-13

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e3d87b5fcdcf42d058e740ba59df4df3ed99477539d20847304e1f87d43a60db
license declaredunknown
license concludedunknown
authorsdandanaka_hitman
imported2026-08-26

Kernel source

submission_3.py168 lines
import torch
from task import input_t, output_t

import cutlass
import cutlass.cute as cute
from cutlass.cute.runtime import make_ptr
import cutlass.utils.blockscaled_layout as blockscaled_utils

# -----------------------------------------------------------------------------
# Controls (edit constants; no env vars needed)
# -----------------------------------------------------------------------------
RUN_SFA_DIAG = True      # set True for one debug submission
SFA_DIAG_RAISE = False    # if True and mismatch, raise RuntimeError (will fail run intentionally)

# How many elements to compare from the 1D scale vector
SFA_DIAG_N = 4096

# -----------------------------------------------------------------------------
# Known-good scale vector from permuted scales (matches reference.to_blocked)
# -----------------------------------------------------------------------------
def _scale_vec_from_permuted(sf_permuted: torch.Tensor, l_idx: int) -> torch.Tensor:
    # sf_permuted: (32, 4, rest_mn, 4, rest_k, L)
    sf = sf_permuted[..., l_idx]          # (32, 4, rest_mn, 4, rest_k)
    sf = sf.permute(2, 4, 0, 1, 3)        # (rest_mn, rest_k, 32, 4, 4)
    return sf.contiguous().view(-1)

_scaled_mm_supports_out = None

@torch.no_grad()
def _fallback_scaled_mm(a, b, sfa_permuted, sfb_permuted, c) -> torch.Tensor:
    """Always-correct path: compute using torch._scaled_mm."""
    global _scaled_mm_supports_out
    _, _, L = c.shape
    for l_idx in range(L):
        scale_a = _scale_vec_from_permuted(sfa_permuted, l_idx)
        scale_b = _scale_vec_from_permuted(sfb_permuted, l_idx)
        aL = a[:, :, l_idx]
        bTL = b[:, :, l_idx].transpose(0, 1)
        cL = c[:, :, l_idx]

        if _scaled_mm_supports_out is None:
            try:
                torch._scaled_mm(aL, bTL, scale_a, scale_b, bias=None, out_dtype=torch.float16, out=cL)
                _scaled_mm_supports_out = True
            except TypeError:
                _scaled_mm_supports_out = False

        if _scaled_mm_supports_out:
            torch._scaled_mm(aL, bTL, scale_a, scale_b, bias=None, out_dtype=torch.float16, out=cL)
        else:
            cL.copy_(torch._scaled_mm(aL, bTL, scale_a, scale_b, bias=None, out_dtype=torch.float16))
    return c

# -----------------------------------------------------------------------------
# CuTe SFA dump kernel: out[i] = inp[i] for i < N (linear indexing)
# -----------------------------------------------------------------------------
sf_dtype = cutlass.Float8E4M3FN

@cute.kernel
def dump_linear_kernel(inp: cute.Tensor, out: cute.Tensor, N: cutlass.Constexpr[int]):
    tidx, _, _ = cute.arch.thread_idx()
    idx = tidx
    step = 256

    total = cute.size(inp)  # may be dynamic; that's fine

    while idx < N:
        if idx < total:
            out[idx] = inp[idx]
        idx += step

@cute.jit
def dump_sfa_host(
    sfa_ptr: cute.Pointer,
    m: int,
    k: int,
    l: int,
    out_ptr: cute.Pointer,
):
    """
    Construct the same CuTe SFA tensor layout your GEMM would use, and dump its
    linearized first N elements into out.
    """
    # Match your CuTe build: assume(x, 32) form
    m32  = cute.assume(m, 32)
    k32  = cute.assume(k, 32)

    # Fake A shape used only to build the SF layout
    # A tensor is (m, k, l) in the "logical-K" sense for the blockscaled helper.
    a_shape = (m32, k32, l)

    # Build SFA layout that CuTe expects for blockscaled loads
    sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a_shape, 16)  # sf_vec_size=16 (reference)
    sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)

    out_tensor = cute.make_tensor(out_ptr, cute.make_layout((SFA_DIAG_N,), stride=(1,)))

    dump_linear_kernel(sfa_tensor, out_tensor, SFA_DIAG_N).launch(grid=(1,1,1), block=[256,1,1], cluster=(1,1,1))
    return

_dump_compiled = None

def _compile_dump():
    global _dump_compiled
    if _dump_compiled is not None:
        return _dump_compiled

    sfa_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
    out_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)

    # compile with representative dims (m multiple of 128, k multiple of 256, l=1 typical)
    _dump_compiled = cute.compile(dump_sfa_host, sfa_ptr, 128, 256, 1, out_ptr)
    return _dump_compiled

_diag_done = False

@torch.no_grad()
def custom_kernel(data: input_t) -> output_t:
    """
    Always returns correct output via torch._scaled_mm.
    Optionally runs a deterministic SFA layout probe (once) that compares:
      - CuTe's interpretation of SFA GMEM layout via tile_atom_to_shape_SF
      - known-good blocked vector order from sfa_permuted (used by torch._scaled_mm)
    """
    global _diag_done

    a, b, sfa, sfb, sfa_permuted, sfb_permuted, c = data

    # Optional one-time diagnostic
    if RUN_SFA_DIAG and not _diag_done:
        _diag_done = True

        # Use l=0 only (leaderboard uses L=1)
        l_idx = 0

        # known-good vector (what torch._scaled_mm expects)
        ref = _scale_vec_from_permuted(sfa_permuted, l_idx)

        # CuTe-dumped vector: interpret the SAME memory as a CuTe blockscaled SFA tensor and dump linear order
        dump = torch.empty((SFA_DIAG_N,), device=a.device, dtype=torch.float8_e4m3fn)

        m = int(sfa.shape[0])
        # logical K in FP4 elements for reference: k = (a.shape[1] * 2)
        # this is only used to construct the SF layout; reference uses sf_vec_size=16.
        k = int(a.shape[1]) * 2
        l = int(a.shape[2])

        compiled = _compile_dump()
        sfa_ptr = make_ptr(sf_dtype, int(sfa_permuted.contiguous().data_ptr()), cute.AddressSpace.gmem, assumed_align=16)
        out_ptr = make_ptr(sf_dtype, int(dump.data_ptr()), cute.AddressSpace.gmem, assumed_align=16)

        compiled(sfa_ptr, m, k, l, out_ptr)

        # Compare prefix
        N = min(SFA_DIAG_N, ref.numel())
        ref_prefix = ref[:N].to(torch.float16)
        dump_prefix = dump[:N].to(torch.float16)
        ok = torch.allclose(dump_prefix, ref_prefix, rtol=0, atol=0)

        if not ok:
            max_abs = (dump_prefix - ref_prefix).abs().max().item()
            if SFA_DIAG_RAISE:
                raise RuntimeError(f"SFA_DIAG mismatch: CuTe SFA layout != scaled_mm blocked order. max_abs={max_abs}")
            # otherwise: just continue; submission remains correct

    # Always-correct output
    return _fallback_scaled_mm(a, b, sfa_permuted, sfb_permuted, c)
scrolls · 168 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