submission 738819
chenxingqiang · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 251 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-738819?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:10c12111c35dcf1c7c4318a3ed8c0f44ece6bed2a35a97a879a11128e63f2d27
license declaredunknown
license concludedunknown
authorschenxingqiang
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
a_dtype="fp4",Kernel source
submission.py251 lines
import torch
from task import input_t, output_t
import os
import aiter
import aiter.fused_moe as fused_moe_mod
from aiter import ActivationType, QuantType, dtypes
import functools
from aiter.ops.flydsl.utils import is_flydsl_available
from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage1, flydsl_moe_stage2
from aiter.ops.shuffle import shuffle_scale_a16w4, shuffle_weight_a16w4
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
_orig_get_2stage_cfgs = fused_moe_mod.get_2stage_cfgs
_SORT_CACHE = {}
_W2_LAYOUT_CACHE = {}
_CACHE_LIMIT = 8
_FLYDSL_DISABLED = False
def _cache_put(cache, key, value):
if key in cache:
cache[key] = value
return
if len(cache) >= _CACHE_LIMIT:
cache.clear()
cache[key] = value
def _get_cached_sorting(topk_ids, topk_weights, num_experts, model_dim, block_m, out_dtype):
device = topk_ids.device
token_num, topk = topk_ids.shape
max_num_tokens_padded = int(topk_ids.numel() + num_experts * block_m - topk)
max_num_m_blocks = int((max_num_tokens_padded + block_m - 1) // block_m)
key = (device, token_num, topk, num_experts, model_dim, block_m, out_dtype)
bufs = _SORT_CACHE.get(key)
if bufs is None:
bufs = (
torch.empty(max_num_tokens_padded, dtype=dtypes.i32, device=device),
torch.empty(max_num_tokens_padded, dtype=dtypes.fp32, device=device),
torch.empty(max_num_m_blocks, dtype=dtypes.i32, device=device),
torch.empty(2, dtype=dtypes.i32, device=device),
torch.empty((token_num, model_dim), dtype=out_dtype, device=device),
)
_cache_put(_SORT_CACHE, key, bufs)
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_out = bufs
aiter.moe_sorting_opus_fwd(
topk_ids,
topk_weights,
sorted_ids,
sorted_weights,
sorted_expert_ids,
num_valid_ids,
moe_out,
num_experts,
int(block_m),
None,
None,
0,
)
return sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_out
def _get_cached_w2_layout(down_weight, down_weight_scale, num_experts):
key = (
down_weight.data_ptr(),
down_weight_scale.data_ptr(),
tuple(down_weight.shape),
tuple(down_weight_scale.shape),
down_weight.device,
)
cached = _W2_LAYOUT_CACHE.get(key)
if cached is not None:
return cached
w2_a16w4 = shuffle_weight_a16w4(down_weight, 16, False)
w2_scale_a16w4 = shuffle_scale_a16w4(
down_weight_scale.view(num_experts, -1),
num_experts,
False,
)
_cache_put(_W2_LAYOUT_CACHE, key, (w2_a16w4, w2_scale_a16w4))
return w2_a16w4, w2_scale_a16w4
def _flydsl_moe(
hidden_states,
gate_up_weight_shuffled,
down_weight,
gate_up_weight_scale_shuffled,
down_weight_scale,
topk_weights,
topk_ids,
):
token_num = hidden_states.shape[0]
model_dim = hidden_states.shape[1]
topk = topk_ids.shape[1]
num_experts = gate_up_weight_shuffled.shape[0]
inter_dim = gate_up_weight_shuffled.shape[1] // 2
block_m = 128 if token_num >= 512 and inter_dim >= 1024 else 32
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_out = _get_cached_sorting(
topk_ids,
topk_weights,
num_experts,
model_dim,
block_m,
hidden_states.dtype,
)
w2_a16w4, w2_scale_a16w4 = _get_cached_w2_layout(down_weight, down_weight_scale, num_experts)
a1_fp4, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
hidden_states,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
topk=1,
block_size=block_m,
)
stage1_out = flydsl_moe_stage1(
a=a1_fp4,
w1=gate_up_weight_shuffled,
sorted_token_ids=sorted_ids,
sorted_expert_ids=sorted_expert_ids,
num_valid_ids=num_valid_ids,
topk=topk,
tile_m=block_m,
tile_n=256,
tile_k=256,
a_dtype="fp4",
b_dtype="fp4",
out_dtype="bf16",
act="silu",
w1_scale=gate_up_weight_scale_shuffled,
a1_scale=a1_scale,
sorted_weights=None,
)
a2_fp4, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
stage1_out.view(-1, inter_dim),
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
topk=topk,
block_size=block_m,
)
return flydsl_moe_stage2(
inter_states=a2_fp4.view(token_num, topk, -1),
w2=w2_a16w4,
sorted_token_ids=sorted_ids,
sorted_expert_ids=sorted_expert_ids,
num_valid_ids=num_valid_ids,
out=moe_out,
topk=topk,
tile_m=block_m,
tile_n=256,
tile_k=256,
a_dtype="fp4",
b_dtype="fp4",
out_dtype="bf16",
mode="atomic",
w2_scale=w2_scale_a16w4,
a2_scale=a2_scale,
sorted_weights=sorted_weights,
)
def _force_1stage_cfgs(*args, **kwargs):
meta = _orig_get_2stage_cfgs(*args, **kwargs)
meta.run_1stage = True
act = args[10]
q_type = args[8]
dt = args[5]
q_dt_a = args[6]
q_dt_w = args[7]
is_g1u1 = args[9]
dow_s1 = args[11]
meta.stage1 = functools.partial(
fused_moe_mod.fused_moe_1stage,
kernelName="",
activation=act,
quant_type=q_type,
)
return meta
_USE_1STAGE = os.environ.get("RK_FORCE_MOE_1STAGE", "0") == "1"
if _USE_1STAGE:
fused_moe_mod.get_2stage_cfgs = _force_1stage_cfgs
def custom_kernel(data: input_t) -> output_t:
global _FLYDSL_DISABLED
(
hidden_states,
_gate_up_weight, # raw, unused in baseline path
down_weight,
_gate_up_weight_scale, # raw, unused in baseline path
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"]
use_flydsl = os.environ.get("RK_USE_FLYDSL_MOE", "0") == "1"
if use_flydsl and (not _FLYDSL_DISABLED) and is_flydsl_available():
try:
return _flydsl_moe(
hidden_states=hidden_states,
gate_up_weight_shuffled=gate_up_weight_shuffled,
down_weight=down_weight,
gate_up_weight_scale_shuffled=gate_up_weight_scale_shuffled,
down_weight_scale=down_weight_scale,
topk_weights=topk_weights,
topk_ids=topk_ids,
)
except Exception:
# Disable FlyDSL fallback permanently in this process after first failure.
_FLYDSL_DISABLED = True
output = fused_moe_mod.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,
block_size_M=None,
moe_sorting_dispatch_policy=0,
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
)
return output
scrolls · 251 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