Skip to content
KernelIndex
Search⌘K

submission 674077

Maxwell Cipher · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

moe_v73_copy.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-674077?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
107.3µs
#11 of 782
2026-03-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:5520fc0849b7f0748693d2ca9e221375af5f5c053bcf245e8e1bd3913364f729
license declaredunknown
license concludedunknown
authorsMaxwell Cipher
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","out_dtype":"bf16",

Kernel source

moe_v73_copy.py330 lines
# HIP MoE v73 -- v72 base + CK _4WG128 for E=257 bs=128
#
# v72 proved _4WG128 is dramatically better than _4WG64 for E=257 bs=512:
#   196us → 140us (-29%)!
#
# v73 experiment: Apply same improvement to E=257 bs=128:
#   v72: CKtile path (ksplit=2, block_m=16, no FP4 quant) → 175us
#   v73: CK path (_4WG128 + FlyDSL, ksplit=0, block_m=32, FP4 quant) → ???
#
# Trade-off:
#   + _4WG128 GEMM is much faster (proven for bs=512)
#   + FlyDSL stage2 is faster than CKtile stage2
#   - CK path adds FP4 quantization overhead (~10-20us)
#   - block_m=32 has more padding than block_m=16 (4.5 tokens/expert)
#
# CU utilization:
#   128*9=1152 tokens, 257 experts, ~4.5 tokens/expert
#   block_m=32: 257 blocks, _4WG128 → ~65 WGs → 1 CU round (25% utilized)
#   Still low, but the faster GEMM kernel may overcome this.

from task import input_t, output_t
import torch
import os
import sys
import functools

from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe
from aiter.ops.moe_sorting import moe_sorting_fwd
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
import aiter.fused_moe as _fm

# =========================================================================== #
# FlyDSL tile registration (needed for stage2 kernel names)
# =========================================================================== #
try:
    import aiter.ops.flydsl.moe_kernels as _flydsl
    for tm, tn in [(32,128),(32,256),(16,256),(16,128),(64,128),(64,256)]:
        _flydsl._KERNEL_PARAMS[f"flydsl_moe2_afp4_wfp4_bf16_t{tm}x{tn}x128_atomic"] = {
            "stage":2,"a_dtype":"fp4","b_dtype":"fp4","out_dtype":"bf16",
            "tile_m":tm,"tile_n":tn,"tile_k":128,"mode":"atomic","MPerBlock":tm,
        }
except ImportError:
    pass

# =========================================================================== #
# Per-shape configs (injected into AITER's cfg_2stages)
# =========================================================================== #
def _key(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)

_4WG128 = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_F16 = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"

_CC = {}
# E=33, TP=4 shapes (identical to v65)
_CC[_key(16, 512, 33)] = {
    "block_m": 16, "ksplit": 2,
    "kernelName1": "", "kernelName2": "",
    "run_1stage": False, "use_non_temporal_load": True,
}
_CC[_key(128, 512, 33)] = {
    "block_m": 32, "ksplit": 0,
    "kernelName1": _4WG128, "kernelName2": _F16,
    "run_1stage": False, "use_non_temporal_load": True,
}
_CC[_key(512, 512, 33)] = {
    "block_m": 64, "ksplit": 0,
    "kernelName1": _4WG128, "kernelName2": _F16,
    "run_1stage": False,
}
_CC[_key(512, 2048, 33)] = {
    "block_m": 64, "ksplit": 0,
    "kernelName1": _4WG128, "kernelName2": _F16,
    "run_1stage": False,
}
# E=257, TP=8 shapes
_CC[_key(16, 256, 257)] = {
    "block_m": 16, "ksplit": 2,
    "kernelName1": "", "kernelName2": "",
    "run_1stage": False, "use_non_temporal_load": True,
}
# CHANGED: CKtile (ksplit=2) → CK _4WG128 (ksplit=0) for E=257 bs=128
_CC[_key(128, 256, 257)] = {
    "block_m": 32, "ksplit": 0,
    "kernelName1": _4WG128, "kernelName2": _F16,
    "run_1stage": False, "use_non_temporal_load": True,
}
# v72 improvement: _4WG128 for E=257 bs=512
_CC[_key(512, 256, 257)] = {
    "block_m": 32, "ksplit": 0,
    "kernelName1": _4WG128, "kernelName2": _F16,
    "run_1stage": False, "use_non_temporal_load": True,
}

# =========================================================================== #
# Injection + metadata capture with NT load patching
# =========================================================================== #
_injected = False
_captured_meta = {}
_cu = None


