submission 597664
Andrew Choi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 224 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-597664?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:978bbb7ef9003a47c23f9f52a2a3004f05cf5ea3bc51950d59ae2afb998df1e7
license declaredunknown
license concludedunknown
authorsAndrew Choi
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
(7168, 512, 33, 9, 32, 128, 256, True, 'fp4', 'fp4', 'bf16'),persistent-kernel
3. HSA_NO_SCRATCH_RECLAIM=1 (persistent GPU scratch)split-k
kwargs['split_k'] = 1Kernel source
submission.py224 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
import os
import sys
os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ["HSA_NO_SCRATCH_RECLAIM"] = "1"
os.environ["HSA_NO_SCRATCH_THREAD_LIMITER"] = "1"
import torch
import functools
from typing import Dict
from task import input_t, output_t
from aiter import ActivationType, QuantType, dtypes
import aiter
import aiter.fused_moe as _aiter_fused_moe_module
# Import flydsl
from aiter.ops.flydsl import flydsl_moe_stage1, flydsl_moe_stage2
# ---- Pre-compile flydsl stage2 kernels for all benchmark shapes ----
# This moves the FlyDSL MLIR compilation out of the first benchmark call,
# similar to how AITER pre-builds CK modules at import time.
try:
from aiter.ops.flydsl import moe_kernels as _moe_kernels
if hasattr(_moe_kernels, 'compile_flydsl_moe_stage2'):
_precompile_configs = [
# E=33/bs=512/d=512: tile_m=32 (forced override from heuristic block_m=64)
(7168, 512, 33, 9, 32, 128, 256, True, 'fp4', 'fp4', 'bf16'),
# E=33/bs=512/d=2048: tile_m=32 (forced override from heuristic block_m=128)
(7168, 2048, 33, 9, 32, 128, 256, True, 'fp4', 'fp4', 'bf16'),
# E=257/d=256: tile_m=32 (same config as CK CSV block_m=32)
(7168, 256, 257, 9, 32, 128, 256, True, 'fp4', 'fp4', 'bf16'),
]
for cfg in _precompile_configs:
try:
_moe_kernels.compile_flydsl_moe_stage2(*cfg)
except Exception:
pass # silently skip if compilation fails for some shape
except Exception:
pass # flydsl not available or compilation error
# ---- Shape-dependent ksplit monkey-patch (for cktile path on small-batch E=33) ----
_original_get_ksplit = _aiter_fused_moe_module.get_ksplit
def _custom_get_ksplit(token, topk, expert, inter_dim, model_dim):
if expert <= 64 and token <= 128:
return 2
return _original_get_ksplit(token, topk, expert, inter_dim, model_dim)
_aiter_fused_moe_module.get_ksplit = _custom_get_ksplit
# ---- cktile_moe_stage1 split_k=1 monkey-patch ----
if hasattr(_aiter_fused_moe_module, 'cktile_moe_stage1'):
_original_cktile_moe_stage1 = _aiter_fused_moe_module.cktile_moe_stage1
def _patched_cktile_moe_stage1(*args, **kwargs):
kwargs['split_k'] = 1
return _original_cktile_moe_stage1(*args, **kwargs)
_aiter_fused_moe_module.cktile_moe_stage1 = _patched_cktile_moe_stage1
# ---- Replace ck_moe_stage2_fwd with flydsl_moe_stage2 for fp4x2 shapes ----
# FlyDSL generates optimized kernels for the down-projection GEMM.
# Available kernels: tile_m={32,64}, tile_n={128,256}, tile_k=256, mode={atomic,reduce}
# We use tile_m=32, tile_n=128, tile_k=256, mode=atomic for all shapes.
def _flydsl_stage2_replacement(
inter_states, # a2: (token_num, topk, inter_dim) -- fp4x2 quantized
w1, # gate_up_weight (unused in stage2)
w2, # down_weight shuffled
sorted_token_ids,
sorted_expert_ids,
num_valid_ids,
out, # moe_out: (M, model_dim) bf16
topk,
w2_scale=None,
a2_scale=None,
block_m=32,
_tile_m_override=None,
sorted_weights=None,
**kwargs,
):
"""Wrapper to call flydsl_moe_stage2 with ck_moe_stage2_fwd's call signature."""
# Use override tile_m if provided, otherwise use block_m (from CK heuristic)
tile_m = _tile_m_override if _tile_m_override is not None else block_m
# Clamp to flydsl's supported values {32, 64}
if tile_m > 64:
tile_m = 64
elif tile_m < 32:
tile_m = 32
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=128, # default
tile_k=256, # default
a_dtype='fp4',
b_dtype='fp4',
out_dtype='bf16',
mode='atomic',
w2_scale=w2_scale,
a2_scale=a2_scale,
sorted_weights=sorted_weights,
)
# Monkey-patch get_2stage_cfgs to replace stage2 with flydsl for standard 2-stage path
_original_get_2stage_cfgs = _aiter_fused_moe_module.get_2stage_cfgs.__wrapped__ \
if hasattr(_aiter_fused_moe_module.get_2stage_cfgs, '__wrapped__') \
else _aiter_fused_moe_module.get_2stage_cfgs
@functools.lru_cache(maxsize=2048)
def _custom_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 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,
)
# For the standard 2-stage path (not cktile, not 1-stage), replace stage2 with flydsl
# for ALL expert counts on fp4x2 per_1x32 workloads.
# E=33/bs=512: proven -20% (d=512), -15% (d=2048) with tile_m=32
# E=257: flydsl stage2 replaces CSV-tuned CK stage2 with tile_m=32
if (
not metadata.run_1stage
and metadata.ksplit <= 1 # not cktile path
and q_type == QuantType.per_1x32
and q_dtype_w == dtypes.fp4x2
and metadata.stage2 is not None
):
# Force tile_m=32 for all shapes:
# - E=33/bs=512: overrides heuristic block_m=64/128 for better CU utilization
# - E=257: matches CSV block_m=32, flydsl replaces CK stage2
tile_m_override = 32
# Replace stage2 with flydsl
from aiter.fused_moe import MOEMetadata
return MOEMetadata(
metadata.stage1,
functools.partial(
_flydsl_stage2_replacement,
block_m=metadata.block_m,
_tile_m_override=tile_m_override,
),
metadata.block_m,
metadata.ksplit,
metadata.run_1stage,
metadata.has_bias,
metadata.use_non_temporal_load,
)
return metadata
_aiter_fused_moe_module.get_2stage_cfgs = _custom_get_2stage_cfgs
from aiter.fused_moe import fused_moe
def custom_kernel(data: input_t) -> output_t:
"""
Optimized DeepSeek-R1 MXFP4 MoE kernel.
Key optimizations:
1. AITER_USE_OPUS_MOE_SORTING=1 (faster sorting kernel)
2. HIP_FORCE_DEV_KERNARG=1 (device-side kernel args)
3. HSA_NO_SCRATCH_RECLAIM=1 (persistent GPU scratch)
4. HSA_NO_SCRATCH_THREAD_LIMITER=1 (full wave occupancy)
5. Shape-dependent ksplit (cktile for E<=64/bs<=128)
6. cktile_moe_stage1 split_k=1
7. flydsl_moe_stage2 for ALL standard 2-stage shapes (E=33 and E=257)
8. tile_m=32 override for flydsl stage2
9. Pre-compile flydsl kernels at import time (including E=257)
"""
(
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"]
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
scrolls · 224 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