Skip to content
KernelIndex
Search⌘K

submission 586370

Ananda Sai A · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-586370?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
124.7µs
#58 of 782
2026-03-18

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:af21695029fa4ce54eed40076e1021ef5d3f3cc45332437377a43dec49a886dc
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.

fp4"stage": 2, "a_dtype": "fp4", "b_dtype": "fp4",
split-ksplitk=0, use_non_temporal_load=nt, dtype=torch.bfloat16)

Kernel source

submission_v3.py215 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)

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 · 215 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 582964.

#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
-
- """
- v160: block_m=64 for bs=512/E=33 (d=512 and d=2048), was block_m=128.
- With ~15.5 tokens/expert, block_m=64 gives better CU distribution.
- """
import os
- import functools
- import torch
- from typing import Dict
+ os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"
+ import functools, torch
from task import input_t, output_t
-
- from aiter import ActivationType, QuantType
+ from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe
- import aiter.fused_moe as _fused_moe_module
import aiter
- import aiter.ops.flydsl.moe_kernels as _flydsl_moe_kernels
+ 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 tile_k=128 kernels that aren't in server's default registration
- _flydsl_moe_kernels._KERNEL_PARAMS["flydsl_moe2_afp4_wfp4_bf16_t32x128x128_atomic"] = {
- "stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16",
- "tile_m": 32, "tile_n": 128, "tile_k": 128, "mode": "atomic", "MPerBlock": 32,
- }
- _flydsl_moe_kernels._KERNEL_PARAMS["flydsl_moe2_afp4_wfp4_bf16_t32x256x128_atomic"] = {
- "stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16",
- "tile_m": 32, "tile_n": 256, "tile_k": 128, "mode": "atomic", "MPerBlock": 32,
- }
- _flydsl_moe_kernels._KERNEL_PARAMS["flydsl_moe2_afp4_wfp4_bf16_t16x256x128_atomic"] = {
- "stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16",
- "tile_m": 16, "tile_n": 256, "tile_k": 128, "mode": "atomic", "MPerBlock": 16,
- }
- _flydsl_moe_kernels._KERNEL_PARAMS["flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"] = {
- "stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16",
- "tile_m": 16, "tile_n": 128, "tile_k": 128, "mode": "atomic", "MPerBlock": 16,
- }
+ # 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,
+ }
- # Inject ksplit=2 configs for shapes that benefit from cktile_moe path
- _CUSTOM_CONFIGS = {}
+ # 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"
- def _make_key(token, inter_dim, expert):
- return (
- 256, token, 7168, inter_dim, expert, 9,
- "ActivationType.Silu", "torch.bfloat16",
- "torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",
- "QuantType.per_1x32", True, False,
- )
+ # 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
- # === E=33 shapes (from v018, proven) ===
- # bs=16/E=33/d=512: cktile_moe gives 59.6us vs 88.7us baseline (-32.8%)
- _CUSTOM_CONFIGS[_make_key(16, 512, 33)] = {
- "block_m": 32,
- "ksplit": 2,
- "kernelName1": "",
- "kernelName2": "",
- "run_1stage": False,
- }
+ # 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()))
- # bs=128/E=33/d=512: 4-WG M128 stage1 + FlyDSL stage2 (v150)
- # v159: block_m=64 to reduce padding waste with ~3.9 tokens/expert
- _CUSTOM_CONFIGS[_make_key(128, 512, 33)] = {
- "block_m": 64,
- "ksplit": 0,
- "kernelName1": "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
- "kernelName2": "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic",
- "run_1stage": False,
- }
+ # 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
- # === E=257 shapes (NEW in v020) ===
- # bs=16/E=257/d=256: try cktile_moe ksplit=2 (overrides tuned CSV config)
- # With 144 token-expert pairs across 257 experts, most experts get 0-1 tokens.
- # Skipping activation quantization + using split-K may help.
- _CUSTOM_CONFIGS[_make_key(16, 256, 257)] = {
- "block_m": 16,
- "ksplit": 2,
- "kernelName1": "",
- "kernelName2": "",
- "run_1stage": False,
- }
+ # Warmup + config injection (first call per shape uses fused_moe)
+ _done = False
+ _warmed = set()
- # bs=128/E=257/d=256: try cktile_moe ksplit=2 (overrides tuned CSV config)
- # bs=128 has ~4.5 tokens/expert avg, similar to E=33 where ksplit=2 helped (-12.9%)
- _CUSTOM_CONFIGS[_make_key(128, 256, 257)] = {
- "block_m": 16,
- "ksplit": 2,
- "kernelName1": "",
- "kernelName2": "",
- "run_1stage": False,
- }
-
- # === bs=512/E=33 shapes: inject 4-WG stage1 kernel ===
- # The 256x64x128x128_1x4 kernel uses 4 workgroups per CU for better utilization.
- # v037 showed d=2048: -3.2% (349->338µs). Now also try d=512 with same 4-WG kernel.
- _4WG_STAGE1 = "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
- _4WG_STAGE1_M128 = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
-
- _FLYDSL_STAGE2 = "flydsl_moe2_afp4_wfp4_bf16_t32x128x256_atomic"
- _FLYDSL_STAGE2_K128 = "flydsl_moe2_afp4_wfp4_bf16_t32x128x128_atomic"
- _FLYDSL_STAGE2_N256_K128 = "flydsl_moe2_afp4_wfp4_bf16_t32x256x128_atomic"
- _FLYDSL_STAGE2_M16_N256_K128 = "flydsl_moe2_afp4_wfp4_bf16_t16x256x128_atomic"
- _FLYDSL_STAGE2_M16_N128_K128 = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"
-
- _CUSTOM_CONFIGS[_make_key(512, 2048, 33)] = {
- "block_m": 64, # v160: was 128, try 64 for better CU distribution
- "ksplit": 0,
- "kernelName1": _4WG_STAGE1_M128,
- "kernelName2": _FLYDSL_STAGE2_M16_N128_K128, # v138: t16x128x128 for d=2048
- "run_1stage": False,
- }
-
- # === bs=512/E=33/d=512: 4-WG M128 stage1 + FlyDSL stage2 ===
- # v160: block_m=64 (was 128) for better CU distribution
- _CUSTOM_CONFIGS[_make_key(512, 512, 33)] = {
- "block_m": 64, # v160: was 128, try 64
- "ksplit": 0,
- "kernelName1": _4WG_STAGE1_M128,
- "kernelName2": _FLYDSL_STAGE2_M16_N128_K128, # v138: t16x128x128 for d=512
- "run_1stage": False,
- }
-
- # === bs=512/E=257: 4-WG CK stage1 + FlyDSL stage2 ===
- # v144: 4-WG (256x32x128x128_1x4) stage1 + FlyDSL stage2.
- # DSV3 tuned CSV uses 4-WG for token>=64/E=257. Block_m=32 matches CSV.
- _4WG_STAGE1_M32 = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
- _CUSTOM_CONFIGS[_make_key(512, 256, 257)] = {
- "block_m": 32,
- "ksplit": 0,
- "kernelName1": _4WG_STAGE1_M32, # v144: 4-WG instead of 1-WG
- "kernelName2": _FLYDSL_STAGE2_M16_N128_K128, # v143: FlyDSL stage2
- "run_1stage": False,
- "use_non_temporal_load": True,
- }
-
- _injected = False
-
- def _inject_configs():
- global _injected
- if _injected:
- return
- _injected = True
-
- if _fused_moe_module.cfg_2stages is None:
+ 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
- tune_file = AITER_CONFIGS.AITER_CONFIG_FMOE_FILE
- if os.path.exists(tune_file):
- _INDEX_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(tune_file)
- if "_tag" in df.columns:
- df = df[df["_tag"].fillna("") == ""]
- _fused_moe_module.cfg_2stages = df.set_index(_INDEX_COLS).to_dict("index")
+ 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:
- _fused_moe_module.cfg_2stages = {}
+ _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)
- _fused_moe_module.cfg_2stages.update(_CUSTOM_CONFIGS)
+ 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)
- # Monkeypatch get_2stage_cfgs to support use_non_temporal_load from config
- _original_get_2stage_cfgs = _fused_moe_module.get_2stage_cfgs
+ # 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)
- @functools.lru_cache(maxsize=2048)
- def _patched_get_2stage_cfgs(
- token, model_dim, inter_dim, expert, topk,
- dtype, q_dtype_a, q_dtype_w, q_type, use_g1u1,
- activation, doweight_stage1, hidden_pad, intermediate_pad, is_shuffled=True,
- ):
- # Get the original metadata
- metadata = _original_get_2stage_cfgs(
- token, model_dim, inter_dim, expert, topk,
- dtype, q_dtype_a, q_dtype_w, q_type, use_g1u1,
- activation, doweight_stage1, hidden_pad, intermediate_pad, is_shuffled,
- )
+ w1s_e8 = w1s.view(dtypes.fp8_e8m0)
+ w2s_e8 = w2s.view(dtypes.fp8_e8m0)
- # Check if this shape has a custom NT setting
- from aiter.jit.utils.chip_info import get_cu_num
- cu_num = get_cu_num()
- keys = (
- cu_num, token, model_dim, inter_dim, expert, topk,
- str(activation), str(dtype), str(q_dtype_a), str(q_dtype_w),
- str(q_type), use_g1u1, doweight_stage1,
- )
- cfg = _fused_moe_module.cfg_2stages.get(keys)
- if cfg and cfg.get("use_non_temporal_load") is not None:
- nt = cfg["use_non_temporal_load"]
- # Rebuild stage1 partial with NT override
- old_s1 = metadata.stage1
- if hasattr(old_s1, 'func') and old_s1.func is not None:
- if 'use_non_temporal_load' in (old_s1.keywords or {}):
- new_kw = dict(old_s1.keywords)
- new_kw['use_non_temporal_load'] = nt
- metadata = _fused_moe_module.MOEMetadata(
- functools.partial(old_s1.func, **{k: v for k, v in new_kw.items()}),
- metadata.stage2,
- metadata.block_m,
- metadata.ksplit,
- metadata.run_1stage,
- metadata.has_bias,
- nt,
- )
- # Also patch stage2 if it's a CK kernel (not FlyDSL)
- old_s2 = metadata.stage2
- if old_s2 and hasattr(old_s2, 'keywords') and 'use_non_temporal_load' in (old_s2.keywords or {}):
- new_kw2 = dict(old_s2.keywords)
- new_kw2['use_non_temporal_load'] = nt
- metadata = _fused_moe_module.MOEMetadata(
- metadata.stage1,
- functools.partial(old_s2.func, **{k: v for k, v in new_kw2.items()}),
- metadata.block_m,
- metadata.ksplit,
- metadata.run_1stage,
- metadata.has_bias,
- nt,
- )
- return metadata
+ # 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
- _fused_moe_module.get_2stage_cfgs = _patched_get_2stage_cfgs
+ 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
- def custom_kernel(data: input_t) -> output_t:
- (
- hidden_states, gate_up_weight, down_weight,
- gate_up_weight_scale, down_weight_scale,
- gate_up_weight_shuffled, down_weight_shuffled,
- gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
- topk_weights, topk_ids, config,
- ) = data
+ 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
- _inject_configs()
-
- hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
- intermediate_pad = config["d_expert_pad"] - config["d_expert"]
-
- output = fused_moe(
- hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
- topk_weights, topk_ids,
- expert_mask=None, activation=ActivationType.Silu,
- quant_type=QuantType.per_1x32, doweight_stage1=False,
- w1_scale=gate_up_weight_scale_shuffled,
- w2_scale=down_weight_scale_shuffled,
- a1_scale=None, a2_scale=None,
- hidden_pad=hidden_pad, intermediate_pad=intermediate_pad,
- )
-
- return output
+ 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 · 437 diff lines total

Best evidence level for this revision: reported

JSON