submission 619457
Coalwood · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 372 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-619457?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:6e1ab7a4f6ac321f45e861a837e927fe65adafee848492a0977f45a639d98553
license declaredunknown
license concludedunknown
authorsCoalwood
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 per-1x32 quant on A, then A4W4 GEMM on MI355X.num-warps = 1
num_warps = 1split-k
_A4W4_TUNED_HEADER = "cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n"stages = 1
NUM_STAGES=1,tile-m = 32
BLOCK_M=32,tile-n = 8
BLOCK_N=8,Kernel source
submission.py372 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
MXFP4 per-1x32 quant on A, then A4W4 GEMM on MI355X.
Formal submission path:
- exact-shape asm dispatch for the fixed benchmark shapes
- specialized quant+shuffle path for those same shapes
- unified aiter fallback for everything else
"""
import importlib.util
import os
from task import input_t, output_t
_ASM_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_ASM_32X128_CO = "f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co"
_A4W4_TUNED_HEADER = "cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n"
_A4W4_TUNED_ROWS = [
(256, 4, 2880, 512, 21, 0, "0.1000", _ASM_32X128, "0.0", "0.0", "0.0"),
(256, 16, 2112, 7168, 21, 0, "0.1000", _ASM_32X128, "0.0", "0.0", "0.0"),
(256, 32, 4096, 512, 21, 0, "0.1000", _ASM_32X128, "0.0", "0.0", "0.0"),
(256, 32, 2880, 512, 21, 0, "0.1000", _ASM_32X128, "0.0", "0.0", "0.0"),
]
_SPECIALIZED_QUANT_ENV_VAR = "MXFP4_ENABLE_SPECIALIZED_QUANT"
_A4W4_TUNED_OVERRIDE = None
_UNIFIED_PLAN = {"kind": "unified"}
_SHAPE_PLANS = {
(4, 2880, 512): {
"kind": "asm",
"kernel_name": _ASM_32X128,
"co_name": _ASM_32X128_CO,
"log2_k_split": None,
"specialized_quant_candidate": True,
},
(16, 2112, 7168): {
"kind": "asm",
"kernel_name": _ASM_32X128,
"co_name": _ASM_32X128_CO,
"log2_k_split": None,
"specialized_quant_candidate": True,
},
(32, 4096, 512): {
"kind": "asm",
"kernel_name": _ASM_32X128,
"co_name": _ASM_32X128_CO,
"log2_k_split": None,
"specialized_quant_candidate": True,
},
(32, 2880, 512): {
"kind": "asm",
"kernel_name": _ASM_32X128,
"co_name": _ASM_32X128_CO,
"log2_k_split": None,
"specialized_quant_candidate": True,
},
}
_RUNTIME = None
_SPECIALIZED_QUANT_RUNTIME = None
_SPECIALIZED_QUANT_WORKSPACES = {}
_SPECIALIZED_QUANT_ERROR = None
_SPECIALIZED_QUANT_INFO_PRINTED = False
def _select_gemm_plan(m: int, n: int, k: int):
return _SHAPE_PLANS.get((m, n, k), _UNIFIED_PLAN)
def _find_aiter_config_path():
try:
spec = importlib.util.find_spec("aiter")
except (ImportError, ValueError):
spec = None
locations = getattr(spec, "submodule_search_locations", None) if spec is not None else None
if locations:
return os.path.join(locations[0], "configs", "a4w4_blockscale_tuned_gemm.csv")
runtime = globals().get("_RUNTIME")
if runtime is None:
return None
_, aiter, _, _, _ = runtime
package_file = getattr(aiter, "__file__", None)
if package_file:
return os.path.join(
os.path.dirname(os.path.abspath(package_file)),
"configs",
"a4w4_blockscale_tuned_gemm.csv",
)
return None
def _render_a4w4_tuned_override():
rows = ["{},{},{},{},{},{},{},{},{},{},{}".format(*row) for row in _A4W4_TUNED_ROWS]
return _A4W4_TUNED_HEADER + "\n".join(rows) + "\n"
def _ensure_a4w4_tuned_override():
global _A4W4_TUNED_OVERRIDE
override_path = _A4W4_TUNED_OVERRIDE or "/tmp/mxfp4_a4w4_tuned_override.csv"
content = _render_a4w4_tuned_override()
try:
existing = None
if os.path.exists(override_path):
with open(override_path, "r", encoding="utf-8") as handle:
existing = handle.read()
if existing != content:
with open(override_path, "w", encoding="utf-8") as handle:
handle.write(content)
except Exception:
return None
_A4W4_TUNED_OVERRIDE = override_path
default_path = _find_aiter_config_path()
if not default_path:
return override_path
os.environ["AITER_CONFIG_GEMM_A4W4"] = os.pathsep.join([default_path, override_path])
return override_path
def _get_runtime():
global _RUNTIME
if _RUNTIME is None:
_ensure_a4w4_tuned_override()
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
_RUNTIME = (torch, aiter, dtypes, dynamic_mxfp4_quant, e8m0_shuffle)
return _RUNTIME
def _alloc_gemm_output(a_q, dtypes, m: int, n: int, zero_init: bool = False):
out_rows = ((m + 31) // 32) * 32
if zero_init:
return a_q.new_zeros((out_rows, n), dtype=dtypes.bf16)
return a_q.new_empty((out_rows, n), dtype=dtypes.bf16)
def _quant_mxfp4(x, dtypes, dynamic_mxfp4_quant, e8m0_shuffle):
x_fp4, scale = dynamic_mxfp4_quant(x)
scale = e8m0_shuffle(scale)
return x_fp4.view(dtypes.fp4x2), scale.view(dtypes.fp8_e8m0)
def _specialized_quant_runtime_enabled(torch, plan):
if not plan.get("specialized_quant_candidate"):
return False
if os.environ.get(_SPECIALIZED_QUANT_ENV_VAR, "1") == "0":
return False
return getattr(getattr(torch, "version", None), "hip", None) is not None
def _get_specialized_quant_runtime():
global _SPECIALIZED_QUANT_RUNTIME
if _SPECIALIZED_QUANT_RUNTIME is not None:
return _SPECIALIZED_QUANT_RUNTIME
import triton
import triton.language as tl
from aiter.ops.triton._triton_kernels.quant.quant import _dynamic_mxfp4_quant_kernel
@triton.jit
def _shuffle_e8m0_scale_kernel(
src_ptr,
dst_ptr,
stride_src_m,
stride_src_n,
M,
N_VALID,
N_PAD,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
mask = (offs_m[:, None] < M) & (offs_n[None, :] < N_VALID)
src_offs = offs_m[:, None] * stride_src_m + offs_n[None, :] * stride_src_n
vals = tl.load(src_ptr + src_offs, mask=mask, other=127)
g0 = offs_m[:, None] // 32
rem_m = offs_m[:, None] % 32
g1 = rem_m // 16
g2 = rem_m % 16
g3 = offs_n[None, :] // 8
rem_n = offs_n[None, :] % 8
g4 = rem_n // 4
g5 = rem_n % 4
dst_offs = (
g1
+ g4 * 2
+ g2 * 4
+ g5 * 64
+ g3 * 256
+ g0 * (32 * N_PAD)
)
tl.store(dst_ptr + dst_offs, vals, mask=mask)
_SPECIALIZED_QUANT_RUNTIME = (triton, _dynamic_mxfp4_quant_kernel, _shuffle_e8m0_scale_kernel)
return _SPECIALIZED_QUANT_RUNTIME
def _get_specialized_quant_workspace(torch, x):
m, n = x.shape
scale_n_valid = (n + 31) // 32
scale_n_pad = ((scale_n_valid + 7) // 8) * 8
scale_m_pad = ((m + 255) // 256) * 256
cache_key = (tuple(x.shape), str(getattr(x, "device", "")), str(getattr(x, "dtype", "")))
workspace = _SPECIALIZED_QUANT_WORKSPACES.get(cache_key)
if workspace is None:
workspace = {
"a_q_raw": torch.empty((m, n // 2), dtype=torch.uint8, device=x.device),
"a_scale_raw": torch.empty((m, scale_n_valid), dtype=torch.uint8, device=x.device),
"a_scale_shuffled_raw": torch.full(
(scale_m_pad, scale_n_pad),
127,
dtype=torch.uint8,
device=x.device,
),
"scale_n_valid": scale_n_valid,
"scale_n_pad": scale_n_pad,
}
_SPECIALIZED_QUANT_WORKSPACES[cache_key] = workspace
return workspace
def _quant_mxfp4_specialized(torch, dtypes, x):
triton, quant_kernel, shuffle_kernel = _get_specialized_quant_runtime()
workspace = _get_specialized_quant_workspace(torch, x)
m, n = x.shape
a_q_raw = workspace["a_q_raw"]
a_scale_raw = workspace["a_scale_raw"]
a_scale_shuffled_raw = workspace["a_scale_shuffled_raw"]
scale_n_valid = workspace["scale_n_valid"]
scale_n_pad = workspace["scale_n_pad"]
if m <= 32:
num_iter = 1
block_size_m = triton.next_power_of_2(m)
block_size_n = 32
num_warps = 1
else:
num_iter = 4
block_size_m = 64
block_size_n = 64
num_warps = 4
if n <= 16384:
block_size_m = 32
block_size_n = 128
if n <= 1024:
num_iter = 1
block_size_n = min(256, triton.next_power_of_2(n))
block_size_n = max(32, block_size_n)
block_size_m = min(8, triton.next_power_of_2(m))
num_warps = 4
grid = (triton.cdiv(m, block_size_m), triton.cdiv(n, block_size_n * num_iter))
quant_kernel[grid](
x,
a_q_raw,
a_scale_raw,
*x.stride(),
*a_q_raw.stride(),
*a_scale_raw.stride(),
M=m,
N=n,
MXFP4_QUANT_BLOCK_SIZE=32,
SCALING_MODE=0,
NUM_ITER=num_iter,
BLOCK_SIZE_M=block_size_m,
BLOCK_SIZE_N=block_size_n,
NUM_STAGES=1,
num_warps=num_warps,
waves_per_eu=0,
num_stages=1,
)
shuffle_grid = (triton.cdiv(m, 32), triton.cdiv(scale_n_valid, 8))
shuffle_kernel[shuffle_grid](
a_scale_raw,
a_scale_shuffled_raw,
*a_scale_raw.stride(),
M=m,
N_VALID=scale_n_valid,
N_PAD=scale_n_pad,
BLOCK_M=32,
BLOCK_N=8,
)
return a_q_raw.view(dtypes.fp4x2), a_scale_shuffled_raw.view(dtypes.fp8_e8m0)
def _maybe_quant_mxfp4_specialized(torch, dtypes, x, plan):
global _SPECIALIZED_QUANT_ERROR, _SPECIALIZED_QUANT_INFO_PRINTED
if not _specialized_quant_runtime_enabled(torch, plan):
return None
try:
result = _quant_mxfp4_specialized(torch, dtypes, x)
except Exception:
if _SPECIALIZED_QUANT_ERROR is None:
_SPECIALIZED_QUANT_ERROR = True
try:
import traceback
print("[mxfp4 quant] falling back after error:")
traceback.print_exc()
except Exception:
pass
return None
if not _SPECIALIZED_QUANT_INFO_PRINTED:
_SPECIALIZED_QUANT_INFO_PRINTED = True
try:
print("[mxfp4 quant] using specialized quant+shuffle path")
except Exception:
pass
return result
def _run_gemm_asm(aiter, dtypes, a_q, b_shuffle, a_scale_sh, b_scale_sh, m: int, n: int, plan):
out = _alloc_gemm_output(a_q, dtypes, m, n)
aiter.gemm_a4w4_asm(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
out,
plan["kernel_name"],
bpreshuffle=True,
log2_k_split=plan["log2_k_split"],
)
return out[:m]
def custom_kernel(data: input_t) -> output_t:
torch, aiter, dtypes, dynamic_mxfp4_quant, e8m0_shuffle = _get_runtime()
a, b, _b_q, b_shuffle, b_scale_sh = data
a = a.contiguous()
b = b.contiguous()
m, k = a.shape
n, _ = b.shape
plan = _select_gemm_plan(m, n, k)
specialized = _maybe_quant_mxfp4_specialized(torch, dtypes, a, plan)
if specialized is None:
a_q, a_scale_sh = _quant_mxfp4(a, dtypes, dynamic_mxfp4_quant, e8m0_shuffle)
else:
a_q, a_scale_sh = specialized
if plan.get("kind") == "asm":
return _run_gemm_asm(
aiter,
dtypes,
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
m,
n,
plan,
)
return aiter.gemm_a4w4(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
scrolls · 372 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