submission 754899
Itay Etelis · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 159 lines, June 9 Researcher Reciprocity License v1.0.
sub_v79_inline_ck2stage.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-754899?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:e998523329a859fa0c4e2f23dcd0e148ad65d10b54b913497f7c41193e244d90
license declaredunknown
license concludedunknown
authorsItay Etelis
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
activation=_S, split_k=plan["ksplit"])Kernel source
sub_v79_inline_ck2stage.py159 lines
"""
V79: Full inline CK 2-stage with pre-allocated a2 buffer.
V76 uses fused_moe_2stages() for CK path which allocates a2 internally.
This version inlines the CK 2-stage path, pre-allocating a2 and calling
stage1/stage2 directly via resolved metadata.
Saves: a2 allocation (~2-3us) + fused_moe_2stages dispatch (~1-2us)
"""
import torch
from task import input_t, output_t
import aiter
from aiter import ActivationType, QuantType, dtypes
import aiter.fused_moe as _fmoe
from aiter.fused_moe import (
fused_moe, get_inter_dim, get_2stage_cfgs, get_padded_M,
fused_dynamic_mxfp4_quant_moe_sort,
cktile_moe_stage1, cktile_moe_stage2,
)
_S = ActivationType.Silu
_Q = QuantType.per_1x32
_BF16 = torch.bfloat16
_CK1_32 = "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_CK1_32_4CU = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_CK1_64 = "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_CK2_d512 = "moe_ck2stages_gemm2_256x32x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
_FLY_32r = "flydsl_moe2_afp4_wfp4_bf16_t32x256x256_reduce"
_FLY_64r = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
_first_call = True
_plans = {}
_sort_bufs = {}
_a2_bufs = {}
def _mk(t, i, e):
return (256, t, 7168, i, e, 9, "ActivationType.Silu", "torch.bfloat16",
"torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",
"QuantType.per_1x32", True, False)
def _inject():
if _fmoe.cfg_2stages is None: return
ek2 = {"run_1stage": False, "ksplit": 2}
e = {"run_1stage": False, "ksplit": 0}
_fmoe.cfg_2stages.update({
_mk(16, 512, 33): {**ek2, "block_m": 32, "kernelName1": "", "kernelName2": ""},
_mk(128, 512, 33): {**ek2, "block_m": 32, "kernelName1": "", "kernelName2": ""},
_mk(512, 512, 33): {**e, "block_m": 32, "kernelName1": _CK1_32, "kernelName2": _FLY_32r},
_mk(512, 2048, 33): {**e, "block_m": 64, "kernelName1": _CK1_64, "kernelName2": _FLY_64r},
_mk(16, 256, 257): {**ek2, "block_m": 16, "kernelName1": "", "kernelName2": ""},
_mk(128, 256, 257): {**ek2, "block_m": 16, "kernelName1": "", "kernelName2": ""},
_mk(512, 256, 257): {**e, "block_m": 32, "kernelName1": _CK1_32_4CU, "kernelName2": _FLY_32r},
})
_fmoe.get_2stage_cfgs.cache_clear()
def _get_sort_bufs(M, topk, E, model_dim, block_m, device):
key = (M, topk, E, block_m)
if key not in _sort_bufs:
mp = int(M * topk + E * block_m - topk); mb = (mp + block_m - 1) // block_m
_sort_bufs[key] = (
torch.empty(mp, dtype=dtypes.i32, device=device),
torch.empty(mp, dtype=dtypes.fp32, device=device),
torch.empty(mb, dtype=dtypes.i32, device=device),
torch.empty(2, dtype=dtypes.i32, device=device),
torch.empty((M, model_dim), dtype=_BF16, device=device))
return _sort_bufs[key]
def _get_a2_buf(M, topk, inter_dim, device):
key = (M, topk, inter_dim)
if key not in _a2_bufs:
_a2_bufs[key] = torch.empty((M, topk, inter_dim), dtype=_BF16, device=device)
return _a2_bufs[key]
def _get_plan(M, topk, E, model_dim, inter_dim, hp, ip):
key = (M, E, inter_dim)
if key in _plans: return _plans[key]
pm = get_padded_M(M)
meta = get_2stage_cfgs(pm, model_dim, inter_dim, E, topk, dtypes.bf16,
dtypes.fp4x2, dtypes.fp4x2, _Q, True, _S, False, hp, ip, True)
_plans[key] = {"block_m": meta.block_m, "ksplit": meta.ksplit,
"use_cktile": meta.ksplit > 1,
"stage1": meta.stage1, "stage2": meta.stage2,
"hp": hp, "ip": ip}
return _plans[key]
def _run(hidden_states, w1, w2, w1_scale, w2_scale, topk_weights, topk_ids, config):
M = hidden_states.shape[0]; topk = topk_ids.shape[1]
E, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)
hp = config["d_hidden_pad"] - config["d_hidden"]; ip = config["d_expert_pad"] - config["d_expert"]
device = hidden_states.device
plan = _get_plan(M, topk, E, model_dim, inter_dim, hp, ip)
block_m = plan["block_m"]
si, sw, se, nv, moe_buf = _get_sort_bufs(M, topk, E, model_dim, block_m, device)
aiter.moe_sorting_fwd(topk_ids, topk_weights, si, sw, se, nv, moe_buf,
E, int(block_m), None, None, 0)
if plan["use_cktile"]:
n1 = ip // 64 * 64 * 2; k1 = hp // 128 * 128
n2 = hp // 64 * 64; k2 = ip // 128 * 128
a2 = cktile_moe_stage1(hidden_states, w1, w2, si, se, nv,
None, topk, block_m, None, w1_scale.view(dtypes.fp8_e8m0),
sorted_weights=None, n_pad_zeros=n1, k_pad_zeros=k1,
activation=_S, split_k=plan["ksplit"])
cktile_moe_stage2(a2, w1, w2, si, se, nv, moe_buf, topk,
w2_scale.view(dtypes.fp8_e8m0), None, block_m,
activation=_S, sorted_weights=sw, n_pad_zeros=n2, k_pad_zeros=k2)
else:
# INLINE CK 2-stage: quant → stage1 → re-quant → stage2
# Step 1: Quant hidden_states
a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
hidden_states, sorted_ids=si, num_valid_ids=nv,
token_num=M, topk=1, block_size=block_m)
# Step 2: Stage1 with pre-allocated a2
a2_buf = _get_a2_buf(M, topk, inter_dim, device)
a2 = plan["stage1"](
a1, w1, w2, si, se, nv, a2_buf, topk,
block_m=block_m, a1_scale=a1_scale,
w1_scale=w1_scale.view(dtypes.fp8_e8m0),
sorted_weights=None)
# Step 3: Re-quant intermediate
a2_flat = a2.view(-1, inter_dim)
a2_q, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
a2_flat, sorted_ids=si, num_valid_ids=nv,
token_num=M, topk=topk, block_size=block_m)
a2_q = a2_q.view(M, topk, -1)
# Step 4: Stage2
plan["stage2"](
a2_q, w1, w2, si, se, nv, moe_buf, topk,
w2_scale=w2_scale.view(dtypes.fp8_e8m0),
a2_scale=a2_scale,
block_m=block_m,
sorted_weights=sw)
return moe_buf
def custom_kernel(data: input_t) -> output_t:
global _first_call
(hidden_states, _, _, _, _, w1, w2, w1_scale, w2_scale,
topk_weights, topk_ids, config) = data
hp = config["d_hidden_pad"] - config["d_hidden"]; ip = config["d_expert_pad"] - config["d_expert"]
if _first_call:
_first_call = False
r = fused_moe(hidden_states, w1, w2, topk_weights, topk_ids,
activation=_S, quant_type=_Q, w1_scale=w1_scale, w2_scale=w2_scale,
hidden_pad=hp, intermediate_pad=ip)
_inject()
return r
try:
return _run(hidden_states, w1, w2, w1_scale, w2_scale, topk_weights, topk_ids, config)
except Exception:
return fused_moe(hidden_states, w1, w2, topk_weights, topk_ids,
activation=_S, quant_type=_Q, w1_scale=w1_scale, w2_scale=w2_scale,
hidden_pad=hp, intermediate_pad=ip)
scrolls · 159 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