submission 614289
anairdrop · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 97 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-614289?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:b7803e151fe942f73f25c0ff97e9fc8da5da717b0d977bcd71eb3f5c5f9396d8
license declaredunknown
license concludedunknown
authorsanairdrop
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Optimized MXFP4-MM: Monkey-patch aiter.gemm_a4w4 to route throughKernel source
submission.py97 lines
"""
Optimized MXFP4-MM: Monkey-patch aiter.gemm_a4w4 to route through
gemm_a4w4_asm with per-shape kernel tile selection.
Key: monkey-patch BEFORE importing ref_kernel so both our call and
the validation call go through the same optimized path.
Sweep results (µs, best tile per shape):
M=4, N=2880, K=512: AUTO=6.5, 32x768=6.9 → use ""
M=16, N=2112, K=7168: 32x128=7.9, AUTO=15.3 → use 32x128
M=32, N=4096, K=512: 256x128=6.8, 32x128=6.9 → use 32x128
M=32, N=2880, K=512: 64x512=6.9, 32x256=6.9 → use 32x128
M=64, N=7168, K=2048: 32x128=6.7, 96x128=7.0 → use 32x128
M=256, N=3072, K=1536: 32x128=6.7, 96x128=6.9 → use 32x128
"""
import torch
import aiter
from aiter import dtypes
from task import input_t, output_t
# Save original before patching
_orig_gemm_a4w4 = aiter.gemm_a4w4
_asm_fn = None
def _get_asm_fn():
global _asm_fn
if _asm_fn is not None:
return _asm_fn
try:
_asm_fn = aiter.gemm_a4w4_asm
except AttributeError:
_asm_fn = False
return _asm_fn
# Build mangled kernel name for a tile
def _mangle(tile):
name = f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{tile}"
return f"_ZN5aiter{len(name)}{name}E"
# Pre-build the ones we need
_K32x128 = _mangle("32x128")
def _select_kernel_name(M, N, K):
"""Select best ASM kernel tile based on sweep data.
Key finding: 32x128 is optimal or near-optimal for all shapes with K >= 1024.
For small K with small M, AUTO (empty string) lets the runtime pick.
"""
if M <= 8 and K <= 1024:
return "" # AUTO is best for tiny M, small K
return _K32x128 # 32x128 wins everywhere else
_out_cache = {}
def _patched_gemm_a4w4(A, B, A_scale, B_scale, bias=None, dtype=15,
alpha=1.0, beta=0.0, bpreshuffle=True):
asm = _get_asm_fn()
if asm and asm is not False:
M = A.shape[0]
N = B.shape[0]
K = A.shape[1] * 2 # fp4x2 packed
out_dtype = torch.bfloat16 if dtype == 15 else torch.float16
# Cache output tensor to avoid torch.empty overhead (~3µs)
cache_key = (M, N, out_dtype)
if cache_key not in _out_cache:
_out_cache[cache_key] = torch.empty((M, N), dtype=out_dtype, device=A.device)
out = _out_cache[cache_key]
kernel_name = _select_kernel_name(M, N, K)
try:
asm(A, B, A_scale, B_scale, out, kernel_name,
bias=bias, alpha=alpha, beta=beta,
bpreshuffle=bpreshuffle, log2_k_split=0)
return out
except Exception:
pass
return _orig_gemm_a4w4(A, B, A_scale, B_scale, bias=bias, dtype=dtype,
alpha=alpha, beta=beta, bpreshuffle=bpreshuffle)
# Apply monkey-patch
aiter.gemm_a4w4 = _patched_gemm_a4w4
# Import ref_kernel AFTER patching
from reference import ref_kernel
def custom_kernel(data: input_t) -> output_t:
return ref_kernel(data)
scrolls · 97 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