Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 MoEsuite of 7 cases
AMD Instinct MI355X
118.2µs
#35 of 782
2026-03-19

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 lines
def 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 lines
return 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 caching
try:
_orig_quant = _fmoe.fused_dynamic_mxfp4_quant_moe_sort
_quant_cache = {}
⋯ 13 unchanged lines
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
- 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 lines
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=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 lines
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 (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