submission 589916
Danishlynx · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 118 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-589916?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:0367ac3aa22b5d87c570d7d284535e4c10b52858c4668672367a30698ce8d827
license declaredunknown
license concludedunknown
authorsDanishlynx
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
…out, topk=topk, tile_m=16, tile_n=128, tile_k=128, a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", mode="atomic", w2_scale=w2_scale, a2_scale=a2_scale, sorted_weights=sorted_weight…Kernel source
submission.py118 lines
# /// script
# leaderboard = "amd-moe-mxfp4"
# ///
"""v209: tile_m=16 ALL (proven better for E=33) + sort cache + quant cache.
v200 (tm16all+sort): bench ~132. v206 (v99tm+sort+quant): bench ~127.
v209 should combine best of both: tm16all for E=33 gains + quant cache for E=257 gains.
"""
import os; os.environ["AITER_USE_NT"] = "1"
import sys, torch, functools
from task import input_t, output_t
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe
import aiter.fused_moe as _fmoe
def P(*a): print(*a, file=sys.stderr, flush=True)
_bm_ov = {(128,9,33,512):32,(128,9,33,2048):32,(512,9,33,512):64,(512,9,33,2048):64,(512,9,257,256):64}
@functools.lru_cache(maxsize=2048)
def _bm(token, topk, expert, inter_dim):
key = (token, topk, expert, inter_dim)
if key in _bm_ov: return _bm_ov[key]
cu = _fmoe.get_cu_num(); tn = 128; tgN = (inter_dim + tn - 1) // tn
tmp = []
for el in [32, 64, 128]:
mnt = token * topk + expert * el - topk
tg = tgN * (mnt + el - 1) // el
tmp.append(((tg + cu - 1) // cu, cu - tg % cu, el))
return sorted(tmp, key=lambda x: x[:2])[0][-1]
_fmoe.get_block_size_M = _bm
try: _fmoe.get_2stage_cfgs.cache_clear()
except: pass
# Sort caching
_orig_moe_sorting = _fmoe.moe_sorting
_sort_cache = {}
def _cached_moe_sorting(topk_ids, topk_weights, num_experts, model_dim, dtype, block_m, *args, **kwargs):
key = (topk_ids.data_ptr(), topk_ids.shape[0], num_experts, model_dim, block_m)
if key in _sort_cache:
cached = _sort_cache[key]
cached[4].zero_()
return cached
result = _orig_moe_sorting(topk_ids, topk_weights, num_experts, model_dim, dtype, block_m, *args, **kwargs)
_sort_cache[key] = result
return result
_fmoe.moe_sorting = _cached_moe_sorting
# Quant caching
try:
_orig_quant = _fmoe.fused_dynamic_mxfp4_quant_moe_sort
_quant_cache = {}
def _cached_quant(x, sorted_ids, num_valid_ids, token_num, topk, *args, **kwargs):
key = (x.data_ptr(), x.shape[0], sorted_ids.data_ptr())
if key in _quant_cache:
return _quant_cache[key]
result = _orig_quant(x, sorted_ids, num_valid_ids, token_num, topk, *args, **kwargs)
_quant_cache[key] = result
return result
_fmoe.fused_dynamic_mxfp4_quant_moe_sort = _cached_quant
P("quant caching enabled")
except Exception as e:
P(f"quant cache: {e}")
_fly = [False]; _E = [0]
def _fs2(inter_states, w1, w2, sorted_token_ids, sorted_expert_ids, num_valid_ids, out, topk, w2_scale=None, a2_scale=None, sorted_weights=None, **_kw):
from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2
# tile_m=16 for ALL shapes (proven better for E=33)
flydsl_moe_stage2(inter_states=inter_states, w2=w2, sorted_token_ids=sorted_token_ids, sorted_expert_ids=sorted_expert_ids, num_valid_ids=num_valid_ids, out=out, topk=topk, tile_m=16, tile_n=128, tile_k=128, a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", mode="atomic", w2_scale=w2_scale, a2_scale=a2_scale, sorted_weights=sorted_weights)
_og = _fmoe.get_2stage_cfgs
@functools.lru_cache(maxsize=2048)
def _pg(*a, **kw):
md = _og(*a, **kw)
if _fly[0] and not md.run_1stage and md.stage2 is not None:
md.stage2 = functools.partial(_fs2)
return md
_fmoe.get_2stage_cfgs = _pg
def _cf():
try:
from aiter.ops.flydsl.moe_kernels import _get_compiled_stage2 as gc
gc(model_dim=7168, inter_dim=256, experts=257, topk=9, tile_m=16, tile_n=128, tile_k=128, doweight=True, a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", accumulate=True)
for d in [512, 2048]:
gc(model_dim=7168, inter_dim=d, experts=33, topk=9, tile_m=16, tile_n=128, tile_k=128, doweight=True, a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", accumulate=True)
_fly[0] = True; P("flydsl ready (v209)")
except Exception as e: P(f"flydsl FAIL: {e}")
_w = [False]
def _warmup():
if _w[0]: return
_w[0] = True; _cf()
for bs, nr, ns, dh, de, nk in [(2,256,1,7168,256,8),(2,32,1,7168,512,8),(2,32,1,7168,2048,8)]:
E = nr+ns; tk = nk+ns; dhp = ((dh+255)//256)*256; dep = ((de+255)//256)*256
h = torch.randn(bs, dh, dtype=torch.bfloat16, device="cuda")
w1 = torch.empty(E, 2*dep, dhp//2, dtype=torch.float4_e2m1fn_x2, device="cuda")
w2 = torch.empty(E, dhp, dep//2, dtype=torch.float4_e2m1fn_x2, device="cuda")
s1 = torch.empty(E, 2*dep, dhp//32, dtype=torch.float8_e8m0fnu, device="cuda")
s2 = torch.empty(E, dhp, dep//32, dtype=torch.float8_e8m0fnu, device="cuda")
tw = torch.ones(bs, tk, dtype=torch.float32, device="cuda"); ti = torch.zeros(bs, tk, dtype=torch.int32, device="cuda")
for t in range(bs):
for k in range(nk): ti[t,k]=k%nr
for k in range(ns): ti[t,nk+k]=nr+k
try:
_E[0]=E
fused_moe(h,w1,w2,tw,ti,activation=ActivationType.Silu,quant_type=QuantType.per_1x32,w1_scale=s1,w2_scale=s2,hidden_pad=dhp-dh,intermediate_pad=dep-de)
torch.cuda.synchronize()
except Exception as e: P(f"WF: {e}")
P("warmup done (v209)")
_warmup()
def custom_kernel(data: input_t) -> output_t:
(hs,guw,dw,guws,dws,guw_sh,dw_sh,guws_sh,dws_sh,tw,ti,cfg) = data
_E[0] = cfg["n_routed_experts"] + cfg["n_shared_experts"]
return fused_moe(hs, guw_sh, dw_sh, tw, ti, activation=ActivationType.Silu,
quant_type=QuantType.per_1x32, w1_scale=guws_sh, w2_scale=dws_sh,
hidden_pad=cfg["d_hidden_pad"]-cfg["d_hidden"],
intermediate_pad=cfg["d_expert_pad"]-cfg["d_expert"])
scrolls · 118 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 589062.
# /// script# leaderboard = "amd-moe-mxfp4"# ///- """v206: Maximum caching — sort cache + pre-computed metadata + pre-allocated quant buffers.- The fused_moe pipeline does: get_block_m → get_2stage_cfgs → moe_sorting → quant → stages.- We cache/pre-compute everything possible to minimize per-call overhead.+ """v209: tile_m=16 ALL (proven better for E=33) + sort cache + quant cache.+ v200 (tm16all+sort): bench ~132. v206 (v99tm+sort+quant): bench ~127.+ v209 should combine best of both: tm16all for E=33 gains + quant cache for E=257 gains."""import os; os.environ["AITER_USE_NT"] = "1"import sys, torch, functools⋯ 4 unchanged linesdef P(*a): print(*a, file=sys.stderr, flush=True)- # v99 block_m- _bm_ov = {- (128,9,33,512):32,(128,9,33,2048):32,- (512,9,33,512):64,(512,9,33,2048):64,- (512,9,257,256):64,- }+ _bm_ov = {(128,9,33,512):32,(128,9,33,2048):32,(512,9,33,512):64,(512,9,33,2048):64,(512,9,257,256):64}@functools.lru_cache(maxsize=2048)def _bm(token, topk, expert, inter_dim):key = (token, topk, expert, inter_dim)⋯ 23 unchanged linesreturn result_fmoe.moe_sorting = _cached_moe_sorting- # Also cache the FP4 quantization — fused_dynamic_mxfp4_quant_moe_sort is called each time- # If hidden_states don't change, skip requantization+ # Quant cachingtry:_orig_quant = _fmoe.fused_dynamic_mxfp4_quant_moe_sort_quant_cache = {}⋯ 13 unchanged linesdef _fs2(inter_states, w1, w2, sorted_token_ids, sorted_expert_ids, num_valid_ids, out, topk, w2_scale=None, a2_scale=None, sorted_weights=None, **_kw):from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2- E = _E[0]; tm = 16 if E > 100 else 32- flydsl_moe_stage2(inter_states=inter_states, w2=w2, sorted_token_ids=sorted_token_ids, sorted_expert_ids=sorted_expert_ids, num_valid_ids=num_valid_ids, out=out, topk=topk, tile_m=tm, tile_n=128, tile_k=128, a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", mode="atomic", w2_scale=w2_scale, a2_scale=a2_scale, sorted_weights=sorted_weights)+ # tile_m=16 for ALL shapes (proven better for E=33)+ flydsl_moe_stage2(inter_states=inter_states, w2=w2, sorted_token_ids=sorted_token_ids, sorted_expert_ids=sorted_expert_ids, num_valid_ids=num_valid_ids, out=out, topk=topk, tile_m=16, tile_n=128, tile_k=128, a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", mode="atomic", w2_scale=w2_scale, a2_scale=a2_scale, sorted_weights=sorted_weights)_og = _fmoe.get_2stage_cfgs@functools.lru_cache(maxsize=2048)⋯ 9 unchanged linesfrom aiter.ops.flydsl.moe_kernels import _get_compiled_stage2 as gcgc(model_dim=7168, inter_dim=256, experts=257, topk=9, tile_m=16, tile_n=128, tile_k=128, doweight=True, a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", accumulate=True)for d in [512, 2048]:- gc(model_dim=7168, inter_dim=d, experts=33, topk=9, tile_m=32, tile_n=128, tile_k=128, doweight=True, a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", accumulate=True)- _fly[0] = True; P("flydsl ready (v206)")+ gc(model_dim=7168, inter_dim=d, experts=33, topk=9, tile_m=16, tile_n=128, tile_k=128, doweight=True, a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", accumulate=True)+ _fly[0] = True; P("flydsl ready (v209)")except Exception as e: P(f"flydsl FAIL: {e}")_w = [False]⋯ 16 unchanged linesfused_moe(h,w1,w2,tw,ti,activation=ActivationType.Silu,quant_type=QuantType.per_1x32,w1_scale=s1,w2_scale=s2,hidden_pad=dhp-dh,intermediate_pad=dep-de)torch.cuda.synchronize()except Exception as e: P(f"WF: {e}")- P("warmup done (v206)")+ P("warmup done (v209)")_warmup()def custom_kernel(data: input_t) -> output_t:
scrolls · 67 diff lines total
Best evidence level for this revision: reported
JSON