submission 563502
Danishlynx · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 211 lines, June 9 Researcher Reciprocity License v1.0.
submission_v59_flydsl_stage2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-563502?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:4ef6febf1bbafa6cde20bd4fc65e64f6e913d8ec978e0f058ddf16f7d37f6be9
license declaredunknown
license concludedunknown
authorsDanishlynx
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MoE MXFP4 v59 — flydsl stage2 replacement.Kernel source
submission_v59_flydsl_stage2.py211 lines
"""
MoE MXFP4 v59 — flydsl stage2 replacement.
Strategy:
1. Standard warmup (preshuffle_off, fast ~102s)
2. Pre-compile flydsl stage2 kernels (~5s total)
3. Wrap get_2stage_cfgs to replace stage2 with flydsl in ALL metadata
4. Benchmark's first call triggers preshuffle_on JIT (stage1 needs it)
but stage2 uses flydsl instead of CK
5. Block_m overrides from v49 included
flydsl stage2 uses MLIR-compiled kernels that may be faster than CK.
"""
import os
os.environ["AITER_USE_NT"] = "1"
import sys
import functools
import torch
from task import input_t, output_t
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe
import aiter.fused_moe as _fmoe
# Block_m overrides (from v49)
_block_m_overrides = {
(128, 9, 33, 512): 32,
(128, 9, 33, 2048): 32,
(512, 9, 33, 2048): 64,
}
@functools.lru_cache(maxsize=2048)
def _custom_get_block_size_M(token, topk, expert, inter_dim):
key = (token, topk, expert, inter_dim)
if key in _block_m_overrides:
return _block_m_overrides[key]
cu_num = _fmoe.get_cu_num()
tileN = 128
tgN = (inter_dim + tileN - 1) // tileN
support_list = [32, 64, 128]
tmp = []
for el in support_list:
max_num_tokens = token * topk + expert * el - topk
tg_num = tgN * (max_num_tokens + el - 1) // el
rnd = (tg_num + cu_num - 1) // cu_num
empty = cu_num - tg_num % cu_num
tmp.append((rnd, empty, el))
return sorted(tmp, key=lambda x: x[:2])[0][-1]
_fmoe.get_block_size_M = _custom_get_block_size_M
try:
_fmoe.get_2stage_cfgs.cache_clear()
except:
pass
# ===== flydsl stage2 wrapper =====
_flydsl_ready = False
def _flydsl_stage2_wrapper(inter_states, w1, w2, sorted_token_ids,
sorted_expert_ids, num_valid_ids, out, topk,
w2_scale=None, a2_scale=None, sorted_weights=None,
**_kwargs):
"""Replace CK stage2 with flydsl_moe_stage2."""
from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2
flydsl_moe_stage2(
inter_states=inter_states,
w2=w2,
sorted_token_ids=sorted_token_ids,
sorted_expert_ids=sorted_expert_ids,
num_valid_ids=num_valid_ids,
out=out,
topk=topk,
tile_m=32,
tile_n=128,
tile_k=256,
a_dtype="fp4",
b_dtype="fp4",
out_dtype="bf16",
mode="atomic",
w2_scale=w2_scale,
a2_scale=a2_scale,
sorted_weights=sorted_weights,
)
# ===== Wrap get_2stage_cfgs to inject flydsl stage2 =====
_orig_get_2stage_cfgs = _fmoe.get_2stage_cfgs
@functools.lru_cache(maxsize=2048)
def _patched_get_2stage_cfgs(*args, **kwargs):
metadata = _orig_get_2stage_cfgs(*args, **kwargs)
if _flydsl_ready and not metadata.run_1stage and metadata.stage2 is not None:
metadata.stage2 = functools.partial(_flydsl_stage2_wrapper)
return metadata
_fmoe.get_2stage_cfgs = _patched_get_2stage_cfgs
# ===== Pre-compile flydsl stage2 kernels =====
def _compile_flydsl():
global _flydsl_ready
try:
from aiter.ops.flydsl.moe_kernels import _get_compiled_stage2
import time
print("Compiling flydsl stage2 kernels...", file=sys.stderr)
shapes = [
(257, 256), # E=257, d_expert_pad=256
(33, 512), # E=33, d_expert_pad=512
(33, 2048), # E=33, d_expert_pad=2048
]
for E, inter_dim in shapes:
t0 = time.time()
_get_compiled_stage2(
model_dim=7168,
inter_dim=inter_dim,
experts=E,
topk=9,
tile_m=32,
tile_n=128,
tile_k=256,
doweight=True,
a_dtype="fp4",
b_dtype="fp4",
out_dtype="bf16",
accumulate=True,
)
t1 = time.time()
print(f" flydsl E={E} inter={inter_dim}: {t1-t0:.1f}s", file=sys.stderr)
_flydsl_ready = True
print("flydsl stage2 compilation complete", file=sys.stderr)
except Exception as e:
print(f"flydsl compilation FAILED: {e}", file=sys.stderr)
import traceback
traceback.print_exc(file=sys.stderr)
# ===== Standard warmup =====
_warmed = False
def _warmup():
global _warmed
if _warmed:
return
_warmed = True
# Compile flydsl first (fast, ~5s total)
_compile_flydsl()
# Standard warmup (preshuffle_off, ~102s)
configs = [
(2, 256, 1, 7168, 256, 8),
(2, 32, 1, 7168, 512, 8),
(2, 32, 1, 7168, 2048, 8),
]
for bs, n_routed, n_shared, d_hidden, d_expert, n_experts_per_token in configs:
E = n_routed + n_shared
total_topk = n_experts_per_token + n_shared
d_hidden_pad = ((d_hidden + 255) // 256) * 256
d_expert_pad = ((d_expert + 255) // 256) * 256
h = torch.randn(bs, d_hidden, dtype=torch.bfloat16, device="cuda")
w1 = torch.empty(E, 2 * d_expert_pad, d_hidden_pad // 2,
dtype=torch.float4_e2m1fn_x2, device="cuda")
w2 = torch.empty(E, d_hidden_pad, d_expert_pad // 2,
dtype=torch.float4_e2m1fn_x2, device="cuda")
w1_s = torch.empty(E, 2 * d_expert_pad, d_hidden_pad // 32,
dtype=torch.float8_e8m0fnu, device="cuda")
w2_s = torch.empty(E, d_hidden_pad, d_expert_pad // 32,
dtype=torch.float8_e8m0fnu, device="cuda")
topk_w = torch.ones(bs, total_topk, dtype=torch.float32, device="cuda")
topk_i = torch.zeros(bs, total_topk, dtype=torch.int32, device="cuda")
for t in range(bs):
for k in range(n_experts_per_token):
topk_i[t, k] = k % n_routed
for k in range(n_shared):
topk_i[t, n_experts_per_token + k] = n_routed + k
try:
fused_moe(h, w1, w2, topk_w, topk_i,
activation=ActivationType.Silu,
quant_type=QuantType.per_1x32,
w1_scale=w1_s, w2_scale=w2_s,
hidden_pad=d_hidden_pad - d_hidden,
intermediate_pad=d_expert_pad - d_expert)
torch.cuda.synchronize()
print(f"Warmup OK: E={E} d_e={d_expert}", file=sys.stderr)
except Exception as e:
print(f"Warmup FAIL: E={E} d_e={d_expert}: {e}", file=sys.stderr)
print("Warmup complete (v59 flydsl_stage2)", file=sys.stderr)
_warmup()
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
return fused_moe(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk_weights,
topk_ids,
activation=ActivationType.Silu,
quant_type=QuantType.per_1x32,
w1_scale=gate_up_weight_scale_shuffled,
w2_scale=down_weight_scale_shuffled,
hidden_pad=config["d_hidden_pad"] - config["d_hidden"],
intermediate_pad=config["d_expert_pad"] - config["d_expert"],
)
scrolls · 211 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 562885.
"""- MoE MXFP4 v49 — Block_m tuning via monkey-patch.+ MoE MXFP4 v59 — flydsl stage2 replacement.- Default block_m values:- M=16, E=257: 32 | M=128, E=257: 32 | M=512, E=257: 128- M=16, E=33: 32 | M=128, E=33: 64 | M=512, E=33 d=512: 128 | M=512, E=33 d=2048: 128+ Strategy:+ 1. Standard warmup (preshuffle_off, fast ~102s)+ 2. Pre-compile flydsl stage2 kernels (~5s total)+ 3. Wrap get_2stage_cfgs to replace stage2 with flydsl in ALL metadata+ 4. Benchmark's first call triggers preshuffle_on JIT (stage1 needs it)+ but stage2 uses flydsl instead of CK+ 5. Block_m overrides from v49 included- The get_block_size_M heuristic selects block_m based on CU utilization.- Let's try:- - E=257 shapes: try block_m=64 for M=128 (default=32)- - E=33 shapes: try block_m=32 for M=128 (default=64), block_m=64 for M=512 (default=128)- These try smaller block_m for better CU occupancy or larger for fewer launches.+ flydsl stage2 uses MLIR-compiled kernels that may be faster than CK."""import osos.environ["AITER_USE_NT"] = "1"import sys+ import functoolsimport torchfrom task import input_t, output_tfrom aiter import ActivationType, QuantTypefrom aiter.fused_moe import fused_moe-- # Monkey-patch get_block_size_Mimport aiter.fused_moe as _fmoe- import functools- _orig_get_block_size_M = _fmoe.get_block_size_M.__wrapped__ if hasattr(_fmoe.get_block_size_M, '__wrapped__') else None+ # Block_m overrides (from v49)+ _block_m_overrides = {+ (128, 9, 33, 512): 32,+ (128, 9, 33, 2048): 32,+ (512, 9, 33, 2048): 64,+ }- # Custom block_m selection- _block_m_overrides = {}-- def _setup_overrides():- """Set block_m overrides for specific shapes."""- # Try different block_m for E=33 shapes- # (token, topk, expert, inter_dim) -> block_m- # Default E=33: M=16→32, M=128→64, M=512→128- # Try: M=128→32 (smaller tiles, more CU parallelism)- _block_m_overrides[(128, 9, 33, 512)] = 32- _block_m_overrides[(128, 9, 33, 2048)] = 32- # Try: M=512→64 for E=33 (more CU parallelism)- _block_m_overrides[(512, 9, 33, 512)] = 64- _block_m_overrides[(512, 9, 33, 2048)] = 64-- _setup_overrides()-@functools.lru_cache(maxsize=2048)def _custom_get_block_size_M(token, topk, expert, inter_dim):key = (token, topk, expert, inter_dim)if key in _block_m_overrides:- result = _block_m_overrides[key]- print(f" block_m override: {key} -> {result}", file=sys.stderr)- return result- # Use original heuristic+ return _block_m_overrides[key]cu_num = _fmoe.get_cu_num()tileN = 128tgN = (inter_dim + tileN - 1) // tileN⋯ 5 unchanged linesrnd = (tg_num + cu_num - 1) // cu_numempty = cu_num - tg_num % cu_numtmp.append((rnd, empty, el))- result = sorted(tmp, key=lambda x: x[:2])[0][-1]- print(f" block_m default: {key} -> {result}", file=sys.stderr)- return result+ return sorted(tmp, key=lambda x: x[:2])[0][-1]- # Apply monkey-patch_fmoe.get_block_size_M = _custom_get_block_size_M-- # Also clear any cached get_2stage_cfgstry:_fmoe.get_2stage_cfgs.cache_clear()except:pass- # Standard warmup+ # ===== flydsl stage2 wrapper =====+ _flydsl_ready = False++ def _flydsl_stage2_wrapper(inter_states, w1, w2, sorted_token_ids,+ sorted_expert_ids, num_valid_ids, out, topk,+ w2_scale=None, a2_scale=None, sorted_weights=None,+ **_kwargs):+ """Replace CK stage2 with flydsl_moe_stage2."""+ from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2+ flydsl_moe_stage2(+ inter_states=inter_states,+ w2=w2,+ sorted_token_ids=sorted_token_ids,+ sorted_expert_ids=sorted_expert_ids,+ num_valid_ids=num_valid_ids,+ out=out,+ topk=topk,+ tile_m=32,+ tile_n=128,+ tile_k=256,+ a_dtype="fp4",+ b_dtype="fp4",+ out_dtype="bf16",+ mode="atomic",+ w2_scale=w2_scale,+ a2_scale=a2_scale,+ sorted_weights=sorted_weights,+ )++ # ===== Wrap get_2stage_cfgs to inject flydsl stage2 =====+ _orig_get_2stage_cfgs = _fmoe.get_2stage_cfgs++ @functools.lru_cache(maxsize=2048)+ def _patched_get_2stage_cfgs(*args, **kwargs):+ metadata = _orig_get_2stage_cfgs(*args, **kwargs)+ if _flydsl_ready and not metadata.run_1stage and metadata.stage2 is not None:+ metadata.stage2 = functools.partial(_flydsl_stage2_wrapper)+ return metadata++ _fmoe.get_2stage_cfgs = _patched_get_2stage_cfgs++ # ===== Pre-compile flydsl stage2 kernels =====+ def _compile_flydsl():+ global _flydsl_ready+ try:+ from aiter.ops.flydsl.moe_kernels import _get_compiled_stage2+ import time+ print("Compiling flydsl stage2 kernels...", file=sys.stderr)++ shapes = [+ (257, 256), # E=257, d_expert_pad=256+ (33, 512), # E=33, d_expert_pad=512+ (33, 2048), # E=33, d_expert_pad=2048+ ]+ for E, inter_dim in shapes:+ t0 = time.time()+ _get_compiled_stage2(+ model_dim=7168,+ inter_dim=inter_dim,+ experts=E,+ topk=9,+ tile_m=32,+ tile_n=128,+ tile_k=256,+ doweight=True,+ a_dtype="fp4",+ b_dtype="fp4",+ out_dtype="bf16",+ accumulate=True,+ )+ t1 = time.time()+ print(f" flydsl E={E} inter={inter_dim}: {t1-t0:.1f}s", file=sys.stderr)++ _flydsl_ready = True+ print("flydsl stage2 compilation complete", file=sys.stderr)+ except Exception as e:+ print(f"flydsl compilation FAILED: {e}", file=sys.stderr)+ import traceback+ traceback.print_exc(file=sys.stderr)++ # ===== Standard warmup =====_warmed = Falsedef _warmup():global _warmed⋯ 1 unchanged linesreturn_warmed = True+ # Compile flydsl first (fast, ~5s total)+ _compile_flydsl()++ # Standard warmup (preshuffle_off, ~102s)configs = [(2, 256, 1, 7168, 256, 8),(2, 32, 1, 7168, 512, 8),(2, 32, 1, 7168, 2048, 8),]-for bs, n_routed, n_shared, d_hidden, d_expert, n_experts_per_token in configs:E = n_routed + n_sharedtotal_topk = n_experts_per_token + n_sharedd_hidden_pad = ((d_hidden + 255) // 256) * 256d_expert_pad = ((d_expert + 255) // 256) * 256- hidden_pad = d_hidden_pad - d_hidden- intermediate_pad = d_expert_pad - d_expert-h = torch.randn(bs, d_hidden, dtype=torch.bfloat16, device="cuda")w1 = torch.empty(E, 2 * d_expert_pad, d_hidden_pad // 2,dtype=torch.float4_e2m1fn_x2, device="cuda")⋯ 10 unchanged linestopk_i[t, k] = k % n_routedfor k in range(n_shared):topk_i[t, n_experts_per_token + k] = n_routed + k-try:- fused_moe(- h, w1, w2, topk_w, topk_i,- activation=ActivationType.Silu,- quant_type=QuantType.per_1x32,- w1_scale=w1_s, w2_scale=w2_s,- hidden_pad=hidden_pad,- intermediate_pad=intermediate_pad,- )+ fused_moe(h, w1, w2, topk_w, topk_i,+ activation=ActivationType.Silu,+ quant_type=QuantType.per_1x32,+ w1_scale=w1_s, w2_scale=w2_s,+ hidden_pad=d_hidden_pad - d_hidden,+ intermediate_pad=d_expert_pad - d_expert)torch.cuda.synchronize()- print(f" Warmup OK: E={E} d_e={d_expert}", file=sys.stderr)+ print(f"Warmup OK: E={E} d_e={d_expert}", file=sys.stderr)except Exception as e:- print(f" Warmup FAIL: E={E} d_e={d_expert}: {e}", file=sys.stderr)+ print(f"Warmup FAIL: E={E} d_e={d_expert}: {e}", file=sys.stderr)- print("Warmup complete (v49 blockm_tune)", file=sys.stderr)+ print("Warmup complete (v59 flydsl_stage2)", file=sys.stderr)_warmup()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_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) = datareturn fused_moe(hidden_states,
scrolls · 257 diff lines total
Best evidence level for this revision: reported
JSON