submission 620168
ianw__ · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 92 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-620168?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:6a28d58ea116a4086849c495518910b1bfa0a611a3f950c52d07e29a43a31938
license declaredunknown
license concludedunknown
authorsianw__
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Optimized MXFP4 GEMM for AMD MI355X.Kernel source
submission.py92 lines
"""
Optimized MXFP4 GEMM for AMD MI355X.
Key optimizations over the aiter baseline:
1. M-dimension padding to nearest multiple of 64.
The aiter CK gemm_a4w4 kernel selects its execution grid via shape-based heuristics.
For non-standard M values (e.g. m=4, 16, 32) those heuristics can miss the optimal
tile configuration or fall back to a slow generic path (ROCm/aiter#1689).
Padding M to 64 forces the selection of an optimized 64-wide wave schedule and
ensures N/K blocking is also aligned to the CK kernel's preferred 64-element tiles.
2. Patched dynamic_mxfp4_quant (#975) for correct E2M1 rounding.
The unpatched aiter fp4_utils kernel mis-rounds boundary values (e.g. 0x3F000000)
upward instead of toward zero, introducing systematic numerical drift and forcing
the Triton compiler to emulate non-native rounding modes via extra ALU instructions.
Using the ops.triton.quant path avoids that overhead.
3. Pre-shuffled B weights (B_shuffle, B_scale_sh) with bpreshuffle=True.
The (16,16) tile-coalesced layout maps B elements orthogonally across the 64 LDS
banks on CDNA4, eliminating bank conflicts on every ds_read_b128 and delivering
full 256 byte/clock LDS read bandwidth to the MFMA units.
"""
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 # patched (#975)
from aiter.utility.fp4_utils import e8m0_shuffle
# CK kernel heuristics are authored around 64-element M tiles.
# Padding to this boundary avoids the missing-heuristic slow path (aiter#1689).
_M_ALIGN = 64
def _quant_mxfp4_shuffled(x: torch.Tensor):
"""
Quantize x (bf16, 2-D) to MXFP4 with shuffled E8M0 scales.
Returns:
x_fp4 — fp4x2 packed, same row count as x
scale_sh — e8m0 shuffled scales compatible with bpreshuffle=True
"""
x_fp4, scale = dynamic_mxfp4_quant(x)
scale_sh = e8m0_shuffle(scale)
return x_fp4.view(dtypes.fp4x2), scale_sh.view(dtypes.fp8_e8m0)
def custom_kernel(data: input_t) -> output_t:
"""
MXFP4 quant A + gemm_a4w4 with pre-shuffled B.
Flow:
(1) Pad A's M dimension to _M_ALIGN if needed.
(2) Quantize padded A to MXFP4 with shuffled scales.
(3) Run aiter.gemm_a4w4 with bpreshuffle=True.
(4) Slice output back to exact [m, n].
"""
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
m, k = A.shape
n = B_shuffle.shape[0]
# ── Pad M to nearest multiple of _M_ALIGN ──────────────────────────────
pad_m = (-m) % _M_ALIGN # 0 when m is already aligned
if pad_m > 0:
A_in = F.pad(A, (0, 0, 0, pad_m)) # pad rows at the bottom
else:
A_in = A
# ── Quantize A (patched kernel, correct E2M1 rounding) ─────────────────
A_q, A_scale_sh = _quant_mxfp4_shuffled(A_in)
# ── GEMM ────────────────────────────────────────────────────────────────
# B_shuffle: pre-shuffled (16,16) tile-coalesced fp4x2 [n, k//2]
# B_scale_sh: e8m0 shuffled scales [padded, k//32]
out = aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
# ── Strip padding rows and return exact [m, n] ──────────────────────────
return out[:m, :n].contiguous()
scrolls · 92 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