def _inject():
    global _injected
    if _injected:
        return
    _injected = True

    if _fm.cfg_2stages is None:
        _fm.cfg_2stages = {}
    for k, v in _CC.items():
        _fm.cfg_2stages[k] = v

    orig = _fm.get_2stage_cfgs

    @functools.lru_cache(maxsize=2048)
    def _p(*a):
        global _cu
        m = orig(*a)

        if _cu is None:
            try:
                from aiter.jit.utils.chip_info import get_cu_num
                _cu = get_cu_num()
            except Exception:
                _cu = 256

        keys = (_cu, a[0], a[1], a[2], a[3], a[4],
                str(a[10]), str(a[5]), str(a[6]), str(a[7]),
                str(a[8]), a[9], a[11])

        c = _CC.get(keys)
        if c and c.get("use_non_temporal_load"):
            s = m.stage1
            kw = getattr(s, 'keywords', None) or {}
            if hasattr(s, 'func') and 'use_non_temporal_load' in kw:
                nk = dict(kw)
                nk['use_non_temporal_load'] = True
                m = _fm.MOEMetadata(
                    functools.partial(s.func, **nk),
                    m.stage2, m.block_m, m.ksplit,
                    m.run_1stage, m.has_bias, True)

        _captured_meta[(a[0], a[1], a[2], a[3], a[4])] = m
        return m

    _fm.get_2stage_cfgs = _p


# =========================================================================== #
# Per-shape cache for direct dispatch
# =========================================================================== #
_shape_cache = {}


def _init_shape(data, shape_key):
    """First call: run fused_moe to populate metadata, build dispatch cache."""
    hs = data[0]; guw_sh = data[5]; dw_sh = data[6]
    guws_sh = data[7]; dws_sh = data[8]
    tw = data[9]; ti = data[10]; cfg = data[11]

    M = cfg['bs']
    E = cfg['n_routed_experts'] + cfg['n_shared_experts']
    topk = ti.shape[1]
    hp = cfg['d_hidden_pad'] - cfg['d_hidden']
    ip = cfg['d_expert_pad'] - cfg['d_expert']
    device = hs.device

    result = fused_moe(
        hs, guw_sh, dw_sh, tw, ti,
        expert_mask=None, activation=ActivationType.Silu,
        quant_type=QuantType.per_1x32, doweight_stage1=False,
        w1_scale=guws_sh, w2_scale=dws_sh,
        a1_scale=None, a2_scale=None,
        hidden_pad=hp, intermediate_pad=ip,
    )

    _, model_dim, inter_dim = _fm.get_inter_dim(guw_sh.shape, dw_sh.shape)
    padded_M = _fm.get_padded_M(M)
    meta_key = (padded_M, model_dim, inter_dim, E, topk)

    metadata = _captured_meta.get(meta_key)
    if metadata is None:
        print(f"[V73] No metadata for {shape_key}, using fallback",
              file=sys.stderr, flush=True)
        _shape_cache[shape_key] = None
        return result

    block_m = int(metadata.block_m)
    ksplit = int(metadata.ksplit)
    max_tok = M * topk + E * block_m - topk
    max_blk = (max_tok + block_m - 1) // block_m

    _shape_cache[shape_key] = {
        'meta': metadata,
        'block_m': block_m,
        'ksplit': ksplit,
        'topk': topk,
        'E': E,
        'M': M,
        'model_dim': model_dim,
        'inter_dim': inter_dim,
        'hp': hp,
        'ip': ip,
        'sid': torch.empty(max_tok, dtype=torch.int32, device=device),
        'sw': torch.empty(max_tok, dtype=torch.float32, device=device),
        'seid': torch.empty(max_blk, dtype=torch.int32, device=device),
        'nvi': torch.empty(2, dtype=torch.int32, device=device),
    }

    s1 = metadata.stage1
    s1k = getattr(s1, 'keywords', {}) if hasattr(s1, 'func') else {}
    nt_val = s1k.get('use_non_temporal_load', 'N/A')
    s1name = s1.func.__name__ if hasattr(s1, 'func') else str(s1)[:40]
    print(f"[V73] Init {shape_key}: block_m={block_m} ksplit={ksplit} "
          f"inter={inter_dim} model={model_dim} nt={nt_val} stage1={s1name}",
          file=sys.stderr, flush=True)

    return result


