submission 581053
RyanWillie · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 272 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-581053?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:8901a1d2672abfef8db321a645a3a9e710c3ddd0e2ec3ae417580922882740dc
license declaredunknown
license concludedunknown
authorsRyanWillie
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
DeepSeek-R1 style MXFP4 MoE: v006 selective ksplit + FlyDSL stage2 (tile_m=32) + JIT warmup.Kernel source
submission.py272 lines
# Zeus Variant: v037_flydsl_tile_m32_warmup
# Experiment: #037
# Technique: FlyDSL stage2 tile_m=32 for bs512 shapes + JIT warmup at import time
# Mechanism: Same as v034 (tile_m=32) but pre-compiles FlyDSL MLIR kernels at import
# time so compilation occurs during test phase. By benchmark/ranked phases
# the kernels are cached in @functools.cache and execute instantly.
# Small shapes use v006 selective ksplit=2 (CKTile path, no stage2 change needed).
# bs512/e256 uses CSV-tuned CK kernel, left unchanged.
"""
DeepSeek-R1 style MXFP4 MoE: v006 selective ksplit + FlyDSL stage2 (tile_m=32) + JIT warmup.
"""
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
import os
import functools
from contextlib import contextmanager
from dataclasses import dataclass
from typing import Callable
from task import input_t, output_t
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe, get_ksplit
import aiter.fused_moe as _fused_moe_mod
import aiter.ops.flydsl as _flydsl_mod
# ---- JIT warmup: pre-compile FlyDSL MLIR kernels in background thread ----
# FlyDSL MLIR compilation for tile_m=32 is too slow to fit in leaderboard timeout
# if it runs inline. By starting compilation in a background thread at import time,
# it overlaps with CK sorting (~25s) + CKTile (~106s) + CK 2-stage (~104s) JIT
# that happens lazily on first custom_kernel() call. By the time the benchmark
# reaches bs512/e512 or bs512/e2048 shapes, the FlyDSL kernels are already cached
# in _get_compiled_stage2's @functools.cache.
import threading as _threading
import inspect as _inspect
from aiter.ops.flydsl.moe_kernels import _get_compiled_stage2 as _warmup_compile
# Detect server AITER signature: older version lacks `persist_m` parameter
_warmup_sig = _inspect.signature(_warmup_compile.__wrapped__)
_has_persist_m = "persist_m" in _warmup_sig.parameters
# Pre-compile for each bs512 shape that uses FlyDSL stage2.
# Parameters must exactly match what flydsl_moe_stage2() will compute at runtime:
# model_dim=7168 (d_hidden_pad), topk=9, experts=33
# a_dtype="fp4" => inter_dim = d_expert_pad (already doubled by flydsl_moe_stage2)
# doweight=True (sorted_weights are passed), accumulate=True (mode="atomic")
# persist_m=4 (sorted_expert_ids.numel() > 256 for bs512) -- only if server supports it
_WARMUP_CONFIGS_FULL = [
# (model_dim, inter_dim, experts, topk, tile_m, tile_n, tile_k,
# doweight, a_dtype, b_dtype, out_dtype, accumulate, persist_m)
(7168, 512, 33, 9, 32, 128, 256, True, "fp4", "fp4", "bf16", True, 4), # bs512/e512
(7168, 2048, 33, 9, 32, 128, 256, True, "fp4", "fp4", "bf16", True, 4), # bs512/e2048
]
# Strip persist_m if server doesn't support it
_WARMUP_CONFIGS = [cfg if _has_persist_m else cfg[:-1] for cfg in _WARMUP_CONFIGS_FULL]
def _do_warmup():
"""Compile FlyDSL kernels. Runs in background thread to overlap with CK JIT."""
for cfg in _WARMUP_CONFIGS:
_warmup_compile(*cfg)
_warmup_thread = _threading.Thread(target=_do_warmup, daemon=True)
_warmup_thread.start()
# NOTE: We do NOT join here — let it run in background while CK/CKTile compile.
# We join before the first FlyDSL call in the stage2 wrapper below.
# ---- Build FlyDSL stage2 replacement ----
def _build_flydsl_stage2(tile_m=64, tile_n=128, tile_k=256, mode="atomic"):
"""Build a FlyDSL stage2 callable matching ck_moe_stage2_fwd interface."""
def _flydsl_stage2(
inter_states, w1, w2,
sorted_token_ids, sorted_expert_ids, num_valid_ids,
out, topk,
kernelName="", w2_scale=None, a2_scale=None,
sorted_weights=None, **_kw,
):
# Ensure warmup compilation is done before first FlyDSL call.
# After first join, _warmup_thread.is_alive() is False and join() is a no-op.
_warmup_thread.join()
_flydsl_mod.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=tile_m,
tile_n=tile_n,
tile_k=tile_k,
a_dtype="fp4",
b_dtype="fp4",
out_dtype="bf16",
mode=mode,
w2_scale=w2_scale,
a2_scale=a2_scale,
sorted_weights=sorted_weights,
)
return _flydsl_stage2
# ---- Patch get_2stage_cfgs to inject FlyDSL for bs512 shapes ----
_orig_get_2stage_cfgs = _fused_moe_mod.get_2stage_cfgs.__wrapped__
# Shape keys (padded_M, inter_dim, E) to apply FlyDSL stage2
_FLYDSL_TARGETS = {
(512, 512, 33), # bs512/e512
(512, 2048, 33), # bs512/e2048
}
_flydsl_stage2_fn = _build_flydsl_stage2(tile_m=32, tile_n=128, tile_k=256, mode="atomic")
@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,
):
orig = _orig_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,
)
# Only replace stage2 for targeted bs512 shapes with ksplit=0
if (token, inter_dim, expert) in _FLYDSL_TARGETS and orig.ksplit == 0:
# Replace stage2 with FlyDSL, keep everything else
@dataclass
class PatchedMetadata:
stage1: Callable
stage2: Callable
block_m: int
ksplit: int
run_1stage: bool = False
has_bias: bool = False
use_non_temporal_load: bool = True
return PatchedMetadata(
stage1=orig.stage1,
stage2=_flydsl_stage2_fn,
block_m=orig.block_m,
ksplit=orig.ksplit,
run_1stage=orig.run_1stage,
)
return orig
_fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs
# ---- v006 selective ksplit ----
FORCED_KSPLIT = "2"
_ENV_OVERRIDES = {
"AITER_BYPASS_TUNE_CONFIG": "1",
"AITER_KSPLIT": FORCED_KSPLIT,
}
@contextmanager
def _temporary_env(overrides: dict[str, str]):
previous = {key: os.environ.get(key) for key in overrides}
try:
for key, value in overrides.items():
os.environ[key] = value
yield
finally:
for key, value in previous.items():
if value is None:
os.environ.pop(key, None)
else:
os.environ[key] = value
def _clear_ksplit_cache() -> None:
get_ksplit.cache_clear()
def _should_force_ksplit(config: dict) -> bool:
return config["bs"] <= 128 and config["d_expert"] <= 512
def _run_fused_moe(
hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
topk_weights, topk_ids,
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
hidden_pad, intermediate_pad,
):
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,
)
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"]
if not _should_force_ksplit(config):
return _run_fused_moe(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk_weights,
topk_ids,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
hidden_pad,
intermediate_pad,
)
_clear_ksplit_cache()
try:
with _temporary_env(_ENV_OVERRIDES):
_clear_ksplit_cache()
return _run_fused_moe(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk_weights,
topk_ids,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
hidden_pad,
intermediate_pad,
)
finally:
_clear_ksplit_cache()
scrolls · 272 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