submission 749778
coderwhisper · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 278 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-749778?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:26ba3775ec1b519f0d8a587038fda5e4db5152e1d5e1eed8a59a981e3c08b12e
license declaredunknown
license concludedunknown
authorscoderwhisper
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
_meta_cache[key] = ('mxfp4', block_m, metadata, tile_m, use_flydsl_s2,Kernel source
submission.py278 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
import os
import sys
import pathlib
import functools
# ── Patch FlyDSL bug BEFORE importing aiter ──
_FLYDSL_PATCH_TARGET = pathlib.Path(
"/home/runner/aiter/aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py"
)
_PATCHED = False
if _FLYDSL_PATCH_TARGET.exists():
_src = _FLYDSL_PATCH_TARGET.read_text()
if "b_tile_in_gate," in _src:
_lines = _src.split('\n')
_fixed_lines = []
_skip_next = 0
for _i, _line in enumerate(_lines):
if _skip_next > 0:
_skip_next -= 1
continue
if 'b_tile_in_gate,' in _line:
_indent = _line[:len(_line) - len(_line.lstrip())]
_fixed_lines.append(f"{_indent}b_tile_in,")
_skip_next = 1
elif 'b_scale_gate=None,' in _line:
_indent = _line[:len(_line) - len(_line.lstrip())]
_fixed_lines.append(f"{_indent}b_scale=None,")
_skip_next = 1
else:
_fixed_lines.append(_line)
_fixed = '\n'.join(_fixed_lines)
if _fixed != _src:
try:
_FLYDSL_PATCH_TARGET.write_text(_fixed)
import shutil
for _pyc in _FLYDSL_PATCH_TARGET.parent.rglob("__pycache__"):
shutil.rmtree(_pyc, ignore_errors=True)
_PATCHED = True
except PermissionError:
pass
else:
_PATCHED = True
import torch
import aiter
from task import input_t, output_t
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import (
get_2stage_cfgs,
get_inter_dim,
get_padded_M,
get_ksplit,
use_nt,
)
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
from aiter.ops.flydsl import is_flydsl_available
_HAS_FLYDSL = is_flydsl_available() and _PATCHED
if _HAS_FLYDSL:
from aiter.ops.flydsl import flydsl_moe_stage2
print(f"[submission] FlyDSL available for stage2", file=sys.stderr)
_meta_cache = {}
def _has_tuned_stage2(metadata):
"""Check if metadata.stage2 has a tuned (non-empty) kernel name."""
if isinstance(metadata.stage2, functools.partial):
kn = metadata.stage2.keywords.get('kernelName', '')
return bool(kn)
return False
def custom_kernel(data: input_t) -> output_t:
(
hidden_states,
_gate_up_weight,
_down_weight,
_gate_up_weight_scale,
_down_weight_scale,
w1,
w2,
w1_scale,
w2_scale,
topk_weights,
topk_ids,
config,
) = data
hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
intermediate_pad = config["d_expert_pad"] - config["d_expert"]
M = hidden_states.shape[0]
E = w1.shape[0]
topk = topk_ids.shape[1]
_, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)
is_shuffled = getattr(w1, "is_shuffled", False)
dtype = hidden_states.dtype
key = (M, E, model_dim, inter_dim)
if key not in _meta_cache:
device = hidden_states.device
if M <= 128:
# ksplit=2 bf16 cktile path
os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1"
os.environ["AITER_KSPLIT"] = "2"
get_ksplit.cache_clear()
use_nt.cache_clear()
metadata = get_2stage_cfgs(
get_padded_M(M), model_dim, inter_dim, E, topk, dtype,
dtypes.fp4x2, dtypes.fp4x2, QuantType.per_1x32,
True, ActivationType.Silu, False,
hidden_pad, intermediate_pad, is_shuffled,
)
block_m = int(metadata.block_m)
ksplit = int(metadata.ksplit)
n_pad_zeros_1 = intermediate_pad // 64 * 64 * 2
k_pad_zeros_1 = hidden_pad // 128 * 128
n_pad_zeros_2 = hidden_pad // 64 * 64
k_pad_zeros_2 = intermediate_pad // 128 * 128
_, n1, k1 = w1.shape
_, k2, n2 = w2.shape
D = n2 if k2 == k1 else n2 * 2
max_num_tokens_padded = int(M * topk + E * block_m - topk)
max_num_m_blocks = (max_num_tokens_padded + block_m - 1) // block_m
s_ids = torch.empty(max_num_tokens_padded, dtype=torch.int32, device=device)
s_weights = torch.empty(max_num_tokens_padded, dtype=torch.float32, device=device)
s_expert_ids = torch.empty(max_num_m_blocks, dtype=torch.int32, device=device)
n_valid = torch.empty(2, dtype=torch.int32, device=device)
moe_buf = torch.empty((M, model_dim), dtype=dtype, device=device)
tmp_out = torch.empty((M, topk, w1.shape[1]), dtype=dtype, device=device)
a2 = torch.empty((M, topk, D), dtype=dtype, device=device)
_meta_cache[key] = ('direct', block_m, ksplit,
s_ids, s_weights, s_expert_ids, n_valid, moe_buf,
n_pad_zeros_1, k_pad_zeros_1, n_pad_zeros_2, k_pad_zeros_2,
tmp_out, a2)
else:
# M>128: MXFP4 with CK stage1
os.environ["AITER_BYPASS_TUNE_CONFIG"] = "0"
os.environ["AITER_KSPLIT"] = "0"
get_ksplit.cache_clear()
use_nt.cache_clear()
metadata = get_2stage_cfgs(
get_padded_M(M), model_dim, inter_dim, E, topk, dtype,
dtypes.fp4x2, dtypes.fp4x2, QuantType.per_1x32,
True, ActivationType.Silu, False,
hidden_pad, intermediate_pad, is_shuffled,
)
block_m = int(metadata.block_m)
# Decision: use tuned CK stage2 if available, else FlyDSL stage2
has_tuned_s2 = _has_tuned_stage2(metadata)
use_flydsl_s2 = _HAS_FLYDSL and not has_tuned_s2
# For FlyDSL stage2: cap block_m at 64 (max supported tile_m)
if use_flydsl_s2 and block_m > 64:
block_m = 64
# Verify FlyDSL compatibility after capping
use_flydsl_s2 = use_flydsl_s2 and block_m <= 64
tile_m = block_m if use_flydsl_s2 else 0
max_num_tokens_padded = int(M * topk + E * block_m - topk)
max_num_m_blocks = (max_num_tokens_padded + block_m - 1) // block_m
s_ids = torch.empty(max_num_tokens_padded, dtype=torch.int32, device=device)
s_weights = torch.empty(max_num_tokens_padded, dtype=torch.float32, device=device)
s_expert_ids = torch.empty(max_num_m_blocks, dtype=torch.int32, device=device)
n_valid = torch.empty(2, dtype=torch.int32, device=device)
moe_buf = torch.empty((M, model_dim), dtype=dtype, device=device)
a2 = torch.empty((M, topk, inter_dim), dtype=dtype, device=device)
_meta_cache[key] = ('mxfp4', block_m, metadata, tile_m, use_flydsl_s2,
s_ids, s_weights, s_expert_ids, n_valid, moe_buf, a2)
cached = _meta_cache[key]
if cached[0] == 'direct':
(_, block_m, ksplit,
s_ids, s_weights, s_expert_ids, n_valid, moe_buf,
n_pad_zeros_1, k_pad_zeros_1, n_pad_zeros_2, k_pad_zeros_2,
tmp_out, a2) = cached
w1s = w1_scale.view(dtypes.fp8_e8m0)
w2s = w2_scale.view(dtypes.fp8_e8m0)
aiter.moe_sorting_fwd(
topk_ids, topk_weights,
s_ids, s_weights, s_expert_ids, n_valid, moe_buf,
E, block_m, None, None, 0,
)
tmp_out.zero_()
aiter.moe_cktile2stages_gemm1(
hidden_states, w1, tmp_out,
s_ids, s_expert_ids, n_valid,
topk, n_pad_zeros_1, k_pad_zeros_1,
None, None, w1s, None,
ActivationType.Silu, block_m, ksplit,
)
aiter.silu_and_mul(a2, tmp_out)
aiter.moe_cktile2stages_gemm2(
a2, w2, moe_buf,
s_ids, s_expert_ids, n_valid,
topk, n_pad_zeros_2, k_pad_zeros_2,
s_weights, None, w2s, None,
ActivationType.Silu, block_m,
)
return moe_buf
else:
(_, block_m, metadata, tile_m, use_flydsl_s2,
s_ids, s_weights, s_expert_ids, n_valid, moe_buf, a2) = cached
w1s = w1_scale.view(dtypes.fp8_e8m0)
w2s = w2_scale.view(dtypes.fp8_e8m0)
aiter.moe_sorting_fwd(
topk_ids, topk_weights,
s_ids, s_weights, s_expert_ids, n_valid, moe_buf,
E, block_m, None, None, 0,
)
# Quantize activations to MXFP4
a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
hidden_states, sorted_ids=s_ids, num_valid_ids=n_valid,
token_num=M, topk=1, block_size=block_m,
)
# Stage 1: CK GEMM (uses tuned kernel if available)
a2 = metadata.stage1(
a1, w1, w2,
s_ids, s_expert_ids, n_valid,
a2, topk,
block_m=block_m,
a1_scale=a1_scale,
w1_scale=w1s,
sorted_weights=None,
)
# Quantize intermediate to MXFP4
a2_flat = a2.view(-1, a2.shape[-1])
a2_q, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
a2_flat, sorted_ids=s_ids, num_valid_ids=n_valid,
token_num=M, topk=topk, block_size=block_m,
)
a2_q = a2_q.view(M, topk, -1)
# Stage 2: tuned CK or FlyDSL
if use_flydsl_s2:
flydsl_moe_stage2(
a2_q, w2, s_ids, s_expert_ids, n_valid,
out=moe_buf, topk=topk,
tile_m=tile_m, tile_n=128, tile_k=256,
a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
mode="reduce",
w2_scale=w2s, a2_scale=a2_scale, sorted_weights=s_weights,
)
else:
metadata.stage2(
a2_q, w1, w2,
s_ids, s_expert_ids, n_valid,
moe_buf, topk,
w2_scale=w2s,
a2_scale=a2_scale,
block_m=block_m,
sorted_weights=s_weights,
)
return moe_buf
scrolls · 278 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