Skip to content
KernelIndex
Search⌘K

submission 733750

bill_97933 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 145 lines, June 9 Researcher Reciprocity License v1.0.

e180_s5_flydsl.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-733750?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
142.7µs
#121 of 782
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:fb9f9b2b7b8b0592dd1f2c5df2e29a7da424ee05afa64e3ee55444747d2f5154
license declaredunknown
license concludedunknown
authorsbill_97933
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

split-kquant_type=qt, dtype=dt, splitk=0, use_non_temporal_load=u),

Kernel source

e180_s5_flydsl.py145 lines
"""
e180: e178 + test FlyDSL t64x256_reduce for S5 stage2.

S5 (bs=128, t=128, topk=9, e=33, id=512, K=512):
  K=512, tile_k=256 → 2 K-tiles → safe for reduce mode (no 1-K-tile bug)
  Same stage2 kernel as S6 (t64x256_reduce), which is stable in leaderboard.

Change from e178:
  - cks: S4 only gets ksplit=2 (t<=16); S5 (t=128) now gets ksplit=0
  - cbm: S5 uses bm=32 (force, same as S6) for fewer stage2 M-tiles
  - cc: S5 uses FlyDSL t64x256_reduce (same as S6)

S5 stage2 with bm=32 and t64x256_reduce:
  mn = 128*9 + 33*32 - 9 = 2199
  M-tiles = ceil(2199/64) = 35, N-tiles = ceil(7168/256) = 28
  CTAs = 980, waves = ceil(980/256) = 4 → only 4 stage2 waves!

Currently S5=108µs (benchmark) / 113µs (leaderboard).
Target: 90-100µs if 4-wave stage2 is fast enough.
"""
import os
os.environ["SELF_BENCH"] = "1"
os.environ["AITER_USE_NT"] = "1"

import functools, sys, math as _m, torch
import aiter, aiter.fused_moe as fm
from aiter.fused_moe import MOEMetadata, _flydsl_stage2_wrapper, fused_moe
from aiter import ActivationType, QuantType
from task import input_t, output_t

KN1_256x32 = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
FLYDSL_64x256_R = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
FLYDSL_32x128_A = "flydsl_moe2_afp4_wfp4_bf16_t32x128x256_atomic"

@functools.lru_cache(maxsize=2048)
def cks(t, tk, e, id, md):
    if e >= 257 and t <= 128: return 2  # S1/S2: cktile
    if e <= 33 and id <= 512 and t <= 16: return 2  # S4 only: cktile
    # S5 (t=128, e<=33, id=512): now ksplit=0 → CK2stages+FlyDSL
    # S3, S6, S7: ksplit=0
    return 0

@functools.lru_cache(maxsize=2048)
def cbm(t, tk, e, id):
    # S5 and S6: force bm=32 for fewer stage2 M-tiles
    if e <= 33 and id == 512 and t > 16: return 32
    cu = fm.get_cu_num(); tN = 128; tgN = (id + tN - 1) // tN
    sl = [32, 64, 128]; tmp = []
    for el in sl:
        mn = t * tk + e * el - tk; tg = tgN * ((mn + el - 1) // el)
        r = (tg + cu - 1) // cu; em = cu - tg % cu; tmp.append((r, em, el))
    return sorted(tmp, key=lambda x: x[:2])[0][-1]

fm.get_ksplit = cks
fm.get_block_size_M = cbm
_o = fm.get_2stage_cfgs.__wrapped__

@functools.lru_cache(maxsize=2048)
def cc(t, md, id, e, tk, dt, qda, qdw, qt, ug, act, dw, hp, ip, ish=True):
    if e >= 257 and t <= 128:
        os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1"
        r = _o(t, md, id, e, tk, dt, qda, qdw, qt, ug, act, dw, hp, ip, ish)
        os.environ.pop("AITER_BYPASS_TUNE_CONFIG", None); return r
    # S5 and S6 (e<=33, id=512, t>=32): FlyDSL t64x256_reduce (K=512, 2 K-tiles, safe)
    if e <= 33 and id == 512 and t > 16:
        u = fm.use_nt(t, tk, e)
        return MOEMetadata(
            functools.partial(fm.ck_moe_stage1, kernelName=KN1_256x32, activation=act,
                              quant_type=qt, dtype=dt, splitk=0, use_non_temporal_load=u),
            functools.partial(_flydsl_stage2_wrapper, kernelName=FLYDSL_64x256_R),
            32, 0, False)
    # S7: FlyDSL t32x128_atomic (stable, same as e178)
    if t >= 512 and e <= 33 and id >= 2048:
        u = fm.use_nt(t, tk, e)
        return MOEMetadata(
            functools.partial(fm.ck_moe_stage1, kernelName="", activation=act,
                              quant_type=qt, dtype=dt, splitk=0, use_non_temporal_load=u),
            functools.partial(_flydsl_stage2_wrapper, kernelName=FLYDSL_32x128_A),
            128, 0, False)
    # S3: falls to _o → CSV KN1_64x32+KN2_64x32 (stable ~245µs)
    return _o(t, md, id, e, tk, dt, qda, qdw, qt, ug, act, dw, hp, ip, ish)

fm.get_2stage_cfgs = cc

def _impl(data):
    (hs, guw, dw, guws, dws, gush, dsh, gussh, dssh, tw, ti, cfg) = data
    hp = cfg["d_hidden_pad"] - cfg["d_hidden"]
    ip = cfg["d_expert_pad"] - cfg["d_expert"]
    return fused_moe(hs, gush, dsh, tw, ti, expert_mask=None,
                     activation=ActivationType.Silu, quant_type=QuantType.per_1x32,
                     doweight_stage1=False, w1_scale=gussh, w2_scale=dssh,
                     a1_scale=None, a2_scale=None, hidden_pad=hp, intermediate_pad=ip)

_SB = os.environ.get("SELF_BENCH", "0") == "1"
_done = False
_BS = [
    {"bs": 16,  "dexpert": 256,  "dhidden": 7168, "nexpertspertoken": 8, "nroutedexperts": 256, "nsharedexperts": 1, "seed": 42},
    {"bs": 128, "dexpert": 256,  "dhidden": 7168, "nexpertspertoken": 8, "nroutedexperts": 256, "nsharedexperts": 1, "seed": 42},
    {"bs": 512, "dexpert": 256,  "dhidden": 7168, "nexpertspertoken": 8, "nroutedexperts": 256, "nsharedexperts": 1, "seed": 42},
    {"bs": 16,  "dexpert": 512,  "dhidden": 7168, "nexpertspertoken": 8, "nroutedexperts": 32,  "nsharedexperts": 1, "seed": 42},
    {"bs": 128, "dexpert": 512,  "dhidden": 7168, "nexpertspertoken": 8, "nroutedexperts": 32,  "nsharedexperts": 1, "seed": 42},
    {"bs": 512, "dexpert": 512,  "dhidden": 7168, "nexpertspertoken": 8, "nroutedexperts": 32,  "nsharedexperts": 1, "seed": 42},
    {"bs": 512, "dexpert": 2048, "dhidden": 7168, "nexpertspertoken": 8, "nroutedexperts": 32,  "nsharedexperts": 1, "seed": 42},
]

def _cl():
    d = torch.randn((16000, 1024, 1024), device="cuda"); del d

def _rsb():
    global _done
    if _done: return
    _done = True
    try:
        from reference import generate_input
        print("\n[SB] e180: S5+S6 FlyDSL t64x256_reduce, bm=32 (S5: 4 waves, ksplit=0)", file=sys.stderr)
        sd = [(s, generate_input(**s)) for s in _BS]
        for _, d in sd: _impl(d); _impl(d)
        gl = 0.0
        print("[SB] " + "=" * 60, file=sys.stderr)
        for s, d in sd:
            k = f"bs{s['bs']}_E{s['nroutedexperts']+s['nsharedexperts']}_d{s['dexpert']}"
            ts = []
            for i in range(30):
                torch.cuda.synchronize(); _cl()
                se = torch.cuda.Event(enable_timing=True)
                ee = torch.cuda.Event(enable_timing=True)
                se.record(); _impl(d); ee.record(); torch.cuda.synchronize()
                ts.append(se.elapsed_time(ee) * 1e3)
                if len(ts) >= 10:
                    m = sum(ts) / len(ts)
                    v = sum((t - m) ** 2 for t in ts) / (len(ts) - 1)
                    if m > 0 and _m.sqrt(v / len(ts)) / m < 0.001: break
            m = sum(ts) / len(ts); gl += _m.log(m)
            print(f"[SB] {k:>24}: mean={m:7.2f}us best={min(ts):7.2f} ({len(ts)} runs)", file=sys.stderr)
        print(f"[SB] {'GEOMEAN':>24} = {_m.exp(gl / len(_BS)):.2f} us", file=sys.stderr)
        print("[SB] " + "=" * 60, file=sys.stderr)
    except Exception as e:
        import traceback
        print(f"[SB] ERR: {e}", file=sys.stderr)
        traceback.print_exc(file=sys.stderr)

def custom_kernel(data: input_t) -> output_t:
    if _SB and not _done: _rsb()
    return _impl(data)
scrolls · 145 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