submission 714808
vuxml · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 460 lines, June 9 Researcher Reciprocity License v1.0.
submission_v16e_auto.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-714808?include=source"interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, fp32, fp8_e8m0, int32, 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:e8c7e7842df60939c32ff7ed56be1944bf95ff93383c194d27d99976433be87e
license declaredunknown
license concludedunknown
authorsvuxml
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
v16e: L2-cold autotune — raw cktile + raw ck2stages, same 5-kernelfp4
cktile = alternate GEMM backend consuming SAME pre-quantized fp4 asnum-warps = 4
_touch_k[(triton.cdiv(n, 8192),)](flat, n, BLK=8192, num_warps=4)tile-m = 16
BLOCK_SIZE_M=16, BLOCK_SIZE_N=4, TOPK=topk,tile-n = 4
BLOCK_SIZE_M=16, BLOCK_SIZE_N=4, TOPK=topk,Kernel source
submission_v16e_auto.py460 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
"""
v16e: L2-cold autotune — raw cktile + raw ck2stages, same 5-kernel
pipeline, swap GEMM backend. NaN-safe correctness check.
cktile = alternate GEMM backend consuming SAME pre-quantized fp4 as
ck2stages. Both are 5-kernel (sort+q1+g1+q2+g2); cktile's tile schedule
is faster at small M (observed: 94µs vs 131µs at bs=16/E=257). Raw
pybind into prealloc for both; sk1=1 only for cktile (sk>1 needs
module_activation for post-SwiGLU → extra 23s build + alloc).
SEARCH (per shape, 8s budget, ck-first so guaranteed valid baseline):
B. ck2stages: bm∈{32,64,128} × s1_kn∈CSV∪{""} × nt × s2_kn × sk2 × nt
A. cktile: bm∈{16,32,64} (sk1=1 only)
+ prefetch-wrap when W<200MB.
"""
import os
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
import sys, re, csv, glob, time as _tm, math, warnings
warnings.filterwarnings("ignore")
import torch
import triton
import triton.language as tl
import aiter
from aiter import ActivationType, QuantType, dtypes
import aiter.fused_moe as _fm
from aiter.fused_moe import (
fused_moe, moe_sorting, fused_dynamic_mxfp4_quant_moe_sort,
cktile_moe_stage1,
)
_SILU, _PER1X32 = ActivationType.Silu, QuantType.per_1x32
_e8m0, _fp4x2, _bf16 = dtypes.fp8_e8m0, dtypes.fp4x2, torch.bfloat16
_L = lambda m: print(m, file=sys.stderr, flush=True)
_orig_sw = sys.stderr.write
sys.stderr.write = lambda s: (len(s) if "ck kernel not found" in s
else _orig_sw(s))
_quant_k = fused_dynamic_mxfp4_quant_moe_sort.__globals__.get(
"_fused_dynamic_mxfp4_quant_moe_sort_kernel")
def _quant_direct(x, x_fp4, sid, nvi, sc5d, M, N, Ls, topk):
scaleN = N // 32
num_pid = (triton.cdiv(M, 128) * scaleN
+ triton.cdiv(Ls, 32) * triton.cdiv(scaleN, 8))
_quant_k[(num_pid,)](
x, x_fp4, sid, nvi, sc5d, M, N, scaleN,
x.stride(0), x.stride(1), x_fp4.stride(0), x_fp4.stride(1),
sc5d.stride(0), sc5d.stride(1), sc5d.stride(2),
sc5d.stride(3), sc5d.stride(4),
token_num=M, M_i=M, N_i=scaleN,
MXFP4_QUANT_BLOCK_SIZE=32, BLOCK_SIZE_Mx=128,
BLOCK_SIZE_M=16, BLOCK_SIZE_N=4, TOPK=topk,
)
@triton.jit
def _touch_k(P, N, BLK: tl.constexpr):
pid = tl.program_id(0)
off = pid * BLK + tl.arange(0, BLK)
_ = tl.load(P + off, mask=off < N, other=0, cache_modifier=".cg")
def _prefetch(tensors):
for t in tensors:
flat = t.reshape(-1).view(torch.int32)
n = flat.numel()
_touch_k[(triton.cdiv(n, 8192),)](flat, n, BLK=8192, num_warps=4)
def _load_csv_kernels():
s1, s2 = set(), set()
for p in glob.glob("/home/runner/aiter/aiter/configs/**/*.csv",
recursive=True):
try:
with open(p, newline="") as f:
for row in csv.DictReader(f):
for v in row.values():
if not isinstance(v, str):
continue
v = v.strip()
if "moe_ck2stages_gemm1_" in v and "FP4X2_FP4X2" in v:
s1.add(v)
elif "moe_ck2stages_gemm2_" in v and "FP4X2_FP4X2" in v:
s2.add(v)
except Exception:
pass
return s1, s2
_CSV_S1, _CSV_S2 = _load_csv_kernels()
def _bm_of(n):
m = re.search(r"_\d+x(\d+)x\d+x\d+_", n)
return int(m.group(1)) if m else -1
def _find_modules():
r = {}
for mn in list(sys.modules):
m = sys.modules.get(mn)
if m is None:
continue
if "moe_ck2stages" in mn and "fp4x2_fp4x2" in mn \
and hasattr(m, "ck_moe_stage1"):
r["ck"] = m
elif "module_moe_sorting" in mn and hasattr(m, "moe_sorting_fwd"):
r["sort"] = m
elif "module_moe_cktile" in mn and hasattr(m, "cktile_moe_gemm1"):
r["cktile"] = m
return r
_l2_buf = None
def _cold(fn, n=7):
global _l2_buf
if _l2_buf is None:
_l2_buf = torch.empty(384 * 1024 * 1024, dtype=torch.int8,
device="cuda")
fn(); fn()
torch.cuda.synchronize()
ts = []
for _ in range(n):
_l2_buf.zero_()
torch.cuda.synchronize()
e0, e1 = torch.cuda.Event(True), torch.cuda.Event(True)
e0.record(); fn(); e1.record()
torch.cuda.synchronize()
ts.append(e0.elapsed_time(e1) * 1000)
ts.sort()
core = ts[1:-1]
return sum(core) / len(core)
_cfg: dict = {}
_mods = {}
_L2_CAP = 200 * 1024 * 1024
_TBUDGET = 8.0
def _sc5d(Ls, N):
scN = N // 32
return (triton.cdiv(Ls, 32), triton.cdiv(scN, 8), 4, 16, 4)
def _warmup_modules(hs, w1sh, w2sh, s1sh, s2sh, tw, ti, config):
global _mods
if _mods:
return
hp = config["d_hidden_pad"] - config["d_hidden"]
ip = config["d_expert_pad"] - config["d_expert"]
_ = fused_moe(hs, w1sh, w2sh, tw, ti, expert_mask=None,
activation=_SILU, quant_type=_PER1X32,
doweight_stage1=False, w1_scale=s1sh, w2_scale=s2sh,
a1_scale=None, a2_scale=None,
hidden_pad=hp, intermediate_pad=ip)
torch.cuda.synchronize()
try:
E = config["n_routed_experts"] + config["n_shared_experts"]
tk = config["n_experts_per_token"] + config["n_shared_experts"]
sid, swt, seid, nvi, _ = moe_sorting(
ti, tw, E, config["d_hidden"], _bf16, 32)
_a2 = cktile_moe_stage1(
hs, w1sh, w2sh, sid, seid, nvi, None, tk,
block_m=32, a1_scale=None, w1_scale=s1sh.view(_e8m0),
sorted_weights=None)
torch.cuda.synchronize()
except Exception as ex:
_L(f"[v16e] cktile warmup EXC: {type(ex).__name__}: "
f"{str(ex)[:160]}")
_mods.update(_find_modules())
_L(f"[v16e] modules: {sorted(_mods.keys())} "
f"CSV s1={len(_CSV_S1)} s2={len(_CSV_S2)}")
def _build(data, config):
(hs, _, _, _, _, w1sh, w2sh, s1sh, s2sh, tw, ti, _) = data
M, dh, dhp = config["bs"], config["d_hidden"], config["d_hidden_pad"]
de, dep = config["d_expert"], config["d_expert_pad"]
tk = config["n_experts_per_token"] + config["n_shared_experts"]
E = config["n_routed_experts"] + config["n_shared_experts"]
hp, ip = dhp - dh, dep - de
dev = hs.device
K = dhp
qt, act = int(_PER1X32), int(_SILU)
_warmup_modules(hs, w1sh, w2sh, s1sh, s2sh, tw, ti, config)
t0 = _tm.time()
ref = fused_moe(hs, w1sh, w2sh, tw, ti, expert_mask=None,
activation=_SILU, quant_type=_PER1X32,
doweight_stage1=False, w1_scale=s1sh, w2_scale=s2sh,
a1_scale=None, a2_scale=None,
hidden_pad=hp, intermediate_pad=ip)
ref_f = ref.float()
torch.cuda.synchronize()
mod_ck = _mods.get("ck")
mod_ckt = _mods.get("cktile")
rsort = getattr(_mods.get("sort"), "moe_sorting_fwd", None)
rs1 = getattr(mod_ck, "ck_moe_stage1", None) if mod_ck else None
rs2 = getattr(mod_ck, "ck_moe_stage2", None) if mod_ck else None
ctg1 = getattr(mod_ckt, "cktile_moe_gemm1", None) if mod_ckt else None
ctg2 = getattr(mod_ckt, "cktile_moe_gemm2", None) if mod_ckt else None
W_bytes = (w1sh.numel() + w2sh.numel() + s1sh.numel() + s2sh.numel())
can_pf = W_bytes <= _L2_CAP
np1 = (ip // 64 * 64) * 2
kp1 = hp // 128 * 128
np2 = hp // 64 * 64
kp2 = ip // 128 * 128
_L(f"\n[v16e M={M} E={E} de={de}] ck={rs1 is not None} "
f"ckt={ctg1 is not None} W={W_bytes/1e6:.0f}MB pf={can_pf}")
best = [float("inf"), "none", None]
exc_cnt = {}
def _try(hot, desc):
if _tm.time() - t0 > _TBUDGET:
return None
try:
out = hot(hs, w1sh, w2sh, s1sh, s2sh, tw, ti)
torch.cuda.synchronize()
e = (out.float() - ref_f).abs().max().item()
if not (e <= 4e-2):
k = f"err({desc[:24]})={e:.2f}"
exc_cnt[k] = exc_cnt.get(k, 0) + 1
return None
t = _cold(lambda: hot(hs, w1sh, w2sh, s1sh, s2sh, tw, ti))
if t < best[0]:
best[0], best[1], best[2] = t, desc, hot
_L(f" * {t:7.1f}us {desc} (err={e:.3f})")
return t
except Exception as ex:
k = f"{type(ex).__name__}:{str(ex)[:80]}"
exc_cnt[k] = exc_cnt.get(k, 0) + 1
return None
Bcache = {}
def _bufs(bm):
if bm in Bcache:
return Bcache[bm]
sid0, swt0, seid0, nvi0, _ = moe_sorting(ti, tw, E, dh, _bf16, bm)
Ls = sid0.shape[0]
B = {
"sid": torch.empty_like(sid0),
"swt": torch.empty_like(swt0),
"seid": torch.empty_like(seid0),
"nvi": torch.empty_like(nvi0),
"mbuf": torch.empty((M, dh), dtype=_bf16, device=dev),
"a2": torch.empty((M, tk, de), dtype=_bf16, device=dev),
"a1f": torch.empty((M, K // 2), dtype=torch.uint8,
device=dev),
"a1sc5d": torch.empty(_sc5d(Ls, K), dtype=torch.uint8,
device=dev),
"a2f": torch.empty((M * tk, de // 2), dtype=torch.uint8,
device=dev),
"a2sc5d": torch.empty(_sc5d(Ls, de), dtype=torch.uint8,
device=dev),
"Ls": Ls,
}
B["a1f_v"] = B["a1f"].view(_fp4x2)
B["a1sc_v"] = B["a1sc5d"].view(_e8m0).view(-1, K // 32)
B["a2f_v"] = B["a2f"].view(_fp4x2).view(M, tk, de // 2)
B["a2sc_v"] = B["a2sc5d"].view(_e8m0).view(-1, de // 32)
B["a2_flat"] = B["a2"].view(M * tk, de)
Bcache[bm] = B
return B
# ── Generic 5-kernel pipeline: sort → q1 → G1 → q2 → G2 ────────────
# backend="ck": G1=rs1(a1_fp4,...,kn,sk,nt), G2=rs2(...)
# backend="ckt": G1=ctg1(a1_fp4,w1,a2,...,a1sc,w1sc,act,bm,1)
# G2=ctg2(a2_fp4,w2,mbuf,...,swt,a2sc,w2sc,act,bm)
def _mk_hot(bm, backend, s1p, s2p, do_pf):
B = _bufs(bm)
Ls = B["Ls"]
if backend == "ck":
s1_kn, sk1, nt1 = s1p
s2_kn, sk2, nt2 = s2p
def hot(hs, w1sh, w2sh, s1sh, s2sh, tw, ti):
s1sv = s1sh.view(_e8m0)
s2sv = s2sh.view(_e8m0)
if do_pf:
_prefetch([w1sh, w2sh, s1sh, s2sh])
rsort(ti, tw, B["sid"], B["swt"], B["seid"], B["nvi"],
B["mbuf"], E, bm)
_quant_direct(hs, B["a1f"], B["sid"], B["nvi"],
B["a1sc5d"], M, K, Ls, 1)
rs1(B["a1f_v"], w1sh, w2sh, B["sid"], B["seid"],
B["nvi"], B["a2"], tk, s1_kn, s1sv, B["a1sc_v"],
bm, None, qt, act, sk1, nt1, None, True)
_quant_direct(B["a2_flat"], B["a2f"], B["sid"],
B["nvi"], B["a2sc5d"],
M * tk, de, Ls, tk)
rs2(B["a2f_v"], w1sh, w2sh, B["sid"], B["seid"],
B["nvi"], B["mbuf"], tk, s2_kn, s2sv, B["a2sc_v"],
bm, B["swt"], qt, act, sk2, nt2, None, True)
return B["mbuf"]
return hot
else:
# cktile: raw gemm1/gemm2 into prealloc. sk1=1 → SwiGLU
# fused in-kernel, output [M,tk,de]. Positional args:
# gemm*(XQ,WQ,Y,sid,seid,nvi,tk,np,kp,swt,xsc,wsc,bias,act,bm,sk)
def hot(hs, w1sh, w2sh, s1sh, s2sh, tw, ti):
s1sv = s1sh.view(_e8m0)
s2sv = s2sh.view(_e8m0)
if do_pf:
_prefetch([w1sh, w2sh, s1sh, s2sh])
rsort(ti, tw, B["sid"], B["swt"], B["seid"], B["nvi"],
B["mbuf"], E, bm)
_quant_direct(hs, B["a1f"], B["sid"], B["nvi"],
B["a1sc5d"], M, K, Ls, 1)
ctg1(B["a1f_v"], w1sh, B["a2"], B["sid"], B["seid"],
B["nvi"], tk, np1, kp1, None, B["a1sc_v"], s1sv,
None, act, bm, 1)
_quant_direct(B["a2_flat"], B["a2f"], B["sid"],
B["nvi"], B["a2sc5d"],
M * tk, de, Ls, tk)
ctg2(B["a2f_v"], w2sh, B["mbuf"], B["sid"], B["seid"],
B["nvi"], tk, np2, kp2, B["swt"], B["a2sc_v"],
s2sv, None, act, bm)
return B["mbuf"]
return hot
# ── Fallback (also establishes a valid baseline if all else fails) ─
def _mk_fm():
def hot(hs, w1sh, w2sh, s1sh, s2sh, tw, ti):
return fused_moe(hs, w1sh, w2sh, tw, ti, expert_mask=None,
activation=_SILU, quant_type=_PER1X32,
doweight_stage1=False, w1_scale=s1sh,
w2_scale=s2sh, a1_scale=None, a2_scale=None,
hidden_pad=hp, intermediate_pad=ip)
return hot
_try(_mk_fm(), "fm[default]")
# ── B: raw ck2stages sweep ─────────────────────────────────────────
if rs1 and rs2 and rsort and _quant_k:
for bm in (32, 64, 128):
s1_kns = [""] + sorted(n for n in _CSV_S1 if _bm_of(n) == bm)
s2_kns = [""] + sorted(n for n in _CSV_S2 if _bm_of(n) == bm)
best_s1_bm = (float("inf"), "", 1, False)
for kn in s1_kns:
for nt in (False, True):
t = _try(
_mk_hot(bm, "ck", (kn, 1, nt),
("", 1, False), False),
f"ck[bm={bm} s1={(kn[18:38] or 'def')}"
f"/nt{int(nt)} s2=def]")
if t is not None and t < best_s1_bm[0]:
best_s1_bm = (t, kn, 1, nt)
if best_s1_bm[0] == float("inf"):
continue
_, bk1, bsk1, bnt1 = best_s1_bm
best_s2_bm = (float("inf"), "", 1, False)
for kn in s2_kns:
for sk in (1, 2, 4):
for nt in (False, True):
t = _try(
_mk_hot(bm, "ck", (bk1, bsk1, bnt1),
(kn, sk, nt), False),
f"ck[bm={bm} s1=best "
f"s2={(kn[18:38] or 'def')}/sk{sk}"
f"/nt{int(nt)}]")
if t is not None and t < best_s2_bm[0]:
best_s2_bm = (t, kn, sk, nt)
if can_pf and best_s2_bm[0] < float("inf"):
_, bk2, bsk2, bnt2 = best_s2_bm
_try(_mk_hot(bm, "ck", (bk1, bsk1, bnt1),
(bk2, bsk2, bnt2), True),
f"ck[bm={bm} best pf]")
# ── A: raw cktile sweep (sk=1 only) ────────────────────────────────
if ctg1 and ctg2 and rsort and _quant_k:
for bm in (16, 32, 64):
for pf in ((False, True) if can_pf else (False,)):
_try(_mk_hot(bm, "ckt", None, None, pf),
f"ckt[bm={bm}{' pf' if pf else ''}]")
# ── Mixed: cktile-g1 + ck-s2 (cktile stage1 faster, CK stage2 has
# explicit large-bn instances). Share bm; a2 format identical. ──
if ctg1 and rs2 and rsort and _quant_k:
def _mk_mix(bm, s2_kn, sk2, nt2, do_pf):
B = _bufs(bm)
Ls = B["Ls"]
def hot(hs, w1sh, w2sh, s1sh, s2sh, tw, ti):
s1sv = s1sh.view(_e8m0)
s2sv = s2sh.view(_e8m0)
if do_pf:
_prefetch([w1sh, w2sh, s1sh, s2sh])
rsort(ti, tw, B["sid"], B["swt"], B["seid"], B["nvi"],
B["mbuf"], E, bm)
_quant_direct(hs, B["a1f"], B["sid"], B["nvi"],
B["a1sc5d"], M, K, Ls, 1)
ctg1(B["a1f_v"], w1sh, B["a2"], B["sid"], B["seid"],
B["nvi"], tk, np1, kp1, None, B["a1sc_v"], s1sv,
None, act, bm, 1)
_quant_direct(B["a2_flat"], B["a2f"], B["sid"],
B["nvi"], B["a2sc5d"],
M * tk, de, Ls, tk)
rs2(B["a2f_v"], w1sh, w2sh, B["sid"], B["seid"],
B["nvi"], B["mbuf"], tk, s2_kn, s2sv, B["a2sc_v"],
bm, B["swt"], qt, act, sk2, nt2, None, True)
return B["mbuf"]
return hot
for bm in (32, 64):
s2_kns = [""] + sorted(n for n in _CSV_S2 if _bm_of(n) == bm)
for kn2 in s2_kns:
for sk2 in (1, 2):
for pf in ((False, True) if can_pf else (False,)):
_try(_mk_mix(bm, kn2, sk2, False, pf),
f"mix[bm={bm} s2={(kn2[18:38] or 'def')}"
f"/sk{sk2}{' pf' if pf else ''}]")
if can_pf and best[2] is not None and " pf" not in best[1] \
and "+pf" not in best[1]:
win = best[2]
def hot_pf(hs, w1sh, w2sh, s1sh, s2sh, tw, ti):
_prefetch([w1sh, w2sh, s1sh, s2sh])
return win(hs, w1sh, w2sh, s1sh, s2sh, tw, ti)
_try(hot_pf, best[1] + "+pf")
for k, c in sorted(exc_cnt.items(), key=lambda kv: -kv[1])[:6]:
_L(f" exc×{c}: {k}")
_L(f" DONE {_tm.time()-t0:.1f}s BEST={best[0]:.1f}us {best[1]}")
if best[2] is None:
best[2] = _mk_fm()
return {"hot": best[2], "desc": best[1]}
def custom_kernel(data):
(hs, _, _, _, _, w1sh, w2sh, s1sh, s2sh, tw, ti, config) = data
M = config["bs"]
E = config["n_routed_experts"] + config["n_shared_experts"]
de = config["d_expert"]
skey = (M, E, de)
C = _cfg.get(skey)
if C is None:
C = _build(data, config)
_cfg[skey] = C
return C["hot"](hs, w1sh, w2sh, s1sh, s2sh, tw, ti)
scrolls · 460 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