def _direct_dispatch(data, c):
    """Direct dispatch: handles both CK and CKtile paths."""
    hs = data[0]; guw_sh = data[5]; dw_sh = data[6]
    guws_sh = data[7]; dws_sh = data[8]
    tw = data[9]; ti = data[10]

    meta = c['meta']
    block_m = c['block_m']
    ksplit = c['ksplit']
    topk = c['topk']
    E = c['E']
    M = c['M']
    model_dim = c['model_dim']
    inter_dim = c['inter_dim']

    sid = c['sid']; sw = c['sw']; seid = c['seid']; nvi = c['nvi']
    device = hs.device

    # 1. Sorting
    moe_buf = torch.empty(M, model_dim, dtype=torch.bfloat16, device=device)
    moe_sorting_fwd(ti, tw, sid, sw, seid, nvi, moe_buf,
                    E, block_m, None, None, 0)

    if ksplit > 0:
        # CKtile path: bf16 activations, no fp4 quantization
        a1 = hs.to(torch.bfloat16)
        a2_placeholder = torch.empty(M, topk, inter_dim,
                                     dtype=torch.bfloat16, device=device)
        a2 = meta.stage1(
            a1, guw_sh, dw_sh, sid, seid, nvi, a2_placeholder, topk,
            block_m=block_m, a1_scale=None,
            w1_scale=guws_sh.view(dtypes.fp8_e8m0),
            sorted_weights=None)
        meta.stage2(
            a2, guw_sh, dw_sh, sid, seid, nvi, moe_buf, topk,
            w2_scale=dws_sh.view(dtypes.fp8_e8m0),
            a2_scale=None, block_m=block_m, sorted_weights=sw)
    else:
        # CK path: fp4 quantized activations
        a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
            hs, sorted_ids=sid, num_valid_ids=nvi,
            token_num=M, topk=1, block_size=block_m)
        a2 = torch.empty(M, topk, inter_dim, dtype=torch.bfloat16, device=device)
        a2 = meta.stage1(
            a1, guw_sh, dw_sh, sid, seid, nvi, a2, topk,
            block_m=block_m, a1_scale=a1_scale,
            w1_scale=guws_sh.view(dtypes.fp8_e8m0),
            sorted_weights=None)
        a2_flat = a2.view(-1, inter_dim)
        a2_q, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
            a2_flat, sorted_ids=sid, num_valid_ids=nvi,
            token_num=M, topk=topk, block_size=block_m)
        a2 = a2_q.view(M, topk, -1)
        meta.stage2(
            a2, guw_sh, dw_sh, sid, seid, nvi, moe_buf, topk,
            w2_scale=dws_sh.view(dtypes.fp8_e8m0),
            a2_scale=a2_scale, block_m=block_m, sorted_weights=sw)

    return moe_buf


def _fallback(data):
    """Safe fallback: full fused_moe dispatch."""
    hs = data[0]; guw_sh = data[5]; dw_sh = data[6]
    guws_sh = data[7]; dws_sh = data[8]
    tw = data[9]; ti = data[10]; cfg = data[11]
    hp = cfg['d_hidden_pad'] - cfg['d_hidden']
    ip = cfg['d_expert_pad'] - cfg['d_expert']
    return fused_moe(
        hs, guw_sh, dw_sh, tw, ti,
        expert_mask=None, activation=ActivationType.Silu,
        quant_type=QuantType.per_1x32, doweight_stage1=False,
        w1_scale=guws_sh, w2_scale=dws_sh,
        a1_scale=None, a2_scale=None,
        hidden_pad=hp, intermediate_pad=ip,
    )


# =========================================================================== #
# Main entry point
# =========================================================================== #
def custom_kernel(data: input_t) -> output_t:
    _inject()

    cfg = data[11]
    shape_key = (cfg['bs'], cfg['d_expert'],
                 cfg['n_routed_experts'] + cfg['n_shared_experts'])

    if shape_key not in _shape_cache:
        return _init_shape(data, shape_key)

    c = _shape_cache[shape_key]

    if c is None:
        return _fallback(data)

    try:
        return _direct_dispatch(data, c)
    except Exception as e:
        print(f"[V73] Dispatch err {shape_key}: {str(e)[:200]}",
              file=sys.stderr, flush=True)
        import traceback
        traceback.print_exc(file=sys.stderr)
        _shape_cache[shape_key] = None
        return _fallback(data)
scrolls · 330 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 673796.

