submission 642218
Frank Liu · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 221 lines, June 9 Researcher Reciprocity License v1.0.
submission_mxfp4_mm_smallm_dynamic.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-642218?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:8483e06bb0aec136ba839bde9ecd6d50728d88b675147a4a98a34cc288574f25
license declaredunknown
license concludedunknown
authorsFrank Liu
imported2026-08-26
Kernel source
submission_mxfp4_mm_smallm_dynamic.py221 lines
from __future__ import annotations
from typing import Callable, Dict, Tuple
import torch
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
try:
from aiter import gemm_a4w4
except Exception: # pragma: no cover
gemm_a4w4 = getattr(aiter, "gemm_a4w4")
try:
from aiter.ops.triton.gemm_afp4wfp4 import gemm_afp4wfp4_preshuffled_weight_scales
except Exception: # pragma: no cover
gemm_afp4wfp4_preshuffled_weight_scales = None
try:
from utils import clear_l2_cache_large as clear_l2_cache
except Exception: # pragma: no cover
def clear_l2_cache() -> None:
return None
DTYPE_OUT = torch.bfloat16
RTOL = 1e-2
ATOL = 1e-2
# Cache chosen backend per (device, M, N, K)
_PLAN_CACHE: Dict[Tuple[int, int, int, int], str] = {}
_OUTBUF_CACHE: Dict[Tuple[int, int, int], torch.Tensor] = {}
@torch.no_grad()
def _shape_key(a: torch.Tensor, b_shuffle: torch.Tensor) -> Tuple[int, int, int, int]:
dev = a.device.index if a.device.index is not None else -1
return (dev, int(a.shape[0]), int(b_shuffle.shape[0]), int(a.shape[1]))
@torch.no_grad()
def _allclose(a: torch.Tensor, b: torch.Tensor) -> bool:
return a.shape == b.shape and torch.allclose(
a.to(torch.float32),
b.to(torch.float32),
rtol=RTOL,
atol=ATOL,
)
@torch.no_grad()
def _get_outbuf(m: int, n: int, device: torch.device) -> torch.Tensor:
padded_m = ((m + 31) // 32) * 32
dev = device.index if device.index is not None else -1
key = (dev, padded_m, n)
out = _OUTBUF_CACHE.get(key)
if out is None or out.shape != (padded_m, n) or out.device != device:
out = torch.empty((padded_m, n), device=device, dtype=DTYPE_OUT)
_OUTBUF_CACHE[key] = out
return out
@torch.no_grad()
def _quant_dynamic_raw(a: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
# Returns raw packed fp4 bytes + raw E8M0 scale bytes.
return dynamic_mxfp4_quant(a)
@torch.no_grad()
def _run_wrapper_dynamic(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor) -> torch.Tensor:
a_q_raw, a_scale_raw = _quant_dynamic_raw(a)
a_q = a_q_raw.view(dtypes.fp4x2)
a_scale_sh = e8m0_shuffle(a_scale_raw).view(dtypes.fp8_e8m0)
return aiter.gemm_a4w4(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
@torch.no_grad()
def _run_outbuf_dynamic(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor) -> torch.Tensor:
a_q_raw, a_scale_raw = _quant_dynamic_raw(a)
a_q = a_q_raw.view(dtypes.fp4x2)
a_scale_sh = e8m0_shuffle(a_scale_raw).view(dtypes.fp8_e8m0)
out = _get_outbuf(a.shape[0], b_shuffle.shape[0], a.device)
result = gemm_a4w4(
a_q,
b_shuffle.view(a_q.dtype),
a_scale_sh,
b_scale_sh.view(a_scale_sh.dtype),
out,
bpreshuffle=True,
)
if isinstance(result, torch.Tensor):
return result[: a.shape[0]]
return out[: a.shape[0]]
@torch.no_grad()
def _run_smallm_dynamic(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor) -> torch.Tensor:
if gemm_afp4wfp4_preshuffled_weight_scales is None:
raise RuntimeError("gemm_afp4wfp4_preshuffled_weight_scales is unavailable")
m = int(a.shape[0])
n = int(b_shuffle.shape[0])
a_q_raw, a_scale_raw = _quant_dynamic_raw(a)
if m >= 32:
a_scale_sh = e8m0_shuffle(a_scale_raw).contiguous()
a_scale_kernel = a_scale_sh.view(torch.uint8).view(a_scale_sh.shape[0] // 32, -1)
else:
a_scale_kernel = a_scale_raw[:m].contiguous().view(torch.uint8)
b_q_kernel = b_shuffle.view(torch.uint8).view(n // 16, -1)
b_scale_kernel = b_scale_sh.view(torch.uint8).view(b_scale_sh.shape[0] // 32, -1)
y = torch.empty((m, n), device=a.device, dtype=DTYPE_OUT)
gemm_afp4wfp4_preshuffled_weight_scales(
a_q_raw.contiguous().view(torch.uint8),
b_q_kernel,
a_scale_kernel,
b_scale_kernel,
DTYPE_OUT,
y,
)
return y
_BACKEND_FNS: Dict[str, Callable[[torch.Tensor, torch.Tensor, torch.Tensor], torch.Tensor]] = {
"wrapper": _run_wrapper_dynamic,
"outbuf": _run_outbuf_dynamic,
"smallm": _run_smallm_dynamic,
}
@torch.no_grad()
def _time_backend(
fn: Callable[[torch.Tensor, torch.Tensor, torch.Tensor], torch.Tensor],
a: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
reps: int = 4,
) -> float:
# One warmup after validation/compile.
_ = fn(a, b_shuffle, b_scale_sh)
torch.cuda.synchronize()
times = []
for _ in range(reps):
clear_l2_cache()
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
out = fn(a, b_shuffle, b_scale_sh)
end.record()
torch.cuda.synchronize()
del out
times.append(start.elapsed_time(end) * 1e6)
return float(sum(times) / len(times))
@torch.no_grad()
def _choose_plan(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor) -> str:
# Always validate against the official reference-style path.
y_ref = _run_wrapper_dynamic(a, b_shuffle, b_scale_sh)
valid = {"wrapper": y_ref}
# Newer AITER builds often support a slightly faster explicit-out GEMM call.
try:
y_out = _run_outbuf_dynamic(a, b_shuffle, b_scale_sh)
if _allclose(y_out, y_ref):
valid["outbuf"] = y_out
except Exception:
pass
# For M <= 64, vLLM's ROCm Quark path prefers the preshuffled-weight-scales
# Triton kernel when it is tuned for the shape. We validate numerics first,
# then benchmark it off the critical path.
if a.shape[0] <= 64 and gemm_afp4wfp4_preshuffled_weight_scales is not None:
try:
y_small = _run_smallm_dynamic(a, b_shuffle, b_scale_sh)
if _allclose(y_small, y_ref):
valid["smallm"] = y_small
except Exception:
pass
if len(valid) == 1:
return "wrapper"
timings: Dict[str, float] = {}
for name in valid.keys():
timings[name] = _time_backend(_BACKEND_FNS[name], a, b_shuffle, b_scale_sh)
# Choose the empirically fastest valid backend for this (device, M, N, K).
best_name = min(timings, key=timings.get)
return best_name
@torch.no_grad()
def custom_kernel(data: input_t) -> output_t:
a, _b, _b_q, b_shuffle, b_scale_sh = data
key = _shape_key(a, b_shuffle)
plan = _PLAN_CACHE.get(key)
if plan is None:
plan = _choose_plan(a, b_shuffle, b_scale_sh)
_PLAN_CACHE[key] = plan
return _BACKEND_FNS[plan](a, b_shuffle, b_scale_sh)
scrolls · 221 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