submission 593539
c3ko · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 165 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-593539?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:ca5c090387390a544c607a4888445ec8a02a63e8501dd056a8c0751f2d223a29
license declaredunknown
license concludedunknown
authorsc3ko
imported2026-08-15
Kernel source
submission.py165 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
import os
import functools as _ft
# ── Performance env vars (set before any HIP/aiter imports) ──
os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"
# ── CK kernel name constants for FP4 MoE 2-stage on gfx950 ──
# Stage1 (gate_up GEMM + SiLU activation):
# Pattern: moe_ck2stages_gemm1_{threads}x{M}x{N}x{K}_{Mw}x{Nw}_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16
_S1_SM = "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_S1_LG = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_S1_XL = "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_S1_XX = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
# Stage2 (down GEMM + weighted reduction):
# Pattern: moe_ck2stages_gemm2_{threads}x{M}x{N}x{K}_{Mw}x{Nw}_MulABScaleExpertWeightShuffled_v{ver}_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16
# 1-wavefront (64-thread) stage2 kernels:
_S2_SM = "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
_S2_M64 = "moe_ck2stages_gemm2_64x64x128x128_1x1_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
_S2_LG = "moe_ck2stages_gemm2_64x128x128x128_1x1_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
# 4-wavefront (256-thread) stage2 kernels:
_S2_256_32 = "moe_ck2stages_gemm2_256x32x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
_S2_256_64 = "moe_ck2stages_gemm2_256x64x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
_S2_256_128 = "moe_ck2stages_gemm2_256x128x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
# FlyDSL stage2 kernels (reduce = separate reduction kernel, atomic = direct output write)
_FD_S2_64 = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
_FD_S2_32 = "flydsl_moe2_afp4_wfp4_bf16_t32x256x256_reduce"
_FD_S2_64A = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_atomic"
_FD_S2_128 = "flydsl_moe2_afp4_wfp4_bf16_t128x256x256_reduce"
_HDR = "cu_num,token,model_dim,inter_dim,expert,topk,act_type,dtype,q_dtype_a,q_dtype_w,q_type,use_g1u1,doweight_stage1,block_m,ksplit,us1,kernelName1,err1,us2,kernelName2,err2,us,run_1stage,tflops,bw,_tag"
_COM = "ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0"
def _row(tok, mdim, idim, E, bm, s1, s2, tag="", ks=0):
return f"256,{tok},{mdim},{idim},{E},9,{_COM},{bm},{ks},0,{s1},0,0,{s2},0,0,0,0,0,{tag}"
_csv_rows = [_HDR]
# ══════════════════════════════════════════════════════════════
# E=257, inter_dim=256 (TP=8 DeepSeek-R1)
# Matches official dsv3_fp4_tuned_fmoe.csv from aiter
# ══════════════════════════════════════════════════════════════
# bs=16: cktile split-K path (extremely sparse: 0.56 tokens/expert avg)
# ksplit=7 valid (7168/7=1024, 1024%256=0) — more CU fill than ks=4 for ultra-sparse
_csv_rows.append(_row(16, 7168, 256, 257, 16, "", "", ks=7))
# bs=32-64: SM/SM block_m=32 (official tuning, very sparse)
for tok in [32, 64]:
_csv_rows.append(_row(tok, 7168, 256, 257, 32, _S1_SM, _S2_SM))
# bs=128: cktile split-K path — try ksplit=4 for more CU fill on sparse E=257
_csv_rows.append(_row(128, 7168, 256, 257, 16, "", "", ks=4))
# bs=256: also try cktile split-K (not benchmarked yet, but likely helps)
_csv_rows.append(_row(256, 7168, 256, 257, 16, "", "", ks=2))
# bs=512: SM stage1 + FlyDSL stage2 (K=256 very small → FlyDSL efficient)
_csv_rows.append(_row(512, 7168, 256, 257, 32, _S1_SM, _FD_S2_32))
# FlyDSL fallback: SM/SM CK 2-stage
_csv_rows.append(_row(512, 7168, 256, 257, 32, _S1_SM, _S2_SM, "flydsl_fallback"))
# bs=1024+: Larger tiles for denser workloads (block_m must match kernel M)
for tok in [1024, 4096, 16384]:
_csv_rows.append(_row(tok, 7168, 256, 257, 64, _S1_XL, _S2_M64))
for tok in [2048, 8192, 32768]:
_csv_rows.append(_row(tok, 7168, 256, 257, 128, _S1_XX, _S2_LG))
# flydsl fallback
for tok in [1024]:
_csv_rows.append(_row(tok, 7168, 256, 257, 32, _S1_SM, _S2_SM, "flydsl_fallback"))
for tok in [2048, 4096, 8192, 16384, 32768]:
_csv_rows.append(_row(tok, 7168, 256, 257, 128, _S1_XX, _S2_LG, "flydsl_fallback"))
# ══════════════════════════════════════════════════════════════
# E=33, inter_dim=512 (TP=4)
# Fewer experts → more tokens/expert → larger tiles benefit
# ══════════════════════════════════════════════════════════════
# bs=16: cktile split-K (4.4 tokens/expert avg, needs CU fill)
# ks=7 for more CU fill on sparse shape
_csv_rows.append(_row(16, 7168, 512, 33, 16, "", "", ks=7))
# bs=128: ~35 tokens/expert. cktile split-K (block_m=16, ks=2).
_csv_rows.append(_row(128, 7168, 512, 33, 16, "", "", ks=2))
# bs=512: ~140 tokens/expert. CK stage1 XL + FlyDSL stage2 reduce (atomic worse: +4.8% contention).
_csv_rows.append(_row(512, 7168, 512, 33, 64, _S1_XL, _FD_S2_64))
# FlyDSL fallback: use CK stage2 if FlyDSL unavailable
_csv_rows.append(_row(512, 7168, 512, 33, 64, _S1_XL, _S2_M64, "flydsl_fallback"))
# ══════════════════════════════════════════════════════════════
# E=33, inter_dim=2048 (EP-enabled)
# Large inter_dim → high compute intensity
# ══════════════════════════════════════════════════════════════
# bs=512: ~140 tokens/expert, inter_dim=2048.
# CK 256_64 stage2 proven best (FlyDSL reduce was equivalent but adds compile risk).
_csv_rows.append(_row(512, 7168, 2048, 33, 64, _S1_XL, _S2_256_64))
# FlyDSL fallback: same CK stage2 (no fallback needed)
# _csv_rows.append(_row(512, 7168, 2048, 33, 64, _S1_XL, _S2_256_64, "flydsl_fallback"))
_csv = "\n".join(_csv_rows)
_cfg_path = "/tmp/_moe_fp4_tuned.csv"
with open(_cfg_path, "w") as _f:
_f.write(_csv)
os.environ["AITER_CONFIG_FMOE"] = _cfg_path
# ── Now import aiter (reads config from env var) ──
import torch
from task import input_t, output_t
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe
from aiter import fused_moe as _fm
# ── Monkey-patch: CKTile block_m=32 for E=33 bs=128 ──
# Default is 16 for token<2048. bm=32 reduces tile count → less launch overhead.
# Tested: bm=16=108µs, bm=32=104µs, bm=64=137µs. bm=32 is the sweet spot.
_orig_cfgs = _fm.get_2stage_cfgs.__wrapped__
@_ft.lru_cache(maxsize=2048)
def _bm_cfgs(*args):
md = _orig_cfgs(*args)
token, expert = args[0], args[3]
if md.ksplit > 0 and token >= 128 and expert <= 64:
md = _fm.MOEMetadata(md.stage1, md.stage2, 32, md.ksplit, md.run_1stage)
return md
_fm.get_2stage_cfgs.cache_clear()
_fm.get_2stage_cfgs = _bm_cfgs
_ACT = ActivationType.Silu
_QT = QuantType.per_1x32
# ── Check FlyDSL availability ──
try:
from aiter.ops.flydsl.utils import is_flydsl_available
print(f"[v71] FlyDSL available: {is_flydsl_available()}")
except Exception as e:
print(f"[v71] FlyDSL check failed: {e}")
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
(
hidden_states, _, _, _, _,
w1, w2, w1_scale, w2_scale,
topk_weights, topk_ids, config,
) = data
return fused_moe(
hidden_states, w1, w2, topk_weights, topk_ids,
activation=_ACT,
quant_type=_QT,
doweight_stage1=False,
w1_scale=w1_scale,
w2_scale=w2_scale,
hidden_pad=config["d_hidden_pad"] - config["d_hidden"],
intermediate_pad=config["d_expert_pad"] - config["d_expert"],
)
scrolls · 165 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