- # HIP MoE v72 -- v65 base + _4WG128 for E=257 bs=512
+ # HIP MoE v73 -- v72 base + CK _4WG128 for E=257 bs=128
#
- # v71 proved CSV configs are MUCH WORSE for E=257 shapes:
- # bs=16: 90.6us (ours) vs 140us (CSV) — CSV 54% slower
- # bs=128: 175us vs 219us — CSV 25% slower
- # bs=512: 196us vs 252us — CSV 29% slower
- # CSV uses 1WG CK kernels (64x32) which are slower than our CKtile/4WG+FlyDSL.
- # NEVER load CSV for E=257.
+ # v72 proved _4WG128 is dramatically better than _4WG64 for E=257 bs=512:
+ # 196us → 140us (-29%)!
#
- # v72 experiment: E=257 bs=512 stage1 _4WG64 → _4WG128
- # - _4WG64: MPerBlock=64, processes 2 sorting blocks (64/32) per WG
- # - _4WG128: MPerBlock=128, processes 4 sorting blocks (128/32) per WG
- # - Both fit in 1 CU round for 257 experts, but _4WG128 does more work/WG
- # - E=33 shapes already use _4WG128 successfully
+ # v73 experiment: Apply same improvement to E=257 bs=128:
+ # v72: CKtile path (ksplit=2, block_m=16, no FP4 quant) → 175us
+ # v73: CK path (_4WG128 + FlyDSL, ksplit=0, block_m=32, FP4 quant) → ???
#
- # All other shapes IDENTICAL to v65.
+ # Trade-off:
+ # + _4WG128 GEMM is much faster (proven for bs=512)
+ # + FlyDSL stage2 is faster than CKtile stage2
+ # - CK path adds FP4 quantization overhead (~10-20us)
+ # - block_m=32 has more padding than block_m=16 (4.5 tokens/expert)
+ #
+ # CU utilization:
+ # 128*9=1152 tokens, 257 experts, ~4.5 tokens/expert
+ # block_m=32: 257 blocks, _4WG128 → ~65 WGs → 1 CU round (25% utilized)
+ # Still low, but the faster GEMM kernel may overcome this.
from task import input_t, output_t
import torch
⋯ 29 unchanged lines
"QuantType.per_1x32", True, False)
_4WG128 = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
- _4WG64 = "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_F16 = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"
_CC = {}
⋯ 24 unchanged lines
"kernelName1": "", "kernelName2": "",
"run_1stage": False, "use_non_temporal_load": True,
}
+ # CHANGED: CKtile (ksplit=2) → CK _4WG128 (ksplit=0) for E=257 bs=128
_CC[_key(128, 256, 257)] = {
- "block_m": 16, "ksplit": 2,
- "kernelName1": "", "kernelName2": "",
+ "block_m": 32, "ksplit": 0,
+ "kernelName1": _4WG128, "kernelName2": _F16,
"run_1stage": False, "use_non_temporal_load": True,
}
- # CHANGED: _4WG64 → _4WG128 for E=257 bs=512
+ # v72 improvement: _4WG128 for E=257 bs=512
_CC[_key(512, 256, 257)] = {
"block_m": 32, "ksplit": 0,
"kernelName1": _4WG128, "kernelName2": _F16,
⋯ 89 unchanged lines
metadata = _captured_meta.get(meta_key)
if metadata is None:
- print(f"[V72] No metadata for {shape_key}, using fallback",
+ print(f"[V73] No metadata for {shape_key}, using fallback",
file=sys.stderr, flush=True)
_shape_cache[shape_key] = None
return result
⋯ 24 unchanged lines
s1k = getattr(s1, 'keywords', {}) if hasattr(s1, 'func') else {}
nt_val = s1k.get('use_non_temporal_load', 'N/A')
s1name = s1.func.__name__ if hasattr(s1, 'func') else str(s1)[:40]
- print(f"[V72] Init {shape_key}: block_m={block_m} ksplit={ksplit} "
+ print(f"[V73] Init {shape_key}: block_m={block_m} ksplit={ksplit} "
f"inter={inter_dim} model={model_dim} nt={nt_val} stage1={s1name}",
file=sys.stderr, flush=True)
⋯ 99 unchanged lines
try:
return _direct_dispatch(data, c)
except Exception as e:
- print(f"[V72] Dispatch err {shape_key}: {str(e)[:200]}",
+ print(f"[V73] Dispatch err {shape_key}: {str(e)[:200]}",
file=sys.stderr, flush=True)
import traceback
traceback.print_exc(file=sys.stderr)
scrolls · 87 diff lines total

Best evidence level for this revision: reported

JSON