Skip to content
KernelIndex
Search⌘K

submission 655339

mumu.0567 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v1-4.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-655339?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
178.0µs
#403 of 782
2026-03-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4442d03bca95170673f92d4e0272cb56d9dd76793461cb45f0dadbb9aca32138
license declaredunknown
license concludedunknown
authorsmumu.0567
imported2026-08-26

Techniques

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

split-ksplitk=0,

Kernel source

v1-4.py190 lines
import torch
import functools
import sys

from task import input_t, output_t
from aiter import ActivationType, QuantType
import aiter
import aiter.fused_moe as _fm
from aiter.fused_moe import fused_moe, MOEMetadata
from aiter.ops.flydsl.utils import is_flydsl_available
from aiter.jit.utils.chip_info import get_cu_num as _get_cu_num

_CU_NUM = _get_cu_num()
_FLYDSL_OK = is_flydsl_available()

print(f"[custom] FlyDSL={_FLYDSL_OK}, CU_NUM={_CU_NUM}", file=sys.stderr)

# ── 验证 flydsl 模块可正常使用(版本精确匹配检查)─────────────────────────────
if _FLYDSL_OK:
    try:
        import aiter.ops.flydsl as _flydsl_mod
        assert hasattr(_flydsl_mod, 'flydsl_moe_stage1'), "flydsl_moe_stage1 missing"
        print(f"[custom] flydsl module fully loaded OK", file=sys.stderr)
    except Exception as e:
        print(f"[custom] WARNING: flydsl import failed: {e}", file=sys.stderr)
        _FLYDSL_OK = False

