submission 602624
Ananda Sai A · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 216 lines, June 9 Researcher Reciprocity License v1.0.
submission_v4c_nograd.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-602624?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:ebf77fa907783d5f849a3e8173f98a3b18c942c73cce306223bfca4284821674
license declaredunknown
license concludedunknown
authorsAnanda Sai A
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
submission_v4c_nograd.py216 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
import os
os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"
import functools, torch
from task import input_t, output_t
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe
import aiter
import aiter.fused_moe as _fm
import aiter.ops.flydsl.moe_kernels as _flydsl
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
from aiter.ops.flydsl.moe_kernels import get_flydsl_kernel_params, _get_compiled_stage2
# Register FlyDSL stage2 t16
for _tm in (16, 32):
for _tn in (128, 256):
for _tk in (128, 256):
_n = f"flydsl_moe2_afp4_wfp4_bf16_t{_tm}x{_tn}x{_tk}_atomic"
if _n not in _flydsl._KERNEL_PARAMS:
_flydsl._KERNEL_PARAMS[_n] = {
"stage": 2, "a_dtype": "fp4", "b_dtype": "fp4",
"out_dtype": "bf16", "tile_m": _tm, "tile_n": _tn,
"tile_k": _tk, "mode": "atomic", "MPerBlock": _tm,
}
# Kernel names
_M128 = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_M32 = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_F2 = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"
# Sort workspace (cached + double-buffered moe_buf)
_W = {}
def _do_sort(ti, tw, E, model_dim, block_m):
d = ti.device
M, topk = ti.shape
p = int(ti.numel() + E * block_m - topk)
b = (p + block_m - 1) // block_m
k = (p, b, M, model_dim, str(d))
w = _W.get(k)
if w is None:
w = [torch.empty(p, dtype=dtypes.i32, device=d),
torch.empty(p, dtype=dtypes.fp32, device=d),
torch.empty(b, dtype=dtypes.i32, device=d),
torch.empty(2, dtype=dtypes.i32, device=d),
torch.empty((M, model_dim), dtype=torch.bfloat16, device=d),
torch.empty((M, model_dim), dtype=torch.bfloat16, device=d),
0]
_W[k] = w
o = w[4 + w[6]]
w[6] ^= 1
aiter.moe_sorting_opus_fwd(ti, tw, w[0], w[1], w[2], w[3], o,
E, block_m, None, None, 0)
return w[0], w[1], w[2], w[3], o
# FlyDSL stage2 direct call (bypass wrapper)
def _fly2(a2, w2, sid, seid, nvid, out, topk, a2s, w2s, sw, name):
p = get_flydsl_kernel_params(name)
inter_dim = a2.shape[2]
if p["a_dtype"] == "fp4":
inter_dim = inter_dim * 2
fn = _get_compiled_stage2(
w2.shape[1], inter_dim, w2.shape[0], topk,
p["tile_m"], p["tile_n"], p["tile_k"],
(sw is not None), p["a_dtype"], p["b_dtype"], p["out_dtype"],
(p.get("mode", "atomic") != "reduce"))
if sw is None:
sw = torch.empty(sid.shape, dtype=torch.float32, device=sid.device)
fn(out, a2, w2, a2s, w2s, sid, seid, sw, nvid, a2.shape[0], int(seid.numel()))
# Prebound CK stage1 partials
_s1_cache = {}
def _get_s1(kn, nt):
k = (kn, nt)
s = _s1_cache.get(k)
if s is None:
s = functools.partial(
_fm.ck_moe_stage1, kernelName=kn,
activation=ActivationType.Silu, quant_type=QuantType.per_1x32,
splitk=0, use_non_temporal_load=nt, dtype=torch.bfloat16)
_s1_cache[k] = s
return s
# Warmup + config injection (first call per shape uses fused_moe)
_done = False
_warmed = set()
def _init():
global _done
if _done: return
_done = True
# Patch sorting for warmup (accept all positional + keyword args from moe_sorting)
def _sort_shim(ti, tw, E, md, dt, bs, em=None, nlt=None, dp=0, use_opus=True):
return _do_sort(ti, tw, E, md, bs)
_fm._moe_sorting_impl = _sort_shim
if _fm.cfg_2stages is None:
import pandas as pd
from aiter.jit.core import AITER_CONFIGS
f = AITER_CONFIGS.AITER_CONFIG_FMOE_FILE
if os.path.exists(f):
cols = ["cu_num","token","model_dim","inter_dim","expert","topk",
"act_type","dtype","q_dtype_a","q_dtype_w","q_type","use_g1u1","doweight_stage1"]
df = pd.read_csv(f)
if "_tag" in df.columns: df = df[df["_tag"].fillna("") == ""]
_fm.cfg_2stages = df.set_index(cols).to_dict("index")
else:
_fm.cfg_2stages = {}
def _k(t, i, e):
return (256, t, 7168, i, e, 9, "ActivationType.Silu", "torch.bfloat16",
"torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",
"QuantType.per_1x32", True, False)
_C = {
_k(16,256,257): {"block_m":16,"ksplit":2,"kernelName1":"","kernelName2":"","run_1stage":False},
_k(128,256,257): {"block_m":16,"ksplit":2,"kernelName1":"","kernelName2":"","run_1stage":False},
_k(512,256,257): {"block_m":32,"ksplit":0,"kernelName1":_M32,"kernelName2":_F2,"run_1stage":False,"use_non_temporal_load":True},
_k(16,512,33): {"block_m":32,"ksplit":2,"kernelName1":"","kernelName2":"","run_1stage":False},
_k(128,512,33): {"block_m":64,"ksplit":0,"kernelName1":_M128,"kernelName2":_F2,"run_1stage":False,"use_non_temporal_load":True},
_k(512,512,33): {"block_m":64,"ksplit":0,"kernelName1":_M128,"kernelName2":_F2,"run_1stage":False},
_k(512,2048,33): {"block_m":64,"ksplit":0,"kernelName1":_M128,"kernelName2":_F2,"run_1stage":False},
}
_fm.cfg_2stages.update(_C)
@torch.no_grad()
def custom_kernel(data: input_t) -> output_t:
hs, _, _, _, _, w1, w2, w1s, w2s, tw, ti, cfg = data
_init()
M = hs.shape[0]
E = int(cfg["n_routed_experts"]) + int(cfg["n_shared_experts"])
inter = int(cfg["d_expert"])
h_pad = cfg["d_hidden_pad"] - cfg["d_hidden"]
i_pad = cfg["d_expert_pad"] - cfg["d_expert"]
sk = (M, E, inter)
# First call: warmup via fused_moe (triggers JIT)
if sk not in _warmed:
_warmed.add(sk)
return fused_moe(hs, w1, w2, tw, ti, expert_mask=None,
activation=ActivationType.Silu, quant_type=QuantType.per_1x32,
doweight_stage1=False, w1_scale=w1s, w2_scale=w2s,
a1_scale=None, a2_scale=None, hidden_pad=h_pad, intermediate_pad=i_pad)
w1s_e8 = w1s.view(dtypes.fp8_e8m0)
w2s_e8 = w2s.view(dtypes.fp8_e8m0)
# Shape-specialized fast paths (no fused_moe dispatch)
if E == 257 and M <= 128:
# Shapes 1,2: cktile ksplit=2, block_m=16
bm = 16
sid, sw, seid, nvid, out = _do_sort(ti, tw, E, 7168, bm)
n_pad = i_pad // 64 * 64 * 2
k_pad = h_pad // 128 * 128
_, n1, _ = w1.shape
D = (w2.shape[2]) * 2
tmp = torch.zeros((M, 9, n1), dtype=torch.bfloat16, device=hs.device)
a2 = torch.empty((M, 9, D), dtype=torch.bfloat16, device=hs.device)
aiter.moe_cktile2stages_gemm1(hs, w1, tmp, sid, seid, nvid, 9,
n_pad, k_pad, None, None, w1s_e8, None, ActivationType.Silu, bm, 2)
aiter.silu_and_mul(a2, tmp)
n2 = h_pad // 64 * 64
k2 = i_pad // 128 * 128
aiter.moe_cktile2stages_gemm2(a2, w2, out, sid, seid, nvid, 9,
n2, k2, sw, None, w2s_e8, None, ActivationType.Silu, bm)
return out
elif E == 33 and M == 16:
# Shape 4: use fused_moe (cktile direct path regresses for this shape)
return fused_moe(hs, w1, w2, tw, ti, expert_mask=None,
activation=ActivationType.Silu, quant_type=QuantType.per_1x32,
doweight_stage1=False, w1_scale=w1s, w2_scale=w2s,
a1_scale=None, a2_scale=None, hidden_pad=h_pad, intermediate_pad=i_pad)
elif E == 257 and M == 512:
# Shape 3: CK M32 + FlyDSL, block_m=32, NT=True
bm = 32
sid, sw, seid, nvid, out = _do_sort(ti, tw, E, 7168, bm)
a1, a1s = fused_dynamic_mxfp4_quant_moe_sort(hs, sorted_ids=sid,
num_valid_ids=nvid, token_num=M, topk=1, block_size=bm)
a2 = torch.empty((M, 9, inter), dtype=torch.bfloat16, device=hs.device)
s1 = _get_s1(_M32, True)
a2 = s1(a1, w1, w2, sid, seid, nvid, a2, 9,
block_m=bm, a1_scale=a1s, w1_scale=w1s_e8, sorted_weights=None)
a2q, a2s = fused_dynamic_mxfp4_quant_moe_sort(a2.view(-1, inter),
sorted_ids=sid, num_valid_ids=nvid, token_num=M, topk=9, block_size=bm)
_fly2(a2q.view(M, 9, -1), w2, sid, seid, nvid, out, 9, a2s, w2s_e8, sw, _F2)
return out
elif E == 33 and inter == 512 and M == 128:
# Shape 5: CK M128 + FlyDSL, block_m=64, NT=True
bm = 64
sid, sw, seid, nvid, out = _do_sort(ti, tw, E, 7168, bm)
a1, a1s = fused_dynamic_mxfp4_quant_moe_sort(hs, sorted_ids=sid,
num_valid_ids=nvid, token_num=M, topk=1, block_size=bm)
a2 = torch.empty((M, 9, inter), dtype=torch.bfloat16, device=hs.device)
s1 = _get_s1(_M128, True)
a2 = s1(a1, w1, w2, sid, seid, nvid, a2, 9,
block_m=bm, a1_scale=a1s, w1_scale=w1s_e8, sorted_weights=None)
a2q, a2s = fused_dynamic_mxfp4_quant_moe_sort(a2.view(-1, inter),
sorted_ids=sid, num_valid_ids=nvid, token_num=M, topk=9, block_size=bm)
_fly2(a2q.view(M, 9, -1), w2, sid, seid, nvid, out, 9, a2s, w2s_e8, sw, _F2)
return out
else:
# Shapes 6, 7: CK M128 + FlyDSL, block_m=64
bm = 64
sid, sw, seid, nvid, out = _do_sort(ti, tw, E, 7168, bm)
a1, a1s = fused_dynamic_mxfp4_quant_moe_sort(hs, sorted_ids=sid,
num_valid_ids=nvid, token_num=M, topk=1, block_size=bm)
a2 = torch.empty((M, 9, inter), dtype=torch.bfloat16, device=hs.device)
s1 = _get_s1(_M128, False)
a2 = s1(a1, w1, w2, sid, seid, nvid, a2, 9,
block_m=bm, a1_scale=a1s, w1_scale=w1s_e8, sorted_weights=None)
a2q, a2s = fused_dynamic_mxfp4_quant_moe_sort(a2.view(-1, inter),
sorted_ids=sid, num_valid_ids=nvid, token_num=M, topk=9, block_size=bm)
_fly2(a2q.view(M, 9, -1), w2, sid, seid, nvid, out, 9, a2s, w2s_e8, sw, _F2)
return out
scrolls · 216 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 586370.
⋯ 119 unchanged lines}_fm.cfg_2stages.update(_C)+ @torch.no_grad()def custom_kernel(data: input_t) -> output_t:hs, _, _, _, _, w1, w2, w1s, w2s, tw, ti, cfg = data_init()
Best evidence level for this revision: reported
JSON