submission 739858
gxtzhuxi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 154 lines, June 9 Researcher Reciprocity License v1.0.
mxfp4_gemm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-739858?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:7262cd01e8f6c8664b84f50b8c1c66cba06078ee79405256047de31fe89cbcfc
license declaredunknown
license concludedunknown
authorsgxtzhuxi
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 GEMM — bf16 A × MXFP4 B → bf16 C on MI355X (CDNA4).Kernel source
mxfp4_gemm.py154 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
MXFP4 GEMM — bf16 A × MXFP4 B → bf16 C on MI355X (CDNA4).
== Architecture ==
MI355X: 256 CUs, 8 XCDs, 64KB LDS/CU, 32MB L2, 256MB Infinity Cache,
8 TB/s HBM3E, FP4 MFMA (10 PFLOPS).
== Explicit Kernel Control ==
Calls gemm_a4w4_asm DIRECTLY with pre-allocated output buffer,
bypassing the gemm_a4w4 wrapper to eliminate per-call allocation
and kernel-selection overhead.
Pipeline (3 kernel launches):
1. dynamic_mxfp4_quant(A) → A_fp4, A_scale [Triton kernel]
2. e8m0_shuffle(A_scale) → A_scale_shuffled [CK MFMA layout]
3. gemm_a4w4_asm(...) → C [FP4 MFMA ASM]
First call per shape uses gemm_a4w4 wrapper for JIT warmup and
ASM kernel-name discovery (verified against reference). Subsequent
calls bypass the wrapper with pre-allocated output + padded A buffers.
"""
import torch
import torch.nn.functional as F
from task import input_t, output_t
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
# ===== Global state =====
_CACHE = {}
_DIRECT_FNS = {}
_KERNEL_PREFIX = "f4gemm_bf16_per1x32Fp4_BpreShuffle_"
_KNOWN_TILES = [32, 192]
def _discover_direct_config(M, N, K, A_q, B_shuffle, A_sc, B_scale_sh,
reference, device):
"""Try CK ASM kernel names to find one matching the wrapper's output.
Verifies correctness with torch.equal before committing. On success,
pre-allocates output + A-padding buffers for the direct path.
"""
asm = _DIRECT_FNS.get("asm")
get_pm = _DIRECT_FNS.get("get_padded_m")
if asm is None or get_pm is None:
return None
padded_m = get_pm(M, N, K, 1)
if padded_m > M:
a_pad = F.pad(A_q.view(torch.uint8),
(0, 0, 0, padded_m - M)).view(dtypes.fp4x2)
sc_pad = F.pad(A_sc.view(torch.uint8),
(0, 0, 0, padded_m - M)).view(dtypes.fp8_e8m0)
else:
a_pad = A_q
sc_pad = A_sc
out = torch.empty(padded_m, N, dtype=torch.bfloat16, device=device)
candidates = sorted(set([padded_m] + _KNOWN_TILES), reverse=True)
for tile_m in candidates:
kname = f"{_KERNEL_PREFIX}{tile_m}x128"
for log2_ks in [None, 0, 1, 2]:
try:
out.zero_()
asm(a_pad, B_shuffle, sc_pad, B_scale_sh, out, kname,
bpreshuffle=True, log2_k_split=log2_ks)
if torch.equal(out[:M, :], reference):
K_half = A_q.view(torch.uint8).shape[1]
K_sc = A_sc.view(torch.uint8).shape[1]
return {
"kname": kname,
"log2_ks": log2_ks,
"padded_m": padded_m,
"needs_pad": padded_m > M,
"out": out,
"A_q_buf": torch.zeros(padded_m, K_half,
dtype=torch.uint8, device=device),
"A_sc_buf": torch.zeros(padded_m, K_sc,
dtype=torch.uint8, device=device),
}
except Exception:
continue
return None
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.shape[0]
device = A.device
# Stage 1: Quantize A (dynamic activations — every call)
A_fp4, A_scale = dynamic_mxfp4_quant(A)
# Stage 2: Shuffle A's E8M0 scale to CK's MFMA-friendly layout
A_scale_sh = e8m0_shuffle(A_scale)
A_q = A_fp4.view(dtypes.fp4x2)
A_sc = A_scale_sh.view(dtypes.fp8_e8m0)
key = (M, N, K)
if key not in _CACHE:
# JIT warmup: use wrapper for reference result + kernel auto-selection
ref = aiter.gemm_a4w4(
A_q, B_shuffle, A_sc, B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True,
)
if "asm" not in _DIRECT_FNS:
_DIRECT_FNS["asm"] = getattr(aiter, "gemm_a4w4_asm", None)
_DIRECT_FNS["get_padded_m"] = getattr(aiter, "get_padded_m", None)
_CACHE[key] = _discover_direct_config(
M, N, K, A_q, B_shuffle, A_sc, B_scale_sh, ref, device,
)
return ref
c = _CACHE[key]
# === Direct ASM path: pre-allocated output, no wrapper overhead ===
if c is not None:
if c["needs_pad"]:
c["A_q_buf"][:M].copy_(A_q.view(torch.uint8))
c["A_sc_buf"][:M].copy_(A_sc.view(torch.uint8))
a_in = c["A_q_buf"].view(dtypes.fp4x2)
sc_in = c["A_sc_buf"].view(dtypes.fp8_e8m0)
else:
a_in = A_q
sc_in = A_sc
_DIRECT_FNS["asm"](
a_in, B_shuffle, sc_in, B_scale_sh,
c["out"], c["kname"],
bpreshuffle=True, log2_k_split=c["log2_ks"],
)
return c["out"][:M, :]
# === Wrapper fallback (if kernel discovery failed) ===
return aiter.gemm_a4w4(
A_q, B_shuffle, A_sc, B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True,
)
scrolls · 154 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