# ── Kernel names ──────────────────────────────────────────────────────────────
K1_SMALL = (
    "moe_ck2stages_gemm1_64x32x32x128_1x1"
    "_MulABScaleShuffled_v3_Nswizzle0_Quant3"
    "_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
K1_MED = (
    "moe_ck2stages_gemm1_256x32x128x128_1x4"
    "_MulABScaleShuffled_v3_Nswizzle0_Quant3"
    "_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
K1_LARGE = (
    "moe_ck2stages_gemm1_256x128x128x128_1x4"
    "_MulABScaleShuffled_v3_Nswizzle0_Quant3"
    "_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
K2_FLY32 = "flydsl_moe2_afp4_wfp4_bf16_t32x256x256_atomic"
K2_FLY64 = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_atomic"

# ── 新增:针对 MI355X 的 stage2 tile_n=128 变体(更适合 d_model=7168)─────────
K2_FLY32_128 = "flydsl_moe2_afp4_wfp4_bf16_t32x128x256_atomic"
K2_FLY64_128 = "flydsl_moe2_afp4_wfp4_bf16_t64x128x256_atomic"

# ── Custom config table ───────────────────────────────────────────────────────
_ACT = "ActivationType.Silu"
_DT  = "torch.bfloat16"
_QA  = "torch.float4_e2m1fn_x2"
_QW  = "torch.float4_e2m1fn_x2"
_QT  = "QuantType.per_1x32"

_CUSTOM_CONFIGS = {}

def _add(token, model_dim, inter_dim, expert, topk, block_m, k1, k2):
    """
    注意:key 中的 expert 必须与 fused_moe_ 实际传入 get_2stage_cfgs 的 E 完全一致。
    对于 EP-off 场景,E = w1.shape[0](即 nroutedexperts,不含 shared expert)。
    对于 EP-on 场景,E = expert_mask.numel()(local expert 数)。
    topk 必须与 nexpertspertoken 完全一致。
    """
    key = (_CU_NUM, token, model_dim, inter_dim, expert, topk,
           _ACT, _DT, _QA, _QW, _QT, True, False)
    _CUSTOM_CONFIGS[key] = {
        "block_m": block_m,
        "kernelName1": k1,
        "kernelName2": k2,
    }

# ────────────────────────────────────────────────────────────────────────────
# 关键修复:topk 从 9 改为 8(nexpertspertoken=8),expert 从 257 改为 256
# ────────────────────────────────────────────────────────────────────────────

# EP-off: E=256, topk=8, d_expert=256
# MI355X 有 304 CU,bs=16 时 token*topk=128,tile_m=32 最优
_add( 16, 7168,  256, 256, 8, 32, K1_SMALL, K2_FLY32)
# bs=128: token*topk=1024,tile_m=64,用 tile_n=256 充分利用带宽
_add(128, 7168,  256, 256, 8, 64, K1_LARGE, K2_FLY64)
# bs=512: token*topk=4096,tile_m=64
_add(512, 7168,  256, 256, 8, 64, K1_LARGE, K2_FLY64)

# EP-on: E=32, topk=8, d_expert=512
_add( 16, 7168,  512, 32, 8, 32, K1_SMALL, K2_FLY32)
_add(128, 7168,  512, 32, 8, 32, K1_MED,   K2_FLY64)
_add(512, 7168,  512, 32, 8, 64, K1_LARGE, K2_FLY64)

# EP-on: E=32, topk=8, d_expert=2048
# d_expert=2048 的 stage2 K 维度很大,tile_n=128 可能更优(减少 atomic 冲突)
_add( 16, 7168, 2048, 32, 8, 32, K1_SMALL, K2_FLY32_128)
_add(128, 7168, 2048, 32, 8, 32, K1_SMALL, K2_FLY32_128)
_add(512, 7168, 2048, 32, 8, 64, K1_LARGE, K2_FLY64_128)

# ── 调试:打印所有注册的 key,用于上线前核对 ─────────────────────────────────
print(f"[custom] Registered {len(_CUSTOM_CONFIGS)} configs:", file=sys.stderr)
for k in _CUSTOM_CONFIGS:
    print(f"  cu={k[0]} tok={k[1]} mdim={k[2]} idim={k[3]} E={k[4]} topk={k[5]}", file=sys.stderr)

# ── Monkey-patch get_2stage_cfgs ──────────────────────────────────────────────
_orig = _fm.get_2stage_cfgs.__wrapped__

@functools.lru_cache(maxsize=2048)
def _patched(
    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,
):
    key = (_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 = _CUSTOM_CONFIGS.get(key)

    if cfg is not None and _FLYDSL_OK:
        print(f"[custom] HIT key=cu{key[0]},tok{key[1]},mdim{key[2]},idim{key[3]},E{key[4]},topk{key[5]}"
              f" k1={cfg['kernelName1'][-30:]}, k2={cfg['kernelName2']}", file=sys.stderr)
        stage1_func = functools.partial(
            _fm.ck_moe_stage1,
            kernelName=cfg["kernelName1"],
            activation=activation,
            quant_type=q_type,
            dtype=dtype,
            splitk=0,
            use_non_temporal_load=False,
        )
        stage2_func = functools.partial(
            _fm._flydsl_stage2_wrapper,
            kernelName=cfg["kernelName2"],
        )
        return MOEMetadata(
            stage1_func, stage2_func,
            int(cfg["block_m"]), 0, False,
        )

    # MISS: 打印完整 key 便于诊断
    print(f"[custom] MISS key=cu{key[0]},tok{key[1]},mdim{key[2]},idim{key[3]},E{key[4]},topk{key[5]}"
          f" act={key[6]} dt={key[7]} qa={key[8]} qw={key[9]} qt={key[10]}"
          f" g1u1={key[11]} dws1={key[12]}, flydsl={_FLYDSL_OK}", file=sys.stderr)
    return _orig(
        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,
    )

_fm.get_2stage_cfgs = _patched
aiter.fused_moe.get_2stage_cfgs = _patched

# ── custom_kernel ─────────────────────────────────────────────────────────────
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

    hidden_pad       = config["d_hidden_pad"] - config["d_hidden"]
    intermediate_pad = config["d_expert_pad"] - config["d_expert"]

    return 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,
    )
scrolls · 190 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