submission 616464
maxvel · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 291 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-616464?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:d07839ed2ce7de3ad8905521817d3e93bd5db3c800e2ac0d32549f901d8519e6
license declaredunknown
license concludedunknown
authorsmaxvel
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
submission.py291 lines
"""
MI355X-oriented MXFP4 submission for `mxfp4-mm`.
Research summary embedded here to keep all implementation details in this file:
- The input already includes `B_shuffle` and `B_scale_sh`, so the main hot path is
quantizing `A` to MXFP4 and dispatching `aiter.gemm_a4w4`.
- The AMD evaluator benchmarks the same shape repeatedly in one spawned process,
which makes module-level lazy init, one-time warmup, and per-shape dispatch
caching worthwhile.
- AMD public ROCm materials identify MI355X as `gfx950`, and recent ROCm/aiter
releases mention ongoing FP4, A4W4 blockscale, split-K, and MI35x tuning work.
Tiered strategy implemented below:
- Tier 1: remove per-call imports/helper construction, avoid needless copies, and
cache the selected execution path per shape.
- Tier 2: add hot-shape specializations plus one-time warmup and candidate GEMM
probing.
- Tier 3: keep a safe framework for trying alternative A4W4 entry points when the
local `aiter` build exposes them, while falling back to the stable baseline.
"""
from __future__ import annotations
from typing import Any, Callable
import inspect
from task import input_t, output_t
_HOT_M_VALUES = frozenset((4, 8, 16, 32, 64, 256))
_HOT_K_VALUES = frozenset((512, 1536, 2048, 7168))
_RUNTIME: dict[str, Any] | None = None
_DISPATCH_CACHE: dict[tuple[Any, ...], Callable[..., output_t]] = {}
_WARMED_DEVICES: set[tuple[str, int | None]] = set()
def _supports_gemm_signature(fn: Callable[..., Any]) -> bool:
try:
params = inspect.signature(fn).parameters
except (TypeError, ValueError):
return False
return "dtype" in params and "bpreshuffle" in params
def _init_runtime() -> dict[str, Any]:
global _RUNTIME
if _RUNTIME is not None:
return _RUNTIME
import torch
import aiter
from aiter import dtypes
from aiter.ops.shuffle import shuffle_weight
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
candidate_names = (
"gemm_a4w4",
"gemm_a4w4_tuned",
"gemm_a4w4_blockscale",
"ck_gemm_a4w4_blockscale",
"asm_gemm_a4w4",
)
candidate_specs: list[tuple[str, Callable[..., Any]]] = []
seen_ids: set[int] = set()
for name in candidate_names:
fn = getattr(aiter, name, None)
if fn is None or not callable(fn):
continue
if id(fn) in seen_ids:
continue
if name != "gemm_a4w4" and not _supports_gemm_signature(fn):
continue
candidate_specs.append((name, fn))
seen_ids.add(id(fn))
if not candidate_specs:
raise RuntimeError("No usable A4W4 GEMM entry point found in aiter")
_RUNTIME = {
"torch": torch,
"aiter": aiter,
"dtypes": dtypes,
"shuffle_weight": shuffle_weight,
"dynamic_mxfp4_quant": dynamic_mxfp4_quant,
"e8m0_shuffle": e8m0_shuffle,
"gemm_candidates": tuple(candidate_specs),
}
return _RUNTIME
def _quant_mxfp4(x, runtime: dict[str, Any]):
x_fp4, bs_e8m0 = runtime["dynamic_mxfp4_quant"](x)
bs_e8m0 = runtime["e8m0_shuffle"](bs_e8m0)
dtypes = runtime["dtypes"]
return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)
def _maybe_contiguous(x):
return x if x.is_contiguous() else x.contiguous()
def _candidate_score(name: str, hot_shape: bool) -> int:
score = 0
lowered = name.lower()
if hot_shape:
if "tuned" in lowered:
score += 8
if "asm" in lowered:
score += 6
if "ck" in lowered or "blockscale" in lowered:
score += 4
if name == "gemm_a4w4":
score += 2
return score
def _select_gemm(runtime: dict[str, Any], m: int, k: int) -> Callable[..., Any]:
hot_shape = m in _HOT_M_VALUES and k in _HOT_K_VALUES
best_name, best_fn = max(
runtime["gemm_candidates"],
key=lambda item: _candidate_score(item[0], hot_shape),
)
runtime["selected_gemm_name"] = best_name
return best_fn
def _maybe_warmup(runtime: dict[str, Any], device) -> None:
device_key = (device.type, device.index)
if device_key in _WARMED_DEVICES:
return
_WARMED_DEVICES.add(device_key)
if device.type != "cuda":
return
torch = runtime["torch"]
try:
dummy = torch.zeros((64, 64), dtype=torch.bfloat16, device=device)
q_a, scale_a = _quant_mxfp4(dummy, runtime)
q_b, scale_b = _quant_mxfp4(dummy, runtime)
b_shuffle = runtime["shuffle_weight"](q_b, layout=(16, 16))
_select_gemm(runtime, 64, 64)( # Compile and cache the common path once.
q_a,
b_shuffle,
scale_a,
scale_b,
dtype=runtime["dtypes"].bf16,
bpreshuffle=True,
)
torch.cuda.synchronize(device)
except Exception:
# Warmup is opportunistic only; never fail the submission because of it.
return
def _run_gemm(
runtime: dict[str, Any],
gemm_fn: Callable[..., Any],
A,
B_shuffle,
B_scale_sh,
):
A_q, A_scale_sh = _quant_mxfp4(_maybe_contiguous(A), runtime)
return gemm_fn(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=runtime["dtypes"].bf16,
bpreshuffle=True,
)
def _run_hot_small(
runtime: dict[str, Any],
gemm_fn: Callable[..., Any],
A,
B_shuffle,
B_scale_sh,
):
return _run_gemm(
runtime,
gemm_fn,
A,
B_shuffle,
B_scale_sh,
)
def _run_hot_medium(
runtime: dict[str, Any],
gemm_fn: Callable[..., Any],
A,
B_shuffle,
B_scale_sh,
):
B_shuffle = _maybe_contiguous(B_shuffle)
B_scale_sh = _maybe_contiguous(B_scale_sh)
return _run_gemm(runtime, gemm_fn, A, B_shuffle, B_scale_sh)
def _run_hot_large(
runtime: dict[str, Any],
gemm_fn: Callable[..., Any],
A,
B_shuffle,
B_scale_sh,
):
A = _maybe_contiguous(A)
B_shuffle = _maybe_contiguous(B_shuffle)
B_scale_sh = _maybe_contiguous(B_scale_sh)
return _run_gemm(runtime, gemm_fn, A, B_shuffle, B_scale_sh)
def _run_generic(
runtime: dict[str, Any],
gemm_fn: Callable[..., Any],
A,
B_shuffle,
B_scale_sh,
):
return _run_gemm(
runtime,
gemm_fn,
A,
_maybe_contiguous(B_shuffle),
_maybe_contiguous(B_scale_sh),
)
def _select_dispatch(m: int, n: int, k: int) -> Callable[..., output_t]:
hot_shape = m in _HOT_M_VALUES and k in _HOT_K_VALUES
if hot_shape and m <= 32 and n >= 2048:
return _run_hot_small
if hot_shape and m >= 256:
return _run_hot_large
if hot_shape:
return _run_hot_medium
return _run_generic
def _build_executor(
runtime: dict[str, Any],
m: int,
n: int,
k: int,
) -> Callable[..., output_t]:
dispatch = _select_dispatch(m, n, k)
gemm_fn = _select_gemm(runtime, m, k)
def _executor(runtime: dict[str, Any], A, B_shuffle, B_scale_sh):
return dispatch(runtime, gemm_fn, A, B_shuffle, B_scale_sh)
return _executor
def custom_kernel(data: input_t) -> output_t:
"""
Optimized path: quantize `A` to MXFP4 once per call, then invoke the selected
A4W4 GEMM entry point with preshuffled `B`.
"""
runtime = _init_runtime()
A, B, _B_q, B_shuffle, B_scale_sh = data
del _B_q
m, k = A.shape
n = B.shape[0]
_maybe_warmup(runtime, A.device)
shape_key = (
m,
n,
k,
tuple(A.stride()),
tuple(B_shuffle.stride()),
A.device.type,
A.device.index,
A.dtype,
)
executor = _DISPATCH_CACHE.get(shape_key)
if executor is None:
executor = _build_executor(runtime, m, n, k)
_DISPATCH_CACHE[shape_key] = executor
return executor(runtime, A, B_shuffle, B_scale_sh)
scrolls · 291 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