submission 714256
chu yifan · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 423 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-714256?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:095d5923af617409447f818ded15e208a2d3080d7267cb9f0c0b5e7367aa8be0
license declaredunknown
license concludedunknown
authorschu yifan
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 GEMM submission.Kernel source
submission.py423 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
MXFP4 GEMM submission.
Goal: minimize end-to-end latency of:
bf16 A -> (per-1x32) MXFP4 quant -> A4W4 GEMM (MXFP4 A x MXFP4 W) -> bf16 C
Strategy:
- Prefer AITER Triton fused path (quantize-A inside GEMM) when available.
- Otherwise fall back to the reference path (dynamic_mxfp4_quant + gemm_a4w4).
- Auto-tune per (M,N,K) shape once per process to pick the faster path without
guessing a fixed threshold.
"""
from task import input_t, output_t
import torch
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.ops.triton.gemm.basic.gemm_a16wfp4 import (
gemm_a16wfp4_preshuffle_ as _gemm_a16wfp4_preshuffle_,
)
from aiter.ops.triton.utils.common_utils import serialize_dict as _serialize_dict
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
_get_config as _get_gemm_a16wfp4_config,
)
except Exception:
_gemm_a16wfp4_preshuffle_ = None
_serialize_dict = None
_get_gemm_a16wfp4_config = None
_IMPL_CACHE: dict[tuple, str] = {}
_FUSED_CONFIG_CACHE: dict[tuple, str] = {}
_FUSED_PICK_CACHE: dict[tuple, str] = {}
_FUSED_CONFIG_OVERRIDES: dict[tuple[int, int, int], list[dict]] = {
(4, 2880, 512): [
{
"BLOCK_SIZE_M": 4,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
{
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
],
(16, 2112, 7168): [
{
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 14,
},
{
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 7,
},
],
(32, 4096, 512): [
{
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
{
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
],
(32, 2880, 512): [
{
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
{
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
],
(64, 7168, 2048): [
{
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
{
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
{
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
],
}
_FORCED_FUSED_CONFIGS: dict[tuple[int, int, int], dict] = {
(16, 2112, 7168): {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 7,
},
(64, 7168, 2048): {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(256, 3072, 1536): {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 256,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
}
def _quant_mxfp4(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
bs_e8m0 = e8m0_shuffle(bs_e8m0)
return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)
def _get_fused_config_hash(m: int, n: int, k_packed: int) -> str:
key = (m, n, k_packed)
cached = _FUSED_CONFIG_CACHE.get(key)
if cached is not None:
return cached
if _get_gemm_a16wfp4_config is None or _serialize_dict is None:
raise RuntimeError("gemm_a16wfp4_preshuffle config helpers unavailable")
config, _ = _get_gemm_a16wfp4_config(m, n, k_packed, True)
config_hash = _serialize_dict(config)
_FUSED_CONFIG_CACHE[key] = config_hash
return config_hash
def _serialize_fused_config(config: dict) -> str:
if _serialize_dict is None:
raise RuntimeError("serialize_dict unavailable")
return _serialize_dict(config)
def _get_fused_config_candidates(m: int, n: int, k: int) -> list[str]:
candidates = [_get_fused_config_hash(m, n, k // 2)]
for config in _FUSED_CONFIG_OVERRIDES.get((int(m), int(n), int(k)), []):
candidates.append(_serialize_fused_config(config))
deduped: list[str] = []
seen: set[str] = set()
for config_hash in candidates:
if config_hash not in seen:
deduped.append(config_hash)
seen.add(config_hash)
return deduped
def _run_fused(
A: torch.Tensor,
B_shuffle: torch.Tensor,
B_scale_sh: torch.Tensor,
config_hash: str | None = None,
) -> torch.Tensor:
if _gemm_a16wfp4_preshuffle_ is None:
raise RuntimeError("gemm_a16wfp4_preshuffle unavailable")
m, k = A.shape
n, k_packed = B_shuffle.shape
w_preshuf = B_shuffle.view(torch.uint8).view(n // 16, k_packed * 16)
scale_bytes = B_scale_sh.view(torch.uint8)
w_scales_preshuf = scale_bytes.view(
scale_bytes.shape[0] // 32,
scale_bytes.shape[1] * 32,
)
if config_hash is None:
config_hash = _get_fused_config_hash(m, n, k_packed)
return _gemm_a16wfp4_preshuffle_(
A,
w_preshuf,
w_scales_preshuf,
prequant=True,
dtype=dtypes.bf16,
y=None,
config=config_hash,
skip_reduce=False,
)
def _run_fallback(
A: torch.Tensor,
B_shuffle: torch.Tensor,
B_scale_sh: torch.Tensor,
) -> torch.Tensor:
A_q, A_scale_sh = _quant_mxfp4(A)
return aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
def _time_cuda_us(fn) -> float:
torch.cuda.synchronize()
start_event = torch.cuda.Event(enable_timing=True)
end_event = torch.cuda.Event(enable_timing=True)
start_event.record()
out = fn()
end_event.record()
torch.cuda.synchronize()
del out
return float(start_event.elapsed_time(end_event) * 1e3)
def _pick_impl(
*,
A: torch.Tensor,
B_shuffle: torch.Tensor,
B_scale_sh: torch.Tensor,
) -> str:
m, k = A.shape
n, _ = B_shuffle.shape
key = (A.device, int(m), int(n), int(k))
cached = _IMPL_CACHE.get(key)
if cached is not None:
return cached
if _gemm_a16wfp4_preshuffle_ is None:
_IMPL_CACHE[key] = "fallback"
return "fallback"
fused_candidates: list[str] = []
for config_hash in _get_fused_config_candidates(m, n, k):
try:
_run_fused(A, B_shuffle, B_scale_sh, config_hash=config_hash)
fused_candidates.append(config_hash)
except Exception:
continue
if not fused_candidates:
_IMPL_CACHE[key] = "fallback"
return "fallback"
# Warm-up both codepaths so one-time compilation/module loads don't skew timing.
try:
_run_fallback(A, B_shuffle, B_scale_sh)
except Exception:
# If fallback path fails for any reason, stick to fused (already verified).
_IMPL_CACHE[key] = "fused"
return "fused"
torch.cuda.synchronize()
# Choose the faster of fused vs (quant + a4w4) using a few short runs.
best_fused_hash = min(
fused_candidates,
key=lambda config_hash: min(
_time_cuda_us(
lambda config_hash=config_hash: _run_fused(
A, B_shuffle, B_scale_sh, config_hash=config_hash
)
)
for _ in range(3)
),
)
fused_us = min(
_time_cuda_us(
lambda: _run_fused(A, B_shuffle, B_scale_sh, config_hash=best_fused_hash)
)
for _ in range(3)
)
fallback_us = min(
_time_cuda_us(lambda: _run_fallback(A, B_shuffle, B_scale_sh)) for _ in range(3)
)
chosen = "fused" if fused_us <= fallback_us else "fallback"
_IMPL_CACHE[key] = chosen
if chosen == "fused":
_FUSED_PICK_CACHE[key] = best_fused_hash
return chosen
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_shuffle.shape
forced_fused = _FORCED_FUSED_CONFIGS.get((int(m), int(n), int(k)))
if forced_fused is not None:
return _run_fused(
A,
B_shuffle,
B_scale_sh,
config_hash=_serialize_fused_config(forced_fused),
)
key = (A.device, int(m), int(n), int(k))
impl = _pick_impl(A=A, B_shuffle=B_shuffle, B_scale_sh=B_scale_sh)
if impl == "fused":
return _run_fused(
A,
B_shuffle,
B_scale_sh,
config_hash=_FUSED_PICK_CACHE.get(key),
)
return _run_fallback(A, B_shuffle, B_scale_sh)
scrolls · 423 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