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