submission 755012
flower2123 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 271 lines, June 9 Researcher Reciprocity License v1.0.
submission_flower_moe_dispatch.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-755012?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:7622d61a59ca6ed39d37bda81452cbdfef3c45026060f58fdf04a36d97649729
license declaredunknown
license concludedunknown
authorsflower2123
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Dispatch-tuned revision for amd-moe-mxfp4.Kernel source
submission_flower_moe_dispatch.py271 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
"""
Dispatch-tuned revision for amd-moe-mxfp4.
Per-shape dispatch choice matters: dispatch=2 helps the sparse-heavy
cases, while dispatch=0 is steadier on the larger dense shapes.
"""
# flower's notebook, condensed into comments:
# I did not try to outsmart the vendor kernels everywhere at once.
# The path that kept paying off was narrower:
# 1. keep the CK/FlyDSL stage machinery intact,
# 2. only override the benchmarked shapes,
# 3. let dispatch_policy be shape-specific instead of pretending one
# global switch is best for both E=33 and E=257.
# The result is intentionally plain: fewer moving parts, fewer places
# to accidentally lose a working kernel while chasing one micro-win.
import os
import functools
import torch
import triton
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
import aiter.fused_moe as _fused_moe_module
import aiter.ops.flydsl.moe_kernels as _flydsl_moe_kernels
from aiter.ops.triton._triton_kernels.quant.fused_mxfp4_quant import (
_fused_dynamic_mxfp4_quant_moe_sort_kernel,
)
_flydsl_moe_kernels._KERNEL_PARAMS["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,
}
def _make_key(token, inter_dim, expert):
return (
256, token, 7168, inter_dim, expert, 9,
"ActivationType.Silu", "torch.bfloat16",
"torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",
"QuantType.per_1x32", True, False,
)
_4WG = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
_FLY = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"
_CUSTOM_CONFIGS = {}
_DISPATCH = {}
_CUSTOM_CONFIGS[_make_key(16, 256, 257)] = {
"block_m": 16, "ksplit": 2, "kernelName1": "", "kernelName2": "", "run_1stage": False,
}
_DISPATCH[_make_key(16, 256, 257)] = 2
_CUSTOM_CONFIGS[_make_key(128, 256, 257)] = {
"block_m": 32, "ksplit": 0,
"kernelName1": _4WG, "kernelName2": _FLY, "run_1stage": False,
}
_DISPATCH[_make_key(128, 256, 257)] = 0
_CUSTOM_CONFIGS[_make_key(512, 256, 257)] = {
"block_m": 32, "ksplit": 0,
"kernelName1": _4WG, "kernelName2": _FLY, "run_1stage": False,
}
_DISPATCH[_make_key(512, 256, 257)] = 0
_CUSTOM_CONFIGS[_make_key(16, 512, 33)] = {
"block_m": 32, "ksplit": 2, "kernelName1": "", "kernelName2": "", "run_1stage": False,
}
_DISPATCH[_make_key(16, 512, 33)] = 0
_CUSTOM_CONFIGS[_make_key(128, 512, 33)] = {
"block_m": 32, "ksplit": 0,
"kernelName1": _4WG, "kernelName2": _FLY, "run_1stage": False,
}
_DISPATCH[_make_key(128, 512, 33)] = 2
_CUSTOM_CONFIGS[_make_key(512, 512, 33)] = {
"block_m": 64, "ksplit": 0,
"kernelName1": _4WG, "kernelName2": _FLY, "run_1stage": False,
}
_DISPATCH[_make_key(512, 512, 33)] = 0
_CUSTOM_CONFIGS[_make_key(512, 2048, 33)] = {
"block_m": 32, "ksplit": 0,
"kernelName1": _4WG, "kernelName2": _FLY, "run_1stage": False,
}
_DISPATCH[_make_key(512, 2048, 33)] = 0
# flower: the only per-shape knobs I am willing to touch here are
# block_m, ksplit, and dispatch policy. Everything else stays tied to
# the known-good CK/FlyDSL pair to keep the risk profile low.
# --- Buffer caches ---
_buf = {}
_qbuf = {}
def _get_a2(M, topk, inter_dim, device):
key = ("a2", M, topk, inter_dim)
if key not in _buf:
_buf[key] = torch.empty((M, topk, inter_dim), dtype=torch.bfloat16, device=device)
return _buf[key]
def _get_sorting_bufs(M, E, topk, model_dim, block_m, device):
key = (M, E, topk, model_dim, block_m)
if key not in _buf:
max_pad = M * topk + E * block_m - topk
max_blk = (max_pad + block_m - 1) // block_m
_buf[key] = {
"sid": torch.empty(max_pad, dtype=dtypes.i32, device=device),
"sw": torch.empty(max_pad, dtype=dtypes.fp32, device=device),
"se": torch.empty(max_blk, dtype=dtypes.i32, device=device),
"nv": torch.empty(2, dtype=dtypes.i32, device=device),
"out": torch.empty((M, model_dim), dtype=torch.bfloat16, device=device),
"a2": torch.empty((M, topk, 0), dtype=torch.bfloat16, device=device),
}
return _buf[key]
_injected = False
def _inject():
global _injected
if _injected:
return
_injected = True
if _fused_moe_module.cfg_2stages is None:
import pandas as pd
from aiter.jit.core import AITER_CONFIGS
tf = AITER_CONFIGS.AITER_CONFIG_FMOE_FILE
if os.path.exists(tf):
_IDX = [
"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(tf)
if "_tag" in df.columns:
df = df[df["_tag"].fillna("") == ""]
_fused_moe_module.cfg_2stages = df.set_index(_IDX).to_dict("index")
else:
_fused_moe_module.cfg_2stages = {}
_fused_moe_module.cfg_2stages.update(_CUSTOM_CONFIGS)
def _quant_prealloc(x, sorted_ids, num_valid_ids, token_num, topk, block_m, device):
# flower's trail here:
# I kept quant + sort fused because splitting them made the code look
# cleaner but immediately cost memory traffic. This helper is ugly on
# purpose; the ugliness is cheaper than another round-trip.
M, N = x.shape
QBS = 32
BLK_Mx = 64
BLK_M, BLK_N = 32, 8
BLK_M_u32, BLK_N_u32 = 16, 4
scaleN = triton.cdiv(N, QBS)
M_o = sorted_ids.shape[0]
qk = (M, N, M_o, topk)
if qk not in _qbuf:
_qbuf[qk] = {
"fp4": torch.empty((M, N // 2), dtype=torch.uint8, device=device),
"bs": torch.empty(
(triton.cdiv(M_o, BLK_M), triton.cdiv(scaleN, BLK_N),
BLK_N_u32, BLK_M_u32, 4),
dtype=torch.uint8, device=device,
),
}
qb = _qbuf[qk]
num_pid = triton.cdiv(M, BLK_Mx) * scaleN + triton.cdiv(
M_o, BLK_M
) * triton.cdiv(scaleN, BLK_N)
_fused_dynamic_mxfp4_quant_moe_sort_kernel[(num_pid,)](
x, qb["fp4"], sorted_ids, num_valid_ids, qb["bs"],
M, N, scaleN,
*x.stride(), *qb["fp4"].stride(), *qb["bs"].stride(),
token_num=token_num, M_i=M, N_i=scaleN,
MXFP4_QUANT_BLOCK_SIZE=QBS, BLOCK_SIZE_Mx=BLK_Mx,
BLOCK_SIZE_M=BLK_M // 2, BLOCK_SIZE_N=BLK_N // 2,
TOPK=topk,
)
return (
qb["fp4"].view(dtypes.fp4x2),
qb["bs"].view(dtypes.fp8_e8m0).view(-1, scaleN),
)
def custom_kernel(data: input_t) -> output_t:
(
hidden_states, _w1r, _w2r, _w1sr, _w2sr,
w1, w2, w1s, w2s,
topk_weights, topk_ids, config,
) = data
_inject()
M = hidden_states.shape[0]
topk = topk_ids.shape[1]
device = hidden_states.device
dhp = config["d_hidden_pad"]
dep = config["d_expert_pad"]
hidden_pad = dhp - config["d_hidden"]
intermediate_pad = dep - config["d_expert"]
E, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)
padded_M = get_padded_M(M)
metadata = get_2stage_cfgs(
padded_M, model_dim, inter_dim, E, topk,
torch.bfloat16, dtypes.fp4x2, dtypes.fp4x2,
QuantType.per_1x32, True, ActivationType.Silu,
False, hidden_pad, intermediate_pad, True,
)
block_m = int(metadata.block_m)
cfg_key = _make_key(padded_M, inter_dim, E)
dp = _DISPATCH.get(cfg_key, 0)
b = _get_sorting_bufs(M, E, topk, model_dim, block_m, device)
aiter.moe_sorting_fwd(
topk_ids, topk_weights,
b["sid"], b["sw"], b["se"], b["nv"], b["out"],
E, block_m, None, None, dp,
)
w1sv = w1s.view(dtypes.fp8_e8m0)
w2sv = w2s.view(dtypes.fp8_e8m0)
a2_buf = _get_a2(M, topk, inter_dim, device)
if metadata.ksplit > 1:
# flower: I leave the ksplit branch in BF16 on purpose.
# I tried being more uniform before, but these cases were happier
# when I stopped forcing them through the narrower path.
a1 = hidden_states.to(torch.bfloat16)
a2 = metadata.stage1(
a1, w1, w2, b["sid"], b["se"], b["nv"], a2_buf, topk,
block_m=block_m, a1_scale=None, w1_scale=w1sv, sorted_weights=None,
)
metadata.stage2(
a2, w1, w2, b["sid"], b["se"], b["nv"], b["out"], topk,
w2_scale=w2sv, a2_scale=None, block_m=block_m, sorted_weights=b["sw"],
)
else:
# flower: the non-ksplit path is where the narrower traffic pattern
# finally won consistently, so I keep the re-quant step explicit
# instead of hiding it behind a smarter-looking abstraction.
a1, a1s = _quant_prealloc(
hidden_states, b["sid"], b["nv"], M, 1, block_m, device,
)
a2 = metadata.stage1(
a1, w1, w2, b["sid"], b["se"], b["nv"], a2_buf, topk,
block_m=block_m, a1_scale=a1s, w1_scale=w1sv, sorted_weights=None,
)
a2_flat = a2.view(-1, inter_dim)
a2q, a2s = _quant_prealloc(
a2_flat, b["sid"], b["nv"], M, topk, block_m, device,
)
a2q = a2q.view(M, topk, -1)
metadata.stage2(
a2q, w1, w2, b["sid"], b["se"], b["nv"], b["out"], topk,
w2_scale=w2sv, a2_scale=a2s, block_m=block_m, sorted_weights=b["sw"],
)
return b["out"]
scrolls · 271 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 754773.
- # ============================================================- # flower_moe.py — Mixture-of-Experts MXFP4 kernel- # Author: flower- # Description: CK 2-stage pipeline + FlyDSL stage2 + fused- # dynamic MXFP4 quant with per-shape block_m optimization- # for AMD MI355X (gfx950). Supports E=33 and E=257.- # flower's note: block_m tuning key insight — sparse experts- # (E=257) prefer smaller block_m, dense experts (E=33) with- # small K prefer larger block_m to reduce block overhead.- # ============================================================#!POPCORN leaderboard amd-moe-mxfp4#!POPCORN gpu MI355X- # -- flower: standard imports --- import os # flower: filesystem operations- import functools # flower: partial / lru_cache- import torch # flower: tensor operations- import triton # flower: for cdiv utility- from task import input_t, output_t # flower: competition types+ """+ Dispatch-tuned revision for amd-moe-mxfp4.+ Per-shape dispatch choice matters: dispatch=2 helps the sparse-heavy+ cases, while dispatch=0 is steadier on the larger dense shapes.+ """- # -- flower: aiter library imports --- import aiter # flower: AMD AI tensor engine- from aiter import ActivationType, QuantType, dtypes # flower: enum types and custom dtypes- from aiter.fused_moe import get_2stage_cfgs, get_padded_M, get_inter_dim # flower: MoE utilities- import aiter.fused_moe as _flower_moe_lib # flower: MoE module reference- import aiter.ops.flydsl.moe_kernels as _flower_fly_kernels # flower: FlyDSL kernel registry- from aiter.ops.triton._triton_kernels.quant.fused_mxfp4_quant import ( # flower: fused quant kernel+ # flower's notebook, condensed into comments:+ # I did not try to outsmart the vendor kernels everywhere at once.+ # The path that kept paying off was narrower:+ # 1. keep the CK/FlyDSL stage machinery intact,+ # 2. only override the benchmarked shapes,+ # 3. let dispatch_policy be shape-specific instead of pretending one+ # global switch is best for both E=33 and E=257.+ # The result is intentionally plain: fewer moving parts, fewer places+ # to accidentally lose a working kernel while chasing one micro-win.++ import os+ import functools+ import torch+ import triton+ 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+ import aiter.fused_moe as _fused_moe_module+ import aiter.ops.flydsl.moe_kernels as _flydsl_moe_kernels+ from aiter.ops.triton._triton_kernels.quant.fused_mxfp4_quant import (_fused_dynamic_mxfp4_quant_moe_sort_kernel,)- # ============================================================- # flower: register FlyDSL t16x128x128 kernel for stage2- # This tile config handles the post-SwiGLU down-projection- # with atomic accumulation across expert splits.- # ============================================================- _flower_fly_kernels._KERNEL_PARAMS["flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"] = {- "stage": 2, # flower: second stage of 2-stage pipeline- "a_dtype": "fp4", # flower: activation dtype- "b_dtype": "fp4", # flower: weight dtype- "out_dtype": "bf16", # flower: output dtype- "tile_m": 16, # flower: M-tile size- "tile_n": 128, # flower: N-tile size- "tile_k": 128, # flower: K-tile size- "mode": "atomic", # flower: use atomic accumulation- "MPerBlock": 16, # flower: M per CU block+ _flydsl_moe_kernels._KERNEL_PARAMS["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,}- # ============================================================- # flower: CK kernel names for stage1- # These are the hardware-optimized CK GEMM kernels that handle- # the gate_up projection with fused SwiGLU activation.- # ============================================================- _FLOWER_CK_S1_M128 = ( # flower: 4-workgroup stage1, M128 tile- "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3"- "_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"- )- _FLOWER_FLY_S2 = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic" # flower: FlyDSL stage2-- # ============================================================- # flower: per-shape configuration override table- # Maps (cu_num, token, model_dim, inter_dim, expert, topk, ...)- # to optimal (block_m, ksplit, kernel names) combinations.- # flower's trail of thought:- # 1. I first tried to trust the stock aiter table as-is.- # 2. Then I only varied block_m and immediately saw E=257 and E=33- # want opposite behavior because token density per expert flips.- # 3. After that, the stable pattern was to leave the big CK/FlyDSL- # kernels intact and only override the exact shapes that appear in- # the competition benches instead of pretending one heuristic wins- # for every case.- # ============================================================- _flower_shape_overrides = {} # flower: populated below--- def _flower_make_cfg_key(num_tokens, expert_inter_dim, num_experts):- """flower: construct the canonical config lookup key."""- return ( # flower: full key tuple matching aiter's internal format- 256, num_tokens, 7168, expert_inter_dim, num_experts, 9,+ def _make_key(token, inter_dim, expert):+ return (+ 256, token, 7168, inter_dim, expert, 9,"ActivationType.Silu", "torch.bfloat16","torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2","QuantType.per_1x32", True, False,)+ _4WG = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"+ _FLY = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"- # flower: E=257 shapes — sparse routing, fewer tokens per expert- _flower_shape_overrides[_flower_make_cfg_key(16, 256, 257)] = { # flower: bs=16, ksplit=2+ _CUSTOM_CONFIGS = {}+ _DISPATCH = {}++ _CUSTOM_CONFIGS[_make_key(16, 256, 257)] = {"block_m": 16, "ksplit": 2, "kernelName1": "", "kernelName2": "", "run_1stage": False,}- _flower_shape_overrides[_flower_make_cfg_key(128, 256, 257)] = { # flower: bs=128, block_m=32- "block_m": 32, "ksplit": 0, "kernelName1": _FLOWER_CK_S1_M128,- "kernelName2": _FLOWER_FLY_S2, "run_1stage": False,+ _DISPATCH[_make_key(16, 256, 257)] = 2++ _CUSTOM_CONFIGS[_make_key(128, 256, 257)] = {+ "block_m": 32, "ksplit": 0,+ "kernelName1": _4WG, "kernelName2": _FLY, "run_1stage": False,}- _flower_shape_overrides[_flower_make_cfg_key(512, 256, 257)] = { # flower: bs=512, block_m=32- "block_m": 32, "ksplit": 0, "kernelName1": _FLOWER_CK_S1_M128,- "kernelName2": _FLOWER_FLY_S2, "run_1stage": False,+ _DISPATCH[_make_key(128, 256, 257)] = 0++ _CUSTOM_CONFIGS[_make_key(512, 256, 257)] = {+ "block_m": 32, "ksplit": 0,+ "kernelName1": _4WG, "kernelName2": _FLY, "run_1stage": False,}+ _DISPATCH[_make_key(512, 256, 257)] = 0- # flower: E=33 shapes — dense routing, more tokens per expert- _flower_shape_overrides[_flower_make_cfg_key(16, 512, 33)] = { # flower: bs=16, sparse with ksplit=2+ _CUSTOM_CONFIGS[_make_key(16, 512, 33)] = {"block_m": 32, "ksplit": 2, "kernelName1": "", "kernelName2": "", "run_1stage": False,}- _flower_shape_overrides[_flower_make_cfg_key(128, 512, 33)] = { # flower: bs=128, block_m=32- "block_m": 32, "ksplit": 0, "kernelName1": _FLOWER_CK_S1_M128,- "kernelName2": _FLOWER_FLY_S2, "run_1stage": False,+ _DISPATCH[_make_key(16, 512, 33)] = 0++ _CUSTOM_CONFIGS[_make_key(128, 512, 33)] = {+ "block_m": 32, "ksplit": 0,+ "kernelName1": _4WG, "kernelName2": _FLY, "run_1stage": False,}- _flower_shape_overrides[_flower_make_cfg_key(512, 512, 33)] = { # flower: bs=512, K=512 dense, block_m=64 wins- "block_m": 64, "ksplit": 0, "kernelName1": _FLOWER_CK_S1_M128,- "kernelName2": _FLOWER_FLY_S2, "run_1stage": False,+ _DISPATCH[_make_key(128, 512, 33)] = 2++ _CUSTOM_CONFIGS[_make_key(512, 512, 33)] = {+ "block_m": 64, "ksplit": 0,+ "kernelName1": _4WG, "kernelName2": _FLY, "run_1stage": False,}- _flower_shape_overrides[_flower_make_cfg_key(512, 2048, 33)] = { # flower: bs=512, K=2048, block_m=32 amortizes- "block_m": 32, "ksplit": 0, "kernelName1": _FLOWER_CK_S1_M128,- "kernelName2": _FLOWER_FLY_S2, "run_1stage": False,+ _DISPATCH[_make_key(512, 512, 33)] = 0++ _CUSTOM_CONFIGS[_make_key(512, 2048, 33)] = {+ "block_m": 32, "ksplit": 0,+ "kernelName1": _4WG, "kernelName2": _FLY, "run_1stage": False,}+ _DISPATCH[_make_key(512, 2048, 33)] = 0- # ============================================================- # flower: pre-allocated buffer caches- # ============================================================- _flower_sort_bufs = {} # flower: moe_sorting output buffers- _flower_quant_bufs = {} # flower: fused quant output buffers+ # flower: the only per-shape knobs I am willing to touch here are+ # block_m, ksplit, and dispatch policy. Everything else stays tied to+ # the known-good CK/FlyDSL pair to keep the risk profile low.+ # --- Buffer caches ---+ _buf = {}+ _qbuf = {}- def _flower_get_a2_buffer(num_tokens, top_k, inter_dim, device):- """flower: get or allocate intermediate activation buffer (stage1 output)."""- cache_key = ("a2", num_tokens, top_k, inter_dim) # flower: unique key- if cache_key not in _flower_sort_bufs: # flower: allocate on first use- _flower_sort_bufs[cache_key] = torch.empty(- (num_tokens, top_k, inter_dim), dtype=torch.bfloat16, device=device,- )- return _flower_sort_bufs[cache_key] # flower: return cached buffer+ def _get_a2(M, topk, inter_dim, device):+ key = ("a2", M, topk, inter_dim)+ if key not in _buf:+ _buf[key] = torch.empty((M, topk, inter_dim), dtype=torch.bfloat16, device=device)+ return _buf[key]-- def _flower_get_sort_buffers(num_tokens, num_experts, top_k, model_dim, blk_m, device):- """flower: get or allocate all buffers needed for moe_sorting_fwd."""- cache_key = (num_tokens, num_experts, top_k, model_dim, blk_m) # flower: unique key- if cache_key not in _flower_sort_bufs: # flower: allocate on first use- max_padded = num_tokens * top_k + num_experts * blk_m - top_k # flower: max padded token count- max_blocks = (max_padded + blk_m - 1) // blk_m # flower: max block count- _flower_sort_bufs[cache_key] = { # flower: all sorting tensors- "sorted_ids": torch.empty(max_padded, dtype=dtypes.i32, device=device),- "sorted_weights": torch.empty(max_padded, dtype=dtypes.fp32, device=device),- "sorted_expert_ids": torch.empty(max_blocks, dtype=dtypes.i32, device=device),- "num_valid": torch.empty(2, dtype=dtypes.i32, device=device),- "output": torch.empty((num_tokens, model_dim), dtype=torch.bfloat16, device=device),- "a2_placeholder": torch.empty((num_tokens, top_k, 0), dtype=torch.bfloat16, device=device),+ def _get_sorting_bufs(M, E, topk, model_dim, block_m, device):+ key = (M, E, topk, model_dim, block_m)+ if key not in _buf:+ max_pad = M * topk + E * block_m - topk+ max_blk = (max_pad + block_m - 1) // block_m+ _buf[key] = {+ "sid": torch.empty(max_pad, dtype=dtypes.i32, device=device),+ "sw": torch.empty(max_pad, dtype=dtypes.fp32, device=device),+ "se": torch.empty(max_blk, dtype=dtypes.i32, device=device),+ "nv": torch.empty(2, dtype=dtypes.i32, device=device),+ "out": torch.empty((M, model_dim), dtype=torch.bfloat16, device=device),+ "a2": torch.empty((M, topk, 0), dtype=torch.bfloat16, device=device),}- return _flower_sort_bufs[cache_key] # flower: return cached dict+ return _buf[key]- def _flower_fused_quant(input_tensor, sorted_ids, num_valid, token_count, top_k, blk_m, device):- """flower: inline fused_dynamic_mxfp4_quant_moe_sort with pre-allocated outputs."""- rows, cols = input_tensor.shape # flower: input dimensions- QUANT_GRP = 32 # flower: MXFP4 group size- TILE_Mx = 128 # flower: quant tile M- TILE_M, TILE_N = 32, 8 # flower: sort-scatter tile- TILE_M_u32, TILE_N_u32 = 16, 4 # flower: packed tile dims+ _injected = False- num_scale_cols = triton.cdiv(cols, QUANT_GRP) # flower: scale column count- sorted_len = sorted_ids.shape[0] # flower: padded sorted length+ def _inject():+ global _injected+ if _injected:+ return+ _injected = True+ if _fused_moe_module.cfg_2stages is None:+ import pandas as pd+ from aiter.jit.core import AITER_CONFIGS+ tf = AITER_CONFIGS.AITER_CONFIG_FMOE_FILE+ if os.path.exists(tf):+ _IDX = [+ "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(tf)+ if "_tag" in df.columns:+ df = df[df["_tag"].fillna("") == ""]+ _fused_moe_module.cfg_2stages = df.set_index(_IDX).to_dict("index")+ else:+ _fused_moe_module.cfg_2stages = {}+ _fused_moe_module.cfg_2stages.update(_CUSTOM_CONFIGS)- # flower: get or allocate quant output buffers- quant_key = (rows, cols, sorted_len, top_k)- if quant_key not in _flower_quant_bufs:- _flower_quant_bufs[quant_key] = {- "fp4_out": torch.empty((rows, cols // 2), dtype=torch.uint8, device=device),- "scale_out": torch.empty(- (triton.cdiv(sorted_len, TILE_M), triton.cdiv(num_scale_cols, TILE_N),- TILE_N_u32, TILE_M_u32, 4),+ def _quant_prealloc(x, sorted_ids, num_valid_ids, token_num, topk, block_m, device):+ # flower's trail here:+ # I kept quant + sort fused because splitting them made the code look+ # cleaner but immediately cost memory traffic. This helper is ugly on+ # purpose; the ugliness is cheaper than another round-trip.+ M, N = x.shape+ QBS = 32+ BLK_Mx = 64+ BLK_M, BLK_N = 32, 8+ BLK_M_u32, BLK_N_u32 = 16, 4++ scaleN = triton.cdiv(N, QBS)+ M_o = sorted_ids.shape[0]++ qk = (M, N, M_o, topk)+ if qk not in _qbuf:+ _qbuf[qk] = {+ "fp4": torch.empty((M, N // 2), dtype=torch.uint8, device=device),+ "bs": torch.empty(+ (triton.cdiv(M_o, BLK_M), triton.cdiv(scaleN, BLK_N),+ BLK_N_u32, BLK_M_u32, 4),dtype=torch.uint8, device=device,),}- qb = _flower_quant_bufs[quant_key] # flower: get cached buffers+ qb = _qbuf[qk]- # flower: compute grid size for fused quant+sort kernel- # flower's trail of thought:- # I did not split quantization and sorting back apart after measuring.- # The separate version was easier to read, but it forced one more trip- # through memory and that showed up immediately on MI355X once token- # count got larger. The fused path is uglier, but the data movement is- # the thing that actually mattered here.- total_pids = (- triton.cdiv(rows, TILE_Mx) * num_scale_cols- + triton.cdiv(sorted_len, TILE_M) * triton.cdiv(num_scale_cols, TILE_N)- )+ num_pid = triton.cdiv(M, BLK_Mx) * scaleN + triton.cdiv(+ M_o, BLK_M+ ) * triton.cdiv(scaleN, BLK_N)- # flower: launch the fused quant+sort kernel- _fused_dynamic_mxfp4_quant_moe_sort_kernel[(total_pids,)](- input_tensor, qb["fp4_out"], sorted_ids, num_valid, qb["scale_out"],- rows, cols, num_scale_cols,- *input_tensor.stride(), *qb["fp4_out"].stride(), *qb["scale_out"].stride(),- token_num=token_count, M_i=rows, N_i=num_scale_cols,- MXFP4_QUANT_BLOCK_SIZE=QUANT_GRP, BLOCK_SIZE_Mx=TILE_Mx,- BLOCK_SIZE_M=TILE_M // 2, BLOCK_SIZE_N=TILE_N // 2,- TOPK=top_k,+ _fused_dynamic_mxfp4_quant_moe_sort_kernel[(num_pid,)](+ x, qb["fp4"], sorted_ids, num_valid_ids, qb["bs"],+ M, N, scaleN,+ *x.stride(), *qb["fp4"].stride(), *qb["bs"].stride(),+ token_num=token_num, M_i=M, N_i=scaleN,+ MXFP4_QUANT_BLOCK_SIZE=QBS, BLOCK_SIZE_Mx=BLK_Mx,+ BLOCK_SIZE_M=BLK_M // 2, BLOCK_SIZE_N=BLK_N // 2,+ TOPK=topk,)- # flower: return with proper dtype viewsreturn (- qb["fp4_out"].view(dtypes.fp4x2), # flower: packed FP4 view- qb["scale_out"].view(dtypes.fp8_e8m0).view(-1, num_scale_cols), # flower: E8M0 scale view+ qb["fp4"].view(dtypes.fp4x2),+ qb["bs"].view(dtypes.fp8_e8m0).view(-1, scaleN),)-- # ============================================================- # flower: one-time config injection into aiter's MoE module- # ============================================================- _flower_configs_loaded = False # flower: injection guard--- def _flower_load_configs():- """flower: inject per-shape overrides into aiter's 2-stage config table."""- global _flower_configs_loaded # flower: module-level flag- if _flower_configs_loaded: # flower: already done- return- _flower_configs_loaded = True # flower: mark as done-- # flower: load default tuning CSV if not yet loaded- if _flower_moe_lib.cfg_2stages is None:- import pandas as pd # flower: for CSV parsing- from aiter.jit.core import AITER_CONFIGS # flower: config paths- tune_csv = AITER_CONFIGS.AITER_CONFIG_FMOE_FILE # flower: tuning file path- if os.path.exists(tune_csv): # flower: file exists- idx_cols = [ # flower: index columns for the tuning table- "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_csv) # flower: read CSV- if "_tag" in df.columns: # flower: filter out tagged rows- df = df[df["_tag"].fillna("") == ""]- _flower_moe_lib.cfg_2stages = df.set_index(idx_cols).to_dict("index") # flower: set as dict- else:- _flower_moe_lib.cfg_2stages = {} # flower: empty table-- # flower: merge our per-shape overrides on top of defaults- _flower_moe_lib.cfg_2stages.update(_flower_shape_overrides)--- # ============================================================- # flower: competition entry point- # ============================================================def custom_kernel(data: input_t) -> output_t:- """flower: MoE kernel — 2-stage CK pipeline with FlyDSL stage2."""- # flower: unpack input data tuple(- hidden_states, # flower: input activations [M, model_dim]- _raw_w1, _raw_w2, # flower: raw (non-shuffled) weights — unused- _raw_w1s, _raw_w2s, # flower: raw weight scales — unused- gate_up_w_sh, down_w_sh, # flower: shuffled weights- gate_up_sc_sh, down_sc_sh, # flower: shuffled weight scales- routing_weights, # flower: expert routing weights- routing_ids, # flower: expert routing indices- task_config, # flower: shape configuration dict+ hidden_states, _w1r, _w2r, _w1sr, _w2sr,+ w1, w2, w1s, w2s,+ topk_weights, topk_ids, config,) = data- # flower: inject per-shape configs into aiter- _flower_load_configs()+ _inject()- # flower: extract dimensions- num_tokens = hidden_states.shape[0] # flower: batch size- top_k = routing_ids.shape[1] # flower: experts per token- device = hidden_states.device # flower: GPU device+ M = hidden_states.shape[0]+ topk = topk_ids.shape[1]+ device = hidden_states.device+ dhp = config["d_hidden_pad"]+ dep = config["d_expert_pad"]+ hidden_pad = dhp - config["d_hidden"]+ intermediate_pad = dep - config["d_expert"]- # flower: compute padding amounts- hidden_pad = task_config["d_hidden_pad"] - task_config["d_hidden"] # flower: hidden dim padding- expert_pad = task_config["d_expert_pad"] - task_config["d_expert"] # flower: expert dim padding+ E, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)+ padded_M = get_padded_M(M)- # flower: get model structure dims- num_experts, model_dim, inter_dim = get_inter_dim(gate_up_w_sh.shape, down_w_sh.shape)- padded_m = get_padded_M(num_tokens) # flower: padded M for alignment-- # flower: get 2-stage pipeline metadata (stage1/stage2 functions, block_m, ksplit)- pipeline_meta = get_2stage_cfgs(- padded_m, model_dim, inter_dim, num_experts, top_k,+ metadata = get_2stage_cfgs(+ padded_M, model_dim, inter_dim, E, topk,torch.bfloat16, dtypes.fp4x2, dtypes.fp4x2,QuantType.per_1x32, True, ActivationType.Silu,- False, hidden_pad, expert_pad, True,+ False, hidden_pad, intermediate_pad, True,)- blk_m = int(pipeline_meta.block_m) # flower: block_m for this shape+ block_m = int(metadata.block_m)- # flower: step 1 — sort tokens by expert assignment- sort_bufs = _flower_get_sort_buffers(- num_tokens, num_experts, top_k, model_dim, blk_m, device,+ cfg_key = _make_key(padded_M, inter_dim, E)+ dp = _DISPATCH.get(cfg_key, 0)++ b = _get_sorting_bufs(M, E, topk, model_dim, block_m, device)+ aiter.moe_sorting_fwd(+ topk_ids, topk_weights,+ b["sid"], b["sw"], b["se"], b["nv"], b["out"],+ E, block_m, None, None, dp,)- aiter.moe_sorting_fwd( # flower: sort tokens → sorted_ids, weights, expert_ids- routing_ids, routing_weights,- sort_bufs["sorted_ids"], sort_bufs["sorted_weights"],- sort_bufs["sorted_expert_ids"], sort_bufs["num_valid"],- sort_bufs["output"],- num_experts, blk_m, None, None, 0,- )- # flower: prepare weight scale views as E8M0- w1_scale_view = gate_up_sc_sh.view(dtypes.fp8_e8m0) # flower: gate_up scale- w2_scale_view = down_sc_sh.view(dtypes.fp8_e8m0) # flower: down scale- a2_buffer = _flower_get_a2_buffer(num_tokens, top_k, inter_dim, device) # flower: intermediate buf+ w1sv = w1s.view(dtypes.fp8_e8m0)+ w2sv = w2s.view(dtypes.fp8_e8m0)+ a2_buf = _get_a2(M, topk, inter_dim, device)- # flower: step 2 — execute 2-stage pipeline- if pipeline_meta.ksplit > 1:- # flower: BF16 path (ksplit > 1 uses cktile, no fp4 quant of activations)- # flower's trail of thought:- # I kept this branch explicit because my first instinct was to- # force everything through the same FP4 route. That looked cleaner- # in code, but the ksplit>1 cases were happier when I stopped being- # clever and just let cktile keep the activations in BF16.- bf16_input = hidden_states.to(torch.bfloat16) # flower: ensure BF16- a2_result = pipeline_meta.stage1( # flower: gate_up GEMM + SwiGLU- bf16_input, gate_up_w_sh, down_w_sh,- sort_bufs["sorted_ids"], sort_bufs["sorted_expert_ids"],- sort_bufs["num_valid"], a2_buffer, top_k,- block_m=blk_m, a1_scale=None, w1_scale=w1_scale_view, sorted_weights=None,+ if metadata.ksplit > 1:+ # flower: I leave the ksplit branch in BF16 on purpose.+ # I tried being more uniform before, but these cases were happier+ # when I stopped forcing them through the narrower path.+ a1 = hidden_states.to(torch.bfloat16)+ a2 = metadata.stage1(+ a1, w1, w2, b["sid"], b["se"], b["nv"], a2_buf, topk,+ block_m=block_m, a1_scale=None, w1_scale=w1sv, sorted_weights=None,)- pipeline_meta.stage2( # flower: down GEMM + weighted reduction- a2_result, gate_up_w_sh, down_w_sh,- sort_bufs["sorted_ids"], sort_bufs["sorted_expert_ids"],- sort_bufs["num_valid"], sort_bufs["output"], top_k,- w2_scale=w2_scale_view, a2_scale=None,- block_m=blk_m, sorted_weights=sort_bufs["sorted_weights"],+ metadata.stage2(+ a2, w1, w2, b["sid"], b["se"], b["nv"], b["out"], topk,+ w2_scale=w2sv, a2_scale=None, block_m=block_m, sorted_weights=b["sw"],)else:- # flower: FP4 path (ksplit=0, quantize activations to MXFP4)- # flower's trail of thought:- # This is the branch I arrived at last: once block_m was settled,- # the next real gain came from accepting one re-quant between the- # two stages instead of carrying a wider intermediate around.- # The code looks more busy, but the benchmark liked the narrower- # traffic pattern more than the tidier BF16-only version.- # flower: quantize input activations- a1_quant, a1_scale = _flower_fused_quant(- hidden_states, sort_bufs["sorted_ids"], sort_bufs["num_valid"],- num_tokens, 1, blk_m, device,+ # flower: the non-ksplit path is where the narrower traffic pattern+ # finally won consistently, so I keep the re-quant step explicit+ # instead of hiding it behind a smarter-looking abstraction.+ a1, a1s = _quant_prealloc(+ hidden_states, b["sid"], b["nv"], M, 1, block_m, device,)- # flower: stage1 — gate_up GEMM with quantized activations- a2_result = pipeline_meta.stage1(- a1_quant, gate_up_w_sh, down_w_sh,- sort_bufs["sorted_ids"], sort_bufs["sorted_expert_ids"],- sort_bufs["num_valid"], a2_buffer, top_k,- block_m=blk_m, a1_scale=a1_scale, w1_scale=w1_scale_view, sorted_weights=None,+ a2 = metadata.stage1(+ a1, w1, w2, b["sid"], b["se"], b["nv"], a2_buf, topk,+ block_m=block_m, a1_scale=a1s, w1_scale=w1sv, sorted_weights=None,)- # flower: inter-stage re-quantization (BF16 -> FP4 for stage2 input)- a2_flat = a2_result.view(-1, inter_dim) # flower: flatten for quant- a2_quant, a2_scale = _flower_fused_quant(- a2_flat, sort_bufs["sorted_ids"], sort_bufs["num_valid"],- num_tokens, top_k, blk_m, device,+ a2_flat = a2.view(-1, inter_dim)+ a2q, a2s = _quant_prealloc(+ a2_flat, b["sid"], b["nv"], M, topk, block_m, device,)- a2_quant = a2_quant.view(num_tokens, top_k, -1) # flower: reshape back- # flower: stage2 — down GEMM + weighted reduction- pipeline_meta.stage2(- a2_quant, gate_up_w_sh, down_w_sh,- sort_bufs["sorted_ids"], sort_bufs["sorted_expert_ids"],- sort_bufs["num_valid"], sort_bufs["output"], top_k,- w2_scale=w2_scale_view, a2_scale=a2_scale,- block_m=blk_m, sorted_weights=sort_bufs["sorted_weights"],+ a2q = a2q.view(M, topk, -1)+ metadata.stage2(+ a2q, w1, w2, b["sid"], b["se"], b["nv"], b["out"], topk,+ w2_scale=w2sv, a2_scale=a2s, block_m=block_m, sorted_weights=b["sw"],)- return sort_bufs["output"] # flower: return final MoE output+ return b["out"]
scrolls · 559 diff lines total
Best evidence level for this revision: reported
JSON