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
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-k
splitk=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