submission 566515
Danishlynx · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 171 lines, June 9 Researcher Reciprocity License v1.0.
submission_v91_blockm_tuned.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-566515?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:8d486287be63930132e4f9e273a80983ac237c72abbf5d81f1153b9c64ccd5d3
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 v91 — Aggressive block_m tuning for ALL shapes + flydsl stage2 tk=128.Kernel source
submission_v91_blockm_tuned.py171 lines
"""
MoE MXFP4 v91 — Aggressive block_m tuning for ALL shapes + flydsl stage2 tk=128.
Override block_m heuristic with empirically better values:
- E=257 bs=512: block_m=64 (default=128, distributes better across 256 CUs)
- E=33 bs=128 d=512: block_m=32 (from v75)
- E=33 bs=128 d=2048: block_m=32 (from v75)
- E=33 bs=512 d=2048: block_m=64 (from v75)
- E=33 bs=512 d=512: block_m=64 (try instead of default 128)
"""
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 = {
# E=33 shapes
(128, 9, 33, 512): 32,
(128, 9, 33, 2048): 32,
(512, 9, 33, 512): 64,
(512, 9, 33, 2048): 64,
# E=257 shapes
(512, 9, 257, 256): 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_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):
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=128,
a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", mode="atomic",
w2_scale=w2_scale, a2_scale=a2_scale, sorted_weights=sorted_weights,
)
_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
def _compile_flydsl():
global _flydsl_ready
try:
from aiter.ops.flydsl.moe_kernels import _get_compiled_stage2
for E, inter_dim in [(257, 256), (33, 512), (33, 2048)]:
_get_compiled_stage2(
model_dim=7168, inter_dim=inter_dim, experts=E, topk=9,
tile_m=32, tile_n=128, tile_k=128,
doweight=True, a_dtype="fp4", b_dtype="fp4",
out_dtype="bf16", accumulate=True,
)
_flydsl_ready = True
print("flydsl ready (v91)", file=sys.stderr)
except Exception as e:
print(f"flydsl FAILED: {e}", file=sys.stderr)
_warmed = False
def _warmup():
global _warmed
if _warmed:
return
_warmed = True
_compile_flydsl()
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()
except Exception as e:
print(f"Warmup FAIL: E={E}: {e}", file=sys.stderr)
print("Warmup complete (v91)", 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 · 171 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 566023.
"""- MoE MXFP4 v75 — tile_k=128 for ALL shapes including E=33 d=2048.+ MoE MXFP4 v91 — Aggressive block_m tuning for ALL shapes + flydsl stage2 tk=128.- v71 showed tile_k=128 helps E=33 d=512 (-16% on bs16).- Test if tile_k=128 also helps E=33 d=2048 (K=2048 → 16 vs 8 iterations).+ Override block_m heuristic with empirically better values:+ - E=257 bs=512: block_m=64 (default=128, distributes better across 256 CUs)+ - E=33 bs=128 d=512: block_m=32 (from v75)+ - E=33 bs=128 d=2048: block_m=32 (from v75)+ - E=33 bs=512 d=2048: block_m=64 (from v75)+ - E=33 bs=512 d=512: block_m=64 (try instead of default 128)"""import osos.environ["AITER_USE_NT"] = "1"⋯ 6 unchanged linesfrom aiter.fused_moe import fused_moeimport aiter.fused_moe as _fmoe- # Block_m overrides (from v49)_block_m_overrides = {+ # E=33 shapes(128, 9, 33, 512): 32,(128, 9, 33, 2048): 32,+ (512, 9, 33, 512): 64,(512, 9, 33, 2048): 64,+ # E=257 shapes+ (512, 9, 257, 256): 64,}@functools.lru_cache(maxsize=2048)⋯ 28 unchanged linesw2_scale=None, a2_scale=None, sorted_weights=None,**_kwargs):from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2- # tile_k=128 for everythingflydsl_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=128,- a_dtype="fp4",- b_dtype="fp4",- out_dtype="bf16",- mode="atomic",- w2_scale=w2_scale,- a2_scale=a2_scale,- sorted_weights=sorted_weights,+ 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=128,+ a_dtype="fp4", b_dtype="fp4", out_dtype="bf16", mode="atomic",+ w2_scale=w2_scale, a2_scale=a2_scale, sorted_weights=sorted_weights,)⋯ 13 unchanged linesglobal _flydsl_readytry:from aiter.ops.flydsl.moe_kernels import _get_compiled_stage2- import time- print("Compiling flydsl stage2 (v75 tk=128 everywhere)...", file=sys.stderr)-- t0 = time.time()- _get_compiled_stage2(- model_dim=7168, inter_dim=256, experts=257, topk=9,- tile_m=32, tile_n=128, tile_k=128,- doweight=True, a_dtype="fp4", b_dtype="fp4",- out_dtype="bf16", accumulate=True,- )- print(f" E=257 d=256 tk=128: {time.time()-t0:.1f}s", file=sys.stderr)-- for inter_dim in [512, 2048]:- t0 = time.time()+ for E, inter_dim in [(257, 256), (33, 512), (33, 2048)]:_get_compiled_stage2(- model_dim=7168, inter_dim=inter_dim, experts=33, topk=9,+ model_dim=7168, inter_dim=inter_dim, experts=E, topk=9,tile_m=32, tile_n=128, tile_k=128,doweight=True, a_dtype="fp4", b_dtype="fp4",out_dtype="bf16", accumulate=True,)- print(f" E=33 d={inter_dim} tk=128: {time.time()-t0:.1f}s", file=sys.stderr)-_flydsl_ready = True- print("flydsl ready (v75)", file=sys.stderr)+ print("flydsl ready (v91)", file=sys.stderr)except Exception as e:print(f"flydsl FAILED: {e}", file=sys.stderr)- import traceback- traceback.print_exc(file=sys.stderr)_warmed = False⋯ 37 unchanged lineshidden_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 (v75)", file=sys.stderr)+ print(f"Warmup FAIL: E={E}: {e}", file=sys.stderr)+ print("Warmup complete (v91)", file=sys.stderr)_warmup()
scrolls · 114 diff lines total
Best evidence level for this revision: reported
JSON