submission 754384
flower2123 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 224 lines, June 9 Researcher Reciprocity License v1.0.
submission_v1_flower_moe.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-754384?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:3570996e2cc98ca626ca8eaa376f7e84afd6777e045624a844cebe09d9261513
license declaredunknown
license concludedunknown
authorsflower2123
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
{"stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16", "tile_m": 32, "tile_n": 128, "tile_k": 128, "mode": "atomic", "MPerBlock": 32}),tile-n = 32
BLK, BIG, BM, BN = 32, 128, 32, 8Kernel source
submission_v1_flower_moe.py224 lines
# Author: flower
import os
import functools
import torch
import triton
from typing import Dict, Tuple, Optional
from task import input_t, output_t
import aiter
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import (
get_2stage_cfgs, get_padded_M, get_inter_dim,
ck_moe_stage1, cktile_moe_stage1, cktile_moe_stage2,
_flydsl_stage2_wrapper,
)
import aiter.fused_moe as _fmoe
import aiter.ops.flydsl.moe_kernels as _fkern
from aiter.ops.triton._triton_kernels.quant.fused_mxfp4_quant import (
_fused_dynamic_mxfp4_quant_moe_sort_kernel,
)
from aiter.utility import fp4_utils
for _nm, _cfg in [
("flydsl_moe2_afp4_wfp4_bf16_t32x128x128_atomic",
{"stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16", "tile_m": 32, "tile_n": 128, "tile_k": 128, "mode": "atomic", "MPerBlock": 32}),
("flydsl_moe2_afp4_wfp4_bf16_t32x256x128_atomic",
{"stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16", "tile_m": 32, "tile_n": 256, "tile_k": 128, "mode": "atomic", "MPerBlock": 32}),
("flydsl_moe2_afp4_wfp4_bf16_t16x256x128_atomic",
{"stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16", "tile_m": 16, "tile_n": 256, "tile_k": 128, "mode": "atomic", "MPerBlock": 16}),
("flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic",
{"stage": 2, "a_dtype": "fp4", "b_dtype": "fp4", "out_dtype": "bf16", "tile_m": 16, "tile_n": 128, "tile_k": 128, "mode": "atomic", "MPerBlock": 16}),
]:
_fkern._KERNEL_PARAMS[_nm] = _cfg
_CK_S1_128 = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_CK_S1_32 = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_FLY_16_128 = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"
_FLY_16_256 = "flydsl_moe2_afp4_wfp4_bf16_t16x256x128_atomic"
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)
_SHAPE_MAP = {
_mk(16, 512, 33): dict(block_m=32, ksplit=2, kernelName1="", kernelName2="", run_1stage=False),
_mk(128, 512, 33): dict(block_m=64, ksplit=0, kernelName1=_CK_S1_128, kernelName2=_FLY_16_128, run_1stage=False),
_mk(512, 512, 33): dict(block_m=64, ksplit=0, kernelName1=_CK_S1_128, kernelName2=_FLY_16_128, run_1stage=False),
_mk(512, 2048, 33): dict(block_m=64, ksplit=0, kernelName1=_CK_S1_128, kernelName2=_FLY_16_256, run_1stage=False),
_mk(16, 256, 257): dict(block_m=16, ksplit=2, kernelName1="", kernelName2="", run_1stage=False),
_mk(128, 256, 257): dict(block_m=16, ksplit=2, kernelName1="", kernelName2="", run_1stage=False),
_mk(512, 256, 257): dict(block_m=32, ksplit=0, kernelName1=_CK_S1_32, kernelName2=_FLY_16_128, run_1stage=False, use_non_temporal_load=True),
}
_mem = {}
def _sorting_tensors(n_tok, n_exp, k, dim, bm, dev):
tag = ("s", n_tok, n_exp, k, dim, bm)
if tag in _mem:
return _mem[tag]
cap = int(n_tok * k + n_exp * bm - k)
nb = int((cap + bm - 1) // bm)
d = {
"si": torch.empty(cap, dtype=dtypes.i32, device=dev),
"sw": torch.empty(cap, dtype=dtypes.fp32, device=dev),
"se": torch.empty(nb, dtype=dtypes.i32, device=dev),
"nv": torch.empty(2, dtype=dtypes.i32, device=dev),
"ob": torch.empty((n_tok, dim), dtype=torch.bfloat16, device=dev),
}
_mem[tag] = d
return d
def _inter_buf(n_tok, k, inter, dev):
tag = ("i", n_tok, k, inter)
if tag in _mem:
return _mem[tag]
t = torch.empty((n_tok, k, inter), dtype=torch.bfloat16, device=dev)
_mem[tag] = t
return t
def _qt_alloc(rows, cols, sid_n, k, dev):
tag = ("q", rows, cols, sid_n, k)
if tag in _mem:
return _mem[tag]
BLK, BM, BN, BM2, BN2 = 32, 32, 8, 16, 4
sn = triton.cdiv(cols, BLK)
r = {
"f": torch.empty((rows, cols // 2), dtype=torch.uint8, device=dev),
"s": torch.empty(
(triton.cdiv(sid_n, BM), triton.cdiv(sn, BN), BN2, BM2, 4),
dtype=torch.uint8, device=dev),
}
_mem[tag] = r
return r
def _run_quant(x, sid, nval, ntok, k, bm, dev):
rows, cols = x.shape
BLK, BIG, BM, BN = 32, 128, 32, 8
sn = triton.cdiv(cols, BLK)
sid_n = sid.shape[0]
qb = _qt_alloc(rows, cols, sid_n, k, dev)
n_pid = triton.cdiv(rows, BIG) * sn + triton.cdiv(sid_n, BM) * triton.cdiv(sn, BN)
_fused_dynamic_mxfp4_quant_moe_sort_kernel[(n_pid,)](
x, qb["f"], sid, nval, qb["s"],
rows, cols, sn,
*x.stride(), *qb["f"].stride(), *qb["s"].stride(),
token_num=ntok, M_i=rows, N_i=sn,
MXFP4_QUANT_BLOCK_SIZE=BLK, BLOCK_SIZE_Mx=BIG,
BLOCK_SIZE_M=BM // 2, BLOCK_SIZE_N=BN // 2, TOPK=k,
)
return qb["f"].view(dtypes.fp4x2), qb["s"].view(dtypes.fp8_e8m0).view(-1, sn)
_loaded = False
def _setup():
global _loaded
if _loaded:
return
_loaded = True
if _fmoe.cfg_2stages is None:
import pandas as pd
from aiter.jit.core import AITER_CONFIGS
fp = AITER_CONFIGS.AITER_CONFIG_FMOE_FILE
if os.path.exists(fp):
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(fp)
if "_tag" in df.columns:
df = df[df["_tag"].fillna("") == ""]
_fmoe.cfg_2stages = df.set_index(cols).to_dict("index")
else:
_fmoe.cfg_2stages = {}
_fmoe.cfg_2stages.update(_SHAPE_MAP)
orig = _fmoe.get_2stage_cfgs
@functools.lru_cache(maxsize=2048)
def _hook(token, model_dim, inter_dim, expert, topk,
dtype, q_dtype_a, q_dtype_w, q_type, use_g1u1,
activation, doweight_stage1, hidden_pad, intermediate_pad, is_shuffled=True):
md = orig(token, model_dim, inter_dim, expert, topk,
dtype, q_dtype_a, q_dtype_w, q_type, use_g1u1,
activation, doweight_stage1, hidden_pad, intermediate_pad, is_shuffled)
from aiter.jit.utils.chip_info import get_cu_num
lk = (get_cu_num(), token, model_dim, inter_dim, expert, topk,
str(activation), str(dtype), str(q_dtype_a), str(q_dtype_w),
str(q_type), use_g1u1, doweight_stage1)
entry = _fmoe.cfg_2stages.get(lk)
if entry and entry.get("use_non_temporal_load") is not None:
nt = entry["use_non_temporal_load"]
s1 = md.stage1
if hasattr(s1, 'func') and s1.func is not None and 'use_non_temporal_load' in (s1.keywords or {}):
kw = dict(s1.keywords); kw['use_non_temporal_load'] = nt
md = _fmoe.MOEMetadata(
functools.partial(s1.func, **{k: v for k, v in kw.items()}),
md.stage2, md.block_m, md.ksplit, md.run_1stage, md.has_bias, nt)
s2 = md.stage2
if s2 and hasattr(s2, 'keywords') and 'use_non_temporal_load' in (s2.keywords or {}):
kw2 = dict(s2.keywords); kw2['use_non_temporal_load'] = nt
md = _fmoe.MOEMetadata(
md.stage1,
functools.partial(s2.func, **{k: v for k, v in kw2.items()}),
md.block_m, md.ksplit, md.run_1stage, md.has_bias, nt)
return md
_fmoe.get_2stage_cfgs = _hook
def custom_kernel(data: input_t) -> output_t:
(h, gu_w, d_w, gu_sc, d_sc, gu_ws, d_ws, gu_scs, d_scs, tw, ti, cfg) = data
_setup()
hp = cfg["d_hidden_pad"] - cfg["d_hidden"]
ip = cfg["d_expert_pad"] - cfg["d_expert"]
n = h.shape[0]
k = ti.shape[1]
dev = ti.device
ne, mdim, idim = get_inter_dim(gu_ws.shape, d_ws.shape)
md = get_2stage_cfgs(
get_padded_M(n), mdim, idim, ne, k,
torch.bfloat16, dtypes.fp4x2, dtypes.fp4x2,
QuantType.per_1x32, True, ActivationType.Silu,
False, hp, ip, True)
bm = int(md.block_m)
sb = _sorting_tensors(n, ne, k, mdim, bm, dev)
aiter.moe_sorting_fwd(ti, tw, sb["si"], sb["sw"], sb["se"], sb["nv"], sb["ob"], ne, bm, None, None, 0)
w1s = gu_scs.view(dtypes.fp8_e8m0)
w2s = d_scs.view(dtypes.fp8_e8m0)
if md.ksplit > 1:
a = h.to(torch.bfloat16)
inter = md.stage1(a, gu_ws, d_ws, sb["si"], sb["se"], sb["nv"],
_inter_buf(n, k, idim, dev), k,
block_m=bm, a1_scale=None, w1_scale=w1s, sorted_weights=None)
md.stage2(inter, gu_ws, d_ws, sb["si"], sb["se"], sb["nv"],
sb["ob"], k, w2_scale=w2s, a2_scale=None, block_m=bm, sorted_weights=sb["sw"])
else:
qa, qas = _run_quant(h, sb["si"], sb["nv"], n, 1, bm, dev)
inter = _inter_buf(n, k, idim, dev)
inter = md.stage1(qa, gu_ws, d_ws, sb["si"], sb["se"], sb["nv"],
inter, k, block_m=bm, a1_scale=qas, w1_scale=w1s, sorted_weights=None)
flat = inter.view(-1, idim)
qi, qis = _run_quant(flat, sb["si"], sb["nv"], n, k, bm, dev)
qi = qi.view(n, k, -1)
md.stage2(qi, gu_ws, d_ws, sb["si"], sb["se"], sb["nv"],
sb["ob"], k, w2_scale=w2s, a2_scale=qis, block_m=bm, sorted_weights=sb["sw"])
return sb["ob"]
scrolls · 224 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