Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 MoEsuite of 7 cases
AMD Instinct MI355X
169.8µs
#312 of 782
2026-04-03

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.

autotunev16e: L2-cold autotune — raw cktile + raw ck2stages, same 5-kernel
fp4cktile = alternate GEMM backend consuming SAME pre-quantized fp4 as
num-warps = 4_touch_k[(triton.cdiv(n, 8192),)](flat, n, BLK=8192, num_warps=4)
tile-m = 16BLOCK_SIZE_M=16, BLOCK_SIZE_N=4, TOPK=topk,
tile-n = 4BLOCK_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