submission 673983
rujutafujuta · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 129 lines, June 9 Researcher Reciprocity License v1.0.
mxfp4-mm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-673983?include=source"interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4
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:3b270a148946aff8e401878f680f735860330ebee4cd53e30b0680a983df0147
license declaredunknown
license concludedunknown
authorsrujutafujuta
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
mxfp4-mm.py129 lines
#!POPCORN leaderboard amd-mxfp4-mm
"""
MXFP4-MM optimized submission.
Strategy:
- Discover valid kernelIds for gemm_a4w4_blockscale_tune at startup (server has
different compiled kernel IDs than the local CSV, starting from 0).
- Probe up to 24 IDs (reduced from 64 to avoid timeouts).
- Runtime-tune across all valid kernelIds × splitK {0} per shape.
splitK=1,2 causes numerical errors for large-K shapes (k=7168) — removed.
- Fall back to aiter.gemm_a4w4 (reference path) if no blockscale kernel found.
- Reduced timing iterations (1 warmup, 3 timed) to stay within timeout budget.
"""
import torch
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from task import input_t, output_t
# Cached best (kernelId, splitK) or None (→ fallback) per (m, n, k)
_TUNE_CACHE: dict = {}
# Valid kernelIds discovered once at first use
_VALID_KERNEL_IDS: list | None = None
_SPLIT_KS = [0] # splitk=1,2 causes numerical errors for large-K shapes (k=7168)
_MAX_KERNEL_PROBE = 24 # don't probe all 64 — stop early to avoid timeouts
def _quant_mxfp4_shuffled(x: torch.Tensor):
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
bs_e8m0_sh = e8m0_shuffle(bs_e8m0)
return x_fp4.view(dtypes.fp4x2), bs_e8m0_sh.view(dtypes.fp8_e8m0)
def _discover_kernel_ids(out, A_flat, B_shuffle, A_scale_sh, B_scale_sh) -> list:
"""
Find all valid kernelId values by probing 0, 1, 2, ... until out-of-range.
Stops at the first 'out of range' error or after _MAX_KERNEL_PROBE attempts.
"""
valid = []
for kid in range(_MAX_KERNEL_PROBE):
try:
aiter.gemm_a4w4_blockscale_tune(A_flat, B_shuffle, A_scale_sh, B_scale_sh, out, kid, 0)
torch.cuda.synchronize()
valid.append(kid)
except RuntimeError as e:
if "out of range" in str(e).lower():
break
# Other RuntimeErrors (shape mismatch etc.) — skip this id but keep going
continue
except Exception:
continue
return valid
def _tune_shape(A_flat, B_shuffle, A_scale_sh, B_scale_sh, m, n):
"""Return (kernelId, splitK) for fastest config, or None to use fallback."""
global _VALID_KERNEL_IDS
padded_m = ((m + 31) // 32) * 32
out = torch.empty((padded_m, n), dtype=dtypes.bf16, device="cuda")
# Discover valid IDs once using the first shape we see
if _VALID_KERNEL_IDS is None:
_VALID_KERNEL_IDS = _discover_kernel_ids(out, A_flat, B_shuffle, A_scale_sh, B_scale_sh)
if not _VALID_KERNEL_IDS:
return None # no blockscale kernels available, use fallback
best_us = float("inf")
best_config = None # only set on a successful timed run
for kid in _VALID_KERNEL_IDS:
for sk in _SPLIT_KS:
try:
# 1 warmup run
aiter.gemm_a4w4_blockscale_tune(A_flat, B_shuffle, A_scale_sh, B_scale_sh, out, kid, sk)
torch.cuda.synchronize()
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(3):
aiter.gemm_a4w4_blockscale_tune(A_flat, B_shuffle, A_scale_sh, B_scale_sh, out, kid, sk)
end.record()
torch.cuda.synchronize()
us = start.elapsed_time(end) * 1e3 / 3
if us < best_us:
best_us = us
best_config = (kid, sk)
except Exception:
continue
return best_config
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
m, k = A.shape
n = B_shuffle.shape[0]
A_q, A_scale_sh = _quant_mxfp4_shuffled(A)
A_flat = A_q.view(m, k // 2)
shape_key = (m, n, k)
if shape_key not in _TUNE_CACHE:
_TUNE_CACHE[shape_key] = _tune_shape(A_flat, B_shuffle, A_scale_sh, B_scale_sh, m, n)
config = _TUNE_CACHE[shape_key]
padded_m = ((m + 31) // 32) * 32
out = torch.empty((padded_m, n), dtype=dtypes.bf16, device="cuda")
if config is None:
# No blockscale kernels available — fall back to reference path
return aiter.gemm_a4w4(A_q, B_shuffle, A_scale_sh, B_scale_sh, dtype=dtypes.bf16, bpreshuffle=True)
kid, sk = config
try:
aiter.gemm_a4w4_blockscale_tune(A_flat, B_shuffle, A_scale_sh, B_scale_sh, out, kid, sk)
return out[:m]
except RuntimeError:
_TUNE_CACHE[shape_key] = None
return aiter.gemm_a4w4(A_q, B_shuffle, A_scale_sh, B_scale_sh, dtype=dtypes.bf16, bpreshuffle=True)
scrolls · 129 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