submission 588558
uoedinburgh8089 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 245 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-588558?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:265af3d5395193cbf7514d6dc46190794c36539d3f28312a4cbad3889f4d1f4a
license declaredunknown
license concludedunknown
authorsuoedinburgh8089
imported2026-08-15
Kernel source
submission.py245 lines
import os
import tempfile
import torch
_fp4x2_str = "torch.float4_e2m1fn_x2"
_bf16_str = "torch.bfloat16"
_HDR = "cu_num,token,model_dim,inter_dim,expert,topk,act_type,dtype,q_dtype_a,q_dtype_w,q_type,use_g1u1,doweight_stage1,block_m,ksplit,us1,kernelName1,err1,us2,kernelName2,err2,us,run_1stage,tflops,bw,_tag"
_COMMON = f"ActivationType.Silu,{_bf16_str},{_fp4x2_str},{_fp4x2_str},QuantType.per_1x32,1,0"
_K = "x"
_KN1_SM = "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_KN2_SM = "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
_ROWS = [
f"256,16,7168,256,257,9,{_COMMON},16,2,0,{_K},0,0,{_K},0,0,0,0,",
f"256,128,7168,256,257,9,{_COMMON},16,2,0,{_K},0,0,{_K},0,0,0,0,",
f"256,512,7168,256,257,9,{_COMMON},32,0,0,{_KN1_SM},0,0,{_KN2_SM},0,0,0,0,",
f"256,16,7168,512,33,9,{_COMMON},16,2,0,{_K},0,0,{_K},0,0,0,0,",
f"256,128,7168,512,33,9,{_COMMON},16,2,0,{_K},0,0,{_K},0,0,0,0,",
]
_csv_content = _HDR + "\n" + "\n".join(_ROWS) + "\n"
_cfg_file = tempfile.NamedTemporaryFile(mode="w", suffix=".csv", delete=False, prefix="fmoe_")
_cfg_file.write(_csv_content)
_cfg_file.close()
os.environ["AITER_CONFIG_FMOE"] = _cfg_file.name
import aiter
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import get_inter_dim, get_padded_M, get_2stage_cfgs
from aiter.ops.triton.quant.fused_mxfp4_quant import (
fused_dynamic_mxfp4_quant_moe_sort,
)
from task import input_t, output_t
_cache = {}
def _init_config(M, topk, E, model_dim, inter_dim, w1, w2, hidden_pad, intermediate_pad, device):
key = (M, E, inter_dim)
if key in _cache:
return _cache[key]
metadata = get_2stage_cfgs(
get_padded_M(M), model_dim, inter_dim, E, topk,
torch.bfloat16, dtypes.fp4x2, dtypes.fp4x2,
QuantType.per_1x32, True, ActivationType.Silu,
False, hidden_pad, intermediate_pad, True,
)
block_m = int(metadata.block_m)
ksplit = int(metadata.ksplit)
s1_n_pad = intermediate_pad // 64 * 64 * 2
s1_k_pad = hidden_pad // 128 * 128
s2_n_pad = hidden_pad // 64 * 64
s2_k_pad = intermediate_pad // 128 * 128
_, n1, k1 = w1.shape
_, k2, n2 = w2.shape
D = n2 if k2 == k1 else n2 * 2
if w1.dtype is torch.uint32:
D = D * 8
kn1 = ""
kn2 = ""
use_nt = False
if hasattr(metadata.stage1, "keywords"):
kn1 = metadata.stage1.keywords.get("kernelName", "")
use_nt = metadata.stage1.keywords.get("use_non_temporal_load", False)
if metadata.stage2 is not None and hasattr(metadata.stage2, "keywords"):
kn2 = metadata.stage2.keywords.get("kernelName", "")
max_num_tokens_padded = M * topk + E * block_m - topk
max_num_m_blocks = (max_num_tokens_padded + block_m - 1) // block_m
sorted_ids = torch.empty(max_num_tokens_padded, dtype=torch.int32, device=device)
sorted_weights = torch.empty(max_num_tokens_padded, dtype=torch.float32, device=device)
sorted_expert_ids = torch.empty(max_num_m_blocks, dtype=torch.int32, device=device)
num_valid_ids = torch.empty(2, dtype=torch.int32, device=device)
moe_buf = torch.empty((M, model_dim), dtype=torch.bfloat16, device=device)
if ksplit >= 2:
tmp_out = torch.empty((M, topk, w1.shape[1]), dtype=torch.bfloat16, device=device)
a2_out = torch.empty((M, topk, D), dtype=torch.bfloat16, device=device)
else:
a2_out = torch.empty((M, topk, inter_dim), dtype=torch.bfloat16, device=device)
tmp_out = None
cfg = {
"ksplit": ksplit,
"block_m": block_m,
"s1_n_pad": s1_n_pad,
"s1_k_pad": s1_k_pad,
"s2_n_pad": s2_n_pad,
"s2_k_pad": s2_k_pad,
"kn1": kn1,
"kn2": kn2,
"use_nt": use_nt,
"sorted_ids": sorted_ids,
"sorted_weights": sorted_weights,
"sorted_expert_ids": sorted_expert_ids,
"num_valid_ids": num_valid_ids,
"moe_buf": moe_buf,
"tmp_out": tmp_out,
"a2_out": a2_out,
"inter_dim": inter_dim,
}
_cache[key] = cfg
return cfg
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
M = hidden_states.shape[0]
topk = topk_ids.shape[1]
device = hidden_states.device
hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
intermediate_pad = config["d_expert_pad"] - config["d_expert"]
w1 = gate_up_weight_shuffled
w2 = down_weight_shuffled
w1_scale = gate_up_weight_scale_shuffled
w2_scale = down_weight_scale_shuffled
w1.is_shuffled = True
w2.is_shuffled = True
E, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)
cfg = _init_config(
M, topk, E, model_dim, inter_dim,
w1, w2, hidden_pad, intermediate_pad, device,
)
block_m = cfg["block_m"]
ksplit = cfg["ksplit"]
sorted_ids = cfg["sorted_ids"]
sorted_weights_buf = cfg["sorted_weights"]
sorted_expert_ids = cfg["sorted_expert_ids"]
num_valid_ids = cfg["num_valid_ids"]
moe_buf = cfg["moe_buf"]
aiter.moe_sorting_fwd(
topk_ids, topk_weights,
sorted_ids, sorted_weights_buf, sorted_expert_ids, num_valid_ids,
moe_buf, E, block_m, None, None, 0,
)
w1_scale_e8m0 = w1_scale.view(dtypes.fp8_e8m0)
w2_scale_e8m0 = w2_scale.view(dtypes.fp8_e8m0)
if ksplit >= 2:
tmp_out = cfg["tmp_out"]
a2_out = cfg["a2_out"]
tmp_out.zero_()
aiter.moe_cktile2stages_gemm1(
hidden_states, w1, tmp_out,
sorted_ids, sorted_expert_ids, num_valid_ids, topk,
cfg["s1_n_pad"], cfg["s1_k_pad"],
None, None, w1_scale_e8m0, None,
ActivationType.Silu, block_m, ksplit,
)
aiter.silu_and_mul(a2_out, tmp_out)
aiter.moe_cktile2stages_gemm2(
a2_out, w2, moe_buf,
sorted_ids, sorted_expert_ids, num_valid_ids, topk,
cfg["s2_n_pad"], cfg["s2_k_pad"],
sorted_weights_buf, None, w2_scale_e8m0, None,
ActivationType.Silu, block_m,
)
return moe_buf
else:
a2_out = cfg["a2_out"]
a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
hidden_states,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=M,
topk=1,
block_size=block_m,
)
aiter.ck_moe_stage1_fwd(
a1, w1, w2,
sorted_ids, sorted_expert_ids, num_valid_ids,
a2_out, topk,
cfg["kn1"],
w1_scale_e8m0, a1_scale,
block_m,
None,
QuantType.per_1x32,
ActivationType.Silu,
0,
cfg["use_nt"],
torch.bfloat16,
)
a2_q, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
a2_out.view(-1, inter_dim),
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=M,
topk=topk,
block_size=block_m,
)
a2_q = a2_q.view(M, topk, -1)
aiter.ck_moe_stage2_fwd(
a2_q, w1, w2,
sorted_ids, sorted_expert_ids, num_valid_ids,
moe_buf, topk,
cfg["kn2"],
w2_scale_e8m0, a2_scale,
block_m,
sorted_weights_buf,
QuantType.per_1x32,
ActivationType.Silu,
cfg["use_nt"],
)
return moe_buf
scrolls · 245 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