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
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.
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, torchfrom task import input_t, output_t-- from aiter import ActivationType, QuantType+ from aiter import ActivationType, QuantType, dtypesfrom aiter.fused_moe import fused_moe- import aiter.fused_moe as _fused_moe_moduleimport 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 pdfrom 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