submission 690294
chen5714288089 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 138 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-690294?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:5e03791bc9b60e7c15a40aa7c3efa39d21bc2be089eec65108e26f0fbf3cf0f9
license declaredunknown
license concludedunknown
authorschen5714288089
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
"""MoE MXFP4 — Fully inlined CK pipeline, all buffers cached.Kernel source
submission.py138 lines
"""MoE MXFP4 — Fully inlined CK pipeline, all buffers cached.
Bypasses fused_moe + fused_moe_ + fused_moe_2stages.
Directly calls: sorting → fused_quant → meta.stage1 → fused_quant → meta.stage2
All buffers (sorted, a2, moe_buf) cached and reused between iterations.
"""
import torch
import aiter
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import (
get_2stage_cfgs,
get_padded_M,
get_inter_dim,
_moe_sorting_impl,
_USE_OPUS_MOE_SORTING,
)
from aiter.utility import fp4_utils
try:
from aiter.ops.triton.quant.fused_mxfp4_quant import (
fused_dynamic_mxfp4_quant_moe_sort as _fq,
)
except ImportError:
_fq = None
_meta_cache = {}
_sort_key = None
_sort_val = None
_a2_buf = None
_quant_fn = None
def custom_kernel(data):
(
hidden_states, _, _, _, _,
w1, w2, w1s, w2s,
topk_weights, topk_ids, config,
) = data
global _sort_key, _sort_val, _a2_buf, _quant_fn
M = hidden_states.shape[0]
topk = topk_ids.shape[1]
dh = int(config["d_hidden"])
hp = int(config["d_hidden_pad"]) - dh
ip = int(config["d_expert_pad"]) - int(config["d_expert"])
dtype = hidden_states.dtype
device = hidden_states.device
# ── Derive dims from weight shapes (same as AITER's get_inter_dim) ──
E, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)
# ── Get tuned metadata (cached) ──
mkey = (M, E, topk, model_dim, inter_dim, hp, ip)
if mkey not in _meta_cache:
_meta_cache[mkey] = get_2stage_cfgs(
get_padded_M(M), model_dim, inter_dim,
E, topk, dtype,
q_dtype_a=dtypes.fp4x2, q_dtype_w=dtypes.fp4x2,
q_type=QuantType.per_1x32,
use_g1u1=True, activation=ActivationType.Silu,
doweight_stage1=False,
hidden_pad=hp, intermediate_pad=ip, is_shuffled=True,
)
meta = _meta_cache[mkey]
block_m = meta.block_m
# ── Sorting (cached when same input — benchmark mode reuses data) ──
skey = (topk_ids.data_ptr(), topk_weights.data_ptr(), M, E)
if _sort_key == skey:
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = _sort_val
else:
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = \
_moe_sorting_impl(topk_ids, topk_weights, E, dh, dtype, block_m,
None, None, 0, _USE_OPUS_MOE_SORTING)
_sort_key = skey
_sort_val = (sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf)
# ── Stage 0: Quantize activations → MXFP4 ──
if _fq is not None and M <= 1024:
a1, a1_scale = _fq(hidden_states, sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=M, topk=1, block_size=block_m)
else:
if _quant_fn is None:
_quant_fn = aiter.get_triton_quant(QuantType.per_1x32)
a1, a1_scale = _quant_fn(hidden_states, quant_dtype=dtypes.fp4x2)
a1_scale = fp4_utils.moe_mxfp4_sort(
a1_scale, sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=M, block_size=block_m)
# ── Reuse a2 buffer ──
a2_shape = (M, topk, inter_dim)
if _a2_buf is None or _a2_buf.shape != a2_shape:
_a2_buf = torch.empty(a2_shape, dtype=dtype, device=device)
# ── Stage 1: CK GEMM (gate_up + SwiGLU) ──
a2 = meta.stage1(
a1, w1, w2,
sorted_ids, sorted_expert_ids, num_valid_ids,
_a2_buf, topk,
block_m=block_m,
a1_scale=a1_scale,
w1_scale=w1s.view(dtypes.fp8_e8m0),
sorted_weights=None,
)
# ── Stage 0.5: Quantize intermediate → MXFP4 ──
a2_flat = a2.view(-1, inter_dim)
if _fq is not None and M <= 1024:
a2_q, a2_scale = _fq(a2_flat, 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)
else:
if _quant_fn is None:
_quant_fn = aiter.get_triton_quant(QuantType.per_1x32)
a2_q, a2_scale = _quant_fn(a2_flat, quant_dtype=dtypes.fp4x2)
a2_scale = fp4_utils.moe_mxfp4_sort(
a2_scale[:M * topk, :].view(M, topk, -1),
sorted_ids=sorted_ids, num_valid_ids=num_valid_ids,
token_num=M, block_size=block_m)
a2_q = a2_q.view(M, topk, -1)
# ── Stage 2: CK GEMM (down + weighted reduce) ──
meta.stage2(
a2_q, w1, w2,
sorted_ids, sorted_expert_ids, num_valid_ids,
moe_buf, topk,
w2_scale=w2s.view(dtypes.fp8_e8m0),
a2_scale=a2_scale,
block_m=block_m,
sorted_weights=sorted_weights,
)
return moe_buf
scrolls · 138 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