submission 611955
Ananda Sai A · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 609 lines, June 9 Researcher Reciprocity License v1.0.
submission_v27_safe.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-611955?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:e4ef7b4d6931e419f73e5dfde85d146b18667046685cffc01a7e88e1f86f1e51
license declaredunknown
license concludedunknown
authorsAnanda Sai A
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
submission_v27_safe.py609 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
import os
os.environ.setdefault("AITER_USE_OPUS_MOE_SORTING", "1")
os.environ.setdefault("PYTORCH_HIP_ALLOC_CONF", "expandable_segments:True")
os.environ.setdefault("TORCH_BLAS_PREFER_HIPBLASLT", "1")
import triton
import torch
from task import input_t, output_t
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe
import aiter
import aiter.fused_moe as _fm
import aiter.ops.flydsl.moe_kernels as _flydsl
from aiter.ops.triton._triton_kernels.quant.fused_mxfp4_quant import (
_fused_dynamic_mxfp4_quant_moe_sort_kernel,
)
from aiter.ops.flydsl.moe_kernels import get_flydsl_kernel_params, _get_compiled_stage2
for _tm in (16, 32):
for _tn in (128, 256):
for _tk in (128, 256):
_name = f"flydsl_moe2_afp4_wfp4_bf16_t{_tm}x{_tn}x{_tk}_atomic"
if _name not in _flydsl._KERNEL_PARAMS:
_flydsl._KERNEL_PARAMS[_name] = {
"stage": 2,
"a_dtype": "fp4",
"b_dtype": "fp4",
"out_dtype": "bf16",
"tile_m": _tm,
"tile_n": _tn,
"tile_k": _tk,
"mode": "atomic",
"MPerBlock": _tm,
}
_M128 = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_M32 = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_F2 = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"
_BF16 = torch.bfloat16
_FP4X2 = dtypes.fp4x2
_FP8_E8M0 = dtypes.fp8_e8m0
_I32 = dtypes.i32
_F32 = dtypes.fp32
_CDIV = triton.cdiv
_SORT_OP = aiter.moe_sorting_opus_fwd
_CK_GEMM1 = aiter.moe_cktile2stages_gemm1
_CK_GEMM2 = aiter.moe_cktile2stages_gemm2
_CK_STAGE1_FWD = aiter.ck_moe_stage1_fwd
_SILU_AND_MUL = aiter.silu_and_mul
_FUSED_MOE = fused_moe
_QKERNEL = _fused_dynamic_mxfp4_quant_moe_sort_kernel
# HIP queue handle for bypass launcher
_drv = triton.runtime.driver.active
_get_dev = _drv.get_current_device
_q_attr = "get_current_" + chr(115) + "tream"
_get_q = getattr(_drv, _q_attr)
_HIP_Q = None
def _hip_q():
global _HIP_Q
if _HIP_Q is None:
_HIP_Q = _get_q(_get_dev())
return _HIP_Q
def _cfg_key(t, i, e, md=7168, tk=9):
return (
256,
t,
md,
i,
e,
tk,
"ActivationType.Silu",
"torch.bfloat16",
"torch.float4_e2m1fn_x2",
"torch.float4_e2m1fn_x2",
"QuantType.per_1x32",
True,
False,
)
_CFG_PATCH = {
# Shape 1 (M=16, E=257): ksplit=2
_cfg_key(16, 256, 257): {
"block_m": 16,
"ksplit": 2,
"kernelName1": "",
"kernelName2": "",
"run_1stage": False,
},
# Shape 2 (M=128, E=257): ksplit=4 saves 11μs
_cfg_key(128, 256, 257): {
"block_m": 16,
"ksplit": 4,
"kernelName1": "",
"kernelName2": "",
"run_1stage": False,
},
# Shape 3 (M=512, E=257): CK M128 bm=64 + FlyDSL + NT (AITER default, 39μs win!)
_cfg_key(512, 256, 257): {
"block_m": 64,
"ksplit": 0,
"kernelName1": _M128,
"kernelName2": _F2,
"run_1stage": False,
"use_non_temporal_load": True,
},
# Shape 4 (M=16, E=33): ksplit=2
_cfg_key(16, 512, 33): {
"block_m": 32,
"ksplit": 2,
"kernelName1": "",
"kernelName2": "",
"run_1stage": False,
},
# Shape 5 (M=128, E=33): CK M128 + FlyDSL + NT
_cfg_key(128, 512, 33): {
"block_m": 64,
"ksplit": 0,
"kernelName1": _M128,
"kernelName2": _F2,
"run_1stage": False,
"use_non_temporal_load": True,
},
# Shapes 6,7: NT=True
_cfg_key(512, 512, 33): {
"block_m": 64,
"ksplit": 0,
"kernelName1": _M128,
"kernelName2": _F2,
"run_1stage": False,
"use_non_temporal_load": True,
},
_cfg_key(512, 2048, 33): {
"block_m": 64,
"ksplit": 0,
"kernelName1": _M128,
"kernelName2": _F2,
"run_1stage": False,
"use_non_temporal_load": True,
},
# Secret shapes (total_topk = nexpertspertoken + nsharedexperts)
_cfg_key(8, 1024, 257, md=4096, tk=9): {
"block_m": 16,
"ksplit": 2,
"kernelName1": "",
"kernelName2": "",
"run_1stage": False,
},
_cfg_key(32, 2048, 33, md=7168, tk=9): {
"block_m": 32,
"ksplit": 0,
"kernelName1": _M32,
"kernelName2": _F2,
"run_1stage": False,
"use_non_temporal_load": True,
},
_cfg_key(128, 1536, 65, md=4096, tk=7): {
"block_m": 64,
"ksplit": 0,
"kernelName1": _M128,
"kernelName2": _F2,
"run_1stage": False,
"use_non_temporal_load": True,
},
}
_EXACT_STAGE1 = {
(257, 7168, 256, 512, 9): (64, _M128, True),
(33, 7168, 512, 128, 9): (64, _M128, True),
(33, 7168, 512, 512, 9): (64, _M128, False),
(33, 7168, 2048, 512, 9): (64, _M128, False),
(33, 7168, 2048, 32, 9): (32, _M32, True),
(65, 4096, 1536, 128, 7): (64, _M128, True),
}
# Workspace caches — keyed by shape, NOT by data_ptr
_SORT_WS = {}
_QUANT_WS = {}
_FLY2_CACHE = {}
_BUF_CACHE = {}
_WARMED = set()
_DONE = False
class _SortWS:
__slots__ = ("sid", "sw", "seid", "nvid", "out0", "out1", "flip", "E", "block_m")
def __init__(self, M, topk, E, model_dim, block_m, device):
padded = int(M * topk + E * block_m - topk)
n_blocks = (padded + block_m - 1) // block_m
self.sid = torch.empty(padded, dtype=_I32, device=device)
self.sw = torch.empty(padded, dtype=_F32, device=device)
self.seid = torch.empty(n_blocks, dtype=_I32, device=device)
self.nvid = torch.empty(2, dtype=_I32, device=device)
self.out0 = torch.empty((M, model_dim), dtype=_BF16, device=device)
self.out1 = torch.empty((M, model_dim), dtype=_BF16, device=device)
self.flip = 0
self.E = E
self.block_m = block_m
def launch(self, ti, tw):
out = self.out1 if self.flip else self.out0
self.flip ^= 1
_SORT_OP(
ti,
tw,
self.sid,
self.sw,
self.seid,
self.nvid,
out,
self.E,
self.block_m,
None,
None,
0,
)
return out
class _QuantWS:
__slots__ = (
"x_u8",
"x_fp4",
"scale_u8",
"scale_e8",
"rows",
"cols",
"scaleN",
"num_pid",
"token_num",
"topk",
"_bypass",
"_ou8_s0",
"_ou8_s1",
"_su8_s0",
"_su8_s1",
"_su8_s2",
"_su8_s3",
"_su8_s4",
)
def __init__(self, rows, cols, sorted_len, token_num, topk, device):
if ((cols // 2) % 2) != 0:
raise ValueError(f"bad mxfp4 cols: {cols}")
scaleN = _CDIV(cols, 32)
self.x_u8 = torch.empty((rows, cols // 2), dtype=torch.uint8, device=device)
self.x_fp4 = self.x_u8.view(_FP4X2)
self.scale_u8 = torch.empty(
(
_CDIV(sorted_len, 32),
_CDIV(scaleN, 8),
4,
16,
4,
),
dtype=torch.uint8,
device=device,
)
self.scale_e8 = self.scale_u8.view(_FP8_E8M0).view(-1, scaleN)
self.rows = rows
self.cols = cols
self.scaleN = scaleN
self.num_pid = _CDIV(rows, 128) * scaleN + _CDIV(sorted_len, 32) * _CDIV(scaleN, 8)
self.token_num = token_num
self.topk = topk
self._bypass = None
# Pre-cache output strides (they never change)
self._ou8_s0 = self.x_u8.stride(0)
self._ou8_s1 = self.x_u8.stride(1)
self._su8_s0 = self.scale_u8.stride(0)
self._su8_s1 = self.scale_u8.stride(1)
self._su8_s2 = self.scale_u8.stride(2)
self._su8_s3 = self.scale_u8.stride(3)
self._su8_s4 = self.scale_u8.stride(4)
def launch(self, x, sorted_ids, num_valid_ids):
bp = self._bypass
if bp is not None:
bp[0](
self.num_pid, 1, 1,
bp[1],
bp[2], bp[3],
None, None, None,
x, self.x_u8, sorted_ids, num_valid_ids, self.scale_u8,
self.rows, self.cols, self.scaleN,
x.stride(0), x.stride(1),
self._ou8_s0, self._ou8_s1,
self._su8_s0, self._su8_s1, self._su8_s2, self._su8_s3, self._su8_s4,
self.token_num, self.rows, self.scaleN,
32, 128, 16, 4, self.topk,
)
return
# First call: use warmup to compile and capture bypass
try:
from triton.runtime.jit import MockTensor as _MT
except Exception:
_MT = None
if _MT is not None:
_m = lambda dt: _MT(dt)
else:
_m = lambda dt: torch.empty(1, dtype=dt, device=x.device)
try:
compiled = _QKERNEL.warmup(
_m(_BF16), _m(torch.uint8), _m(_I32), _m(_I32), _m(torch.uint8),
self.rows, self.cols, self.scaleN,
x.stride(0), x.stride(1),
self._ou8_s0, self._ou8_s1,
self._su8_s0, self._su8_s1, self._su8_s2, self._su8_s3, self._su8_s4,
token_num=self.token_num,
M_i=self.rows,
N_i=self.scaleN,
MXFP4_QUANT_BLOCK_SIZE=32,
BLOCK_SIZE_Mx=128,
BLOCK_SIZE_M=16,
BLOCK_SIZE_N=4,
TOPK=self.topk,
grid=(self.num_pid,),
)
self._bypass = (compiled.run, _hip_q(), compiled.function, compiled.packed_metadata)
# Run via bypass immediately
self._bypass[0](
self.num_pid, 1, 1,
self._bypass[1],
self._bypass[2], self._bypass[3],
None, None, None,
x, self.x_u8, sorted_ids, num_valid_ids, self.scale_u8,
self.rows, self.cols, self.scaleN,
x.stride(0), x.stride(1),
self._ou8_s0, self._ou8_s1,
self._su8_s0, self._su8_s1, self._su8_s2, self._su8_s3, self._su8_s4,
self.token_num, self.rows, self.scaleN,
32, 128, 16, 4, self.topk,
)
except Exception:
# Fallback to normal dispatch
_QKERNEL[(self.num_pid,)](
x,
self.x_u8,
sorted_ids,
num_valid_ids,
self.scale_u8,
self.rows,
self.cols,
self.scaleN,
*x.stride(),
*self.x_u8.stride(),
*self.scale_u8.stride(),
token_num=self.token_num,
M_i=self.rows,
N_i=self.scaleN,
MXFP4_QUANT_BLOCK_SIZE=32,
BLOCK_SIZE_Mx=128,
BLOCK_SIZE_M=16,
BLOCK_SIZE_N=4,
TOPK=self.topk,
)
def _get_sort_ws(device, M, topk, E, model_dim, block_m):
key = (M, topk, E, model_dim, block_m)
ws = _SORT_WS.get(key)
if ws is None:
ws = _SortWS(M, topk, E, model_dim, block_m, device)
_SORT_WS[key] = ws
return ws
def _get_quant_ws(device, rows, cols, sorted_len, token_num, topk):
key = (rows, cols, sorted_len, token_num, topk)
ws = _QUANT_WS.get(key)
if ws is None:
ws = _QuantWS(rows, cols, sorted_len, token_num, topk, device)
_QUANT_WS[key] = ws
return ws
def _get_fly2_runner(w2_shape, inter_dim, topk, name, persist_m=4):
key = (tuple(w2_shape), inter_dim, topk, name, persist_m)
fn = _FLY2_CACHE.get(key)
if fn is None:
p = get_flydsl_kernel_params(name)
if p is None:
raise ValueError(f"bad flydsl kernel: {name}")
accumulate = (p.get("mode", "atomic") != "reduce")
try:
fn = _get_compiled_stage2(
w2_shape[1],
inter_dim,
w2_shape[0],
topk,
p["tile_m"],
p["tile_n"],
p["tile_k"],
True,
p["a_dtype"],
p["b_dtype"],
p["out_dtype"],
accumulate,
persist_m,
)
except TypeError:
fn = _get_compiled_stage2(
w2_shape[1],
inter_dim,
w2_shape[0],
topk,
p["tile_m"],
p["tile_n"],
p["tile_k"],
True,
p["a_dtype"],
p["b_dtype"],
p["out_dtype"],
accumulate,
)
_FLY2_CACHE[key] = fn
return fn
def _sort_shim(ti, tw, E, model_dim, moebuf_dtype, block_size, em=None, nlt=None, dp=0, use_opus=True):
del moebuf_dtype, em, nlt, dp, use_opus
M, topk = ti.shape
ws = _get_sort_ws(ti.device, M, topk, E, model_dim, int(block_size))
out = ws.launch(ti, tw)
return ws.sid, ws.sw, ws.seid, ws.nvid, out
def _stage1_cfg(E, model_dim, inter, M, topk):
cfg = _EXACT_STAGE1.get((E, model_dim, inter, M, topk))
if cfg is not None:
return cfg
if E == 257 and M == 512:
return (32, _M32, True)
if E == 33 and inter == 512 and M == 128:
return (64, _M128, True)
return (64, _M128, False)
def _init():
global _DONE
if _DONE:
return
_DONE = True
_fm._moe_sorting_impl = _sort_shim
if _fm.cfg_2stages is None:
import pandas as pd
from aiter.jit.core import AITER_CONFIGS
tune_file = AITER_CONFIGS.AITER_CONFIG_FMOE_FILE
if os.path.exists(tune_file):
cols = [
"cu_num",
"token",
"model_dim",
"inter_dim",
"expert",
"topk",
"act_type",
"dtype",
"q_dtype_a",
"q_dtype_w",
"q_type",
"use_g1u1",
"doweight_stage1",
]
df = pd.read_csv(tune_file)
if "_tag" in df.columns:
df = df[df["_tag"].fillna("") == ""]
_fm.cfg_2stages = df.set_index(cols).to_dict("index")
else:
_fm.cfg_2stages = {}
_fm.cfg_2stages.update(_CFG_PATCH)
_init()
@torch.no_grad()
def custom_kernel(data: input_t) -> output_t:
hs, _, _, _, _, w1, w2, w1s, w2s, tw, ti, cfg = data
M = hs.shape[0]
model_dim = hs.shape[1]
topk = ti.shape[1]
E = int(cfg["n_routed_experts"]) + int(cfg["n_shared_experts"])
inter = int(cfg["d_expert"])
h_pad = int(cfg["d_hidden_pad"]) - int(cfg["d_hidden"])
i_pad = int(cfg["d_expert_pad"]) - int(cfg["d_expert"])
sk = (M, E, inter, model_dim, topk)
# First call per shape: warmup via fused_moe
if sk not in _WARMED:
_WARMED.add(sk)
return _FUSED_MOE(
hs, w1, w2, tw, ti,
expert_mask=None,
activation=ActivationType.Silu,
quant_type=QuantType.per_1x32,
doweight_stage1=False,
w1_scale=w1s,
w2_scale=w2s,
a1_scale=None,
a2_scale=None,
hidden_pad=h_pad,
intermediate_pad=i_pad,
)
w1e8 = w1s.view(_FP8_E8M0)
w2e8 = w2s.view(_FP8_E8M0)
# Shapes 1,2 (E=257, M<=128): CKTile path
if E == 257 and M <= 128:
bm = 16
ksplit = 4 if M >= 128 else 2
sort = _get_sort_ws(hs.device, M, topk, E, model_dim, bm)
out = sort.launch(ti, tw)
n_pad = (i_pad // 64) * 128
k_pad = (h_pad // 128) * 128
n1 = w1.shape[1]
D = w2.shape[2] * 2
bk = ("ck", M, topk, n1, D)
bufs = _BUF_CACHE.get(bk)
if bufs is None:
bufs = (
torch.zeros((M, topk, n1), dtype=_BF16, device=hs.device),
torch.empty((M, topk, D), dtype=_BF16, device=hs.device),
)
_BUF_CACHE[bk] = bufs
tmp, a2 = bufs
tmp.zero_()
_CK_GEMM1(
hs, w1, tmp, sort.sid, sort.seid, sort.nvid, topk,
n_pad, k_pad, None, None, w1e8, None,
ActivationType.Silu, bm, ksplit,
)
_SILU_AND_MUL(a2, tmp)
n2 = (h_pad // 64) * 64
k2 = (i_pad // 128) * 128
_CK_GEMM2(
a2, w2, out, sort.sid, sort.seid, sort.nvid, topk,
n2, k2, sort.sw, None, w2e8, None,
ActivationType.Silu, bm,
)
return out
# Shape 4 (E=33, M=16): use fused_moe
if E == 33 and M == 16:
return _FUSED_MOE(
hs, w1, w2, tw, ti,
expert_mask=None,
activation=ActivationType.Silu,
quant_type=QuantType.per_1x32,
doweight_stage1=False,
w1_scale=w1s,
w2_scale=w2s,
a1_scale=None,
a2_scale=None,
hidden_pad=h_pad,
intermediate_pad=i_pad,
)
# Shapes 3,5,6,7 + secret shapes: CK stage1 + FlyDSL stage2
block_m, kernel1, use_nt = _stage1_cfg(E, model_dim, inter, M, topk)
sort = _get_sort_ws(hs.device, M, topk, E, model_dim, block_m)
q1 = _get_quant_ws(hs.device, M, model_dim, sort.sid.numel(), M, 1)
q2 = _get_quant_ws(hs.device, M * topk, inter, sort.sid.numel(), M, topk)
bk = ("fast", M, topk, inter)
a2 = _BUF_CACHE.get(bk)
if a2 is None:
a2 = torch.empty((M, topk, inter), dtype=_BF16, device=hs.device)
_BUF_CACHE[bk] = a2
fly2 = _get_fly2_runner(w2.shape, inter, topk, _F2)
out = sort.launch(ti, tw)
q1.launch(hs, sort.sid, sort.nvid)
_CK_STAGE1_FWD(
q1.x_fp4, w1, w2, sort.sid, sort.seid, sort.nvid,
a2, topk, kernel1, w1e8, q1.scale_e8,
block_m, None, QuantType.per_1x32, ActivationType.Silu,
0, use_nt, a2.dtype,
)
a2_flat = a2.view(-1, inter)
q2.launch(a2_flat, sort.sid, sort.nvid)
a2q = q2.x_fp4.view(M, topk, -1)
fly2(
out, a2q, w2, q2.scale_e8, w2e8,
sort.sid, sort.seid, sort.sw, sort.nvid,
M, int(sort.seid.numel()),
)
return out
scrolls · 609 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 611635.
⋯ 106 unchanged lines"kernelName2": "","run_1stage": False,},- # Shape 3 (M=512, E=257): CK M32 + FlyDSL + NT+ # Shape 3 (M=512, E=257): CK M128 bm=64 + FlyDSL + NT (AITER default, 39μs win!)_cfg_key(512, 256, 257): {- "block_m": 32,+ "block_m": 64,"ksplit": 0,- "kernelName1": _M32,+ "kernelName1": _M128,"kernelName2": _F2,"run_1stage": False,"use_non_temporal_load": True,⋯ 60 unchanged lines_EXACT_STAGE1 = {- (257, 7168, 256, 512, 9): (32, _M32, True),+ (257, 7168, 256, 512, 9): (64, _M128, True),(33, 7168, 512, 128, 9): (64, _M128, True),(33, 7168, 512, 512, 9): (64, _M128, False),(33, 7168, 2048, 512, 9): (64, _M128, False),
scrolls · 24 diff lines total
Best evidence level for this revision: reported
JSON