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
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