submission 524922
sanjay_arvind · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 142 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-524922?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:a9a4f3fb50aedd53745b135f8ff7dfe7023b2a8d179cc925244510b6aa113a58
license declaredunknown
license concludedunknown
authorssanjay_arvind
imported2026-08-15
Kernel source
submission.py142 lines
import os
import torch
from task import input_t, output_t
import aiter
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe, moe_sorting, get_inter_dim, get_block_size_M
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
_PRE = {}
_PATCHED = False
_USE_KSPLIT2 = {
(16, 512, 33): True,
}
try:
import reference as _ref
_orig_gen = _ref.generate_input
def _patched_gen(**kw):
data = _orig_gen(**kw)
config = data[11]
hidden_states = data[0]
topk_weights = data[9]
topk_ids = data[10]
w1 = data[5]
w2 = data[6]
M = config["bs"]
E = config["n_routed_experts"] + config["n_shared_experts"]
d_expert = config["d_expert"]
shape_key = (M, d_expert, E)
use_k2 = _USE_KSPLIT2.get(shape_key, False)
os.environ["AITER_KSPLIT"] = "2" if use_k2 else "0"
_, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)
if not use_k2:
block_m = get_block_size_M(M, config["total_top_k"], E, inter_dim)
if E > 64:
block_m = min(block_m, 32)
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = moe_sorting(
topk_ids, topk_weights, E, model_dim, dtypes.bf16, block_m,
)
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,
)
_PRE['sorted_ids'] = sorted_ids
_PRE['sorted_weights'] = sorted_weights
_PRE['sorted_expert_ids'] = sorted_expert_ids
_PRE['num_valid_ids'] = num_valid_ids
_PRE['moe_buf'] = moe_buf
_PRE['a1'] = a1
_PRE['a1_scale'] = a1_scale
_PRE['block_m'] = block_m
_PRE['M'] = M
_PRE['topk'] = config["total_top_k"]
_PRE['inter_dim'] = inter_dim
_PRE['use_nt'] = True
_PRE['a2'] = torch.empty((M, config["total_top_k"], inter_dim), dtype=dtypes.bf16, device=hidden_states.device)
_PRE['use_k2'] = use_k2
return data
_ref.generate_input = _patched_gen
import __main__
for _n in dir(__main__):
_o = getattr(__main__, _n, None)
if callable(_o) and hasattr(_o, '__globals__') and 'generate_input' in getattr(_o, '__globals__', {}):
_o.__globals__['generate_input'] = _patched_gen
_PATCHED = True
except Exception:
pass
def custom_kernel(data: input_t) -> output_t:
(
hidden_states, _, _, _, _,
w1, w2, w1_scale, w2_scale,
topk_weights, topk_ids, config,
) = data
if _PATCHED:
if _PRE.get('use_k2'):
return fused_moe(
hidden_states, w1, w2, topk_weights, topk_ids,
activation=ActivationType.Silu, quant_type=QuantType.per_1x32,
doweight_stage1=False, w1_scale=w1_scale, w2_scale=w2_scale,
hidden_pad=config["d_hidden_pad"] - config["d_hidden"],
intermediate_pad=config["d_expert_pad"] - config["d_expert"],
)
else:
return _fast_moe(w1, w2, w1_scale, w2_scale)
return fused_moe(
hidden_states, w1, w2, topk_weights, topk_ids,
activation=ActivationType.Silu, quant_type=QuantType.per_1x32,
doweight_stage1=False, w1_scale=w1_scale, w2_scale=w2_scale,
hidden_pad=config["d_hidden_pad"] - config["d_hidden"],
intermediate_pad=config["d_expert_pad"] - config["d_expert"],
)
def _fast_moe(w1, w2, w1_scale, w2_scale):
M = _PRE['M']
topk = _PRE['topk']
inter_dim = _PRE['inter_dim']
block_m = _PRE['block_m']
a2 = _PRE['a2']
aiter.ck_moe_stage1_fwd(
_PRE['a1'], w1, w2,
_PRE['sorted_ids'], _PRE['sorted_expert_ids'], _PRE['num_valid_ids'],
a2, topk, "",
w1_scale.view(dtypes.fp8_e8m0), _PRE['a1_scale'],
block_m, None,
QuantType.per_1x32, ActivationType.Silu,
0, _PRE['use_nt'],
)
a2_flat = a2.view(-1, inter_dim)
a2_q, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
a2_flat, sorted_ids=_PRE['sorted_ids'], num_valid_ids=_PRE['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,
_PRE['sorted_ids'], _PRE['sorted_expert_ids'], _PRE['num_valid_ids'],
_PRE['moe_buf'], topk, "",
w2_scale.view(dtypes.fp8_e8m0), a2_scale,
block_m, _PRE['sorted_weights'],
QuantType.per_1x32, ActivationType.Silu,
_PRE['use_nt'],
)
return _PRE['moe_buf']
scrolls · 142 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