submission 681568
ooousay · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 579 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-681568?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:b2dcdfa9864b7c4c28be80d74c14b2fc5be0fa1c692ffb71ee25d409518bbff4
license declaredunknown
license concludedunknown
authorsooousay
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",Kernel source
submission.py579 lines
#!POPCORN leaderboard amd-moe-mxfp4
"""Auto-generated by build.py"""
import os
import sys
import functools
import torch
import triton
import triton.language as tl
try:
import triton.experimental.gluon as gluon
import triton.experimental.gluon.language as gl
_HAS_GLUON = True
except Exception:
_HAS_GLUON = False
def _test_gluon_mfma_scaled():
pass
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe
import aiter.fused_moe as _fused_moe_module
import aiter.ops.flydsl.moe_kernels as _flydsl_moe_kernels
# Register FlyDSL tile_k=128 kernels not in server defaults
_flydsl_moe_kernels._KERNEL_PARAMS["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_moe_kernels._KERNEL_PARAMS["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_moe_kernels._KERNEL_PARAMS["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_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,
}
# Collected configs from all shapes ? each kernel.py adds its own
CUSTOM_CONFIGS = {}
# Registry for custom stage2 functions: {(inter_dim, expert): callable}
# kernel.py files register their custom stage2 here before fused_moe runs.
CUSTOM_STAGE2_REGISTRY = {}
_injected = False
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,
)
def inject_configs():
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
tune_file = AITER_CONFIGS.AITER_CONFIG_FMOE_FILE
if os.path.exists(tune_file):
_INDEX_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("") == ""]
_fused_moe_module.cfg_2stages = df.set_index(_INDEX_COLS).to_dict("index")
else:
_fused_moe_module.cfg_2stages = {}
_fused_moe_module.cfg_2stages.update(CUSTOM_CONFIGS)
_original_get_2stage_cfgs = _fused_moe_module.get_2stage_cfgs
@functools.lru_cache(maxsize=2048)
def _patched_get_2stage_cfgs(
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,
):
metadata = _original_get_2stage_cfgs(
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
cu_num = get_cu_num()
keys = (
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,
)
cfg = _fused_moe_module.cfg_2stages.get(keys)
if cfg and cfg.get("use_non_temporal_load") is not None:
nt = cfg["use_non_temporal_load"]
old_s1 = metadata.stage1
if hasattr(old_s1, 'func') and old_s1.func is not None:
if 'use_non_temporal_load' in (old_s1.keywords or {}):
new_kw = dict(old_s1.keywords)
new_kw['use_non_temporal_load'] = nt
metadata = _fused_moe_module.MOEMetadata(
functools.partial(old_s1.func, **{k: v for k, v in new_kw.items()}),
metadata.stage2, metadata.block_m, metadata.ksplit,
metadata.run_1stage, metadata.has_bias, nt,
)
old_s2 = metadata.stage2
if old_s2 and hasattr(old_s2, 'keywords') and 'use_non_temporal_load' in (old_s2.keywords or {}):
new_kw2 = dict(old_s2.keywords)
new_kw2['use_non_temporal_load'] = nt
metadata = _fused_moe_module.MOEMetadata(
metadata.stage1,
functools.partial(old_s2.func, **{k: v for k, v in new_kw2.items()}),
metadata.block_m, metadata.ksplit, metadata.run_1stage,
metadata.has_bias, nt,
)
# Hook for custom stage2 replacement (e.g., Gluon/Triton kernels)
custom_s2 = CUSTOM_STAGE2_REGISTRY.get((inter_dim, expert))
if custom_s2 is not None:
metadata = _fused_moe_module.MOEMetadata(
metadata.stage1, custom_s2,
metadata.block_m, metadata.ksplit, metadata.run_1stage,
metadata.has_bias, getattr(metadata, 'use_non_temporal_load', False),
)
return metadata
_fused_moe_module.get_2stage_cfgs = _patched_get_2stage_cfgs
# ---------- Native FP4 quant monkey-patch ----------
# Only applied by shapes that benefit ? see per-shape kernel.py files.
@triton.jit
def _patched_quant_kernel(
x_ptr,
x_fp4_ptr,
sorted_ids_ptr,
num_valid_ids_ptr,
blockscale_e8m0_sorted_ptr,
Mx,
Nx,
scaleNx,
stride_x_m,
stride_x_n,
stride_x_fp4_m,
stride_x_fp4_n,
stride_o3,
stride_o2,
stride_o1,
stride_o0,
stride_o4,
token_num,
M_i,
N_i,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
BLOCK_SIZE_Mx: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
TOPK: tl.constexpr,
):
pid = tl.program_id(0)
num_pid_x = tl.cdiv(Mx, BLOCK_SIZE_Mx) * scaleNx
stride_x_m = tl.cast(stride_x_m, tl.int64)
stride_x_n = tl.cast(stride_x_n, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n, tl.int64)
if pid < num_pid_x:
pid_m = pid // scaleNx
pid_n = pid % scaleNx
x_offs_m = pid_m * BLOCK_SIZE_Mx + tl.arange(0, BLOCK_SIZE_Mx)
x_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
x_mask = (x_offs_m < Mx)[:, None] & (x_offs_n < Nx)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask).to(tl.float32)
# Calculate scale
amax = tl.max(tl.abs(x), axis=1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
# Native FP4 convert via v_cvt_scalef32_pk_fp4_bf16
# amax is already power-of-2 (line 168). The hardware extracts the
# exponent field from hw_scale, so multiplying by 0.25 shifts the
# exponent by -2, equivalent to exp2(log2(amax).floor() - 2).
hw_scale = amax * 0.25
x_bf16 = x.to(tl.bfloat16)
x_pairs = x_bf16.reshape(BLOCK_SIZE_Mx, MXFP4_QUANT_BLOCK_SIZE // 2, 2)
x_lo, x_hi = tl.split(x_pairs)
x_lo_i32 = x_lo.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF
x_hi_i32 = (x_hi.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF) << 16
packed_bf16x2 = x_lo_i32 | x_hi_i32
hw_scale_bc = tl.broadcast_to(hw_scale, (BLOCK_SIZE_Mx, MXFP4_QUANT_BLOCK_SIZE // 2))
hw_result = tl.inline_asm_elementwise(
"v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",
"=v,v,v",
[packed_bf16x2, hw_scale_bc],
dtype=tl.int32,
is_pure=True,
pack=1,
)
out_tensor = (hw_result & 0xFF).to(tl.uint8)
out_offs_m = pid_m * BLOCK_SIZE_Mx + tl.arange(0, BLOCK_SIZE_Mx)
out_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(
0, MXFP4_QUANT_BLOCK_SIZE // 2
)
out_offs = (
out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
)
out_mask = (out_offs_m < Mx)[:, None] & (out_offs_n < (Nx // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)
return
# Second half: sorted scale store (unchanged from aiter)
pid -= num_pid_x
num_pid_n = tl.cdiv(N_i, BLOCK_SIZE_N * 2)
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
num_valid_ids = tl.load(num_valid_ids_ptr)
if pid_m * BLOCK_SIZE_M * 2 >= num_valid_ids:
return
stride_o0 = tl.cast(stride_o0, tl.int64)
stride_o1 = tl.cast(stride_o1, tl.int64)
stride_o2 = tl.cast(stride_o2, tl.int64)
stride_o3 = tl.cast(stride_o3, tl.int64)
stride_o4 = tl.cast(stride_o4, tl.int64)
BLOCK_SIZE_Nb: tl.constexpr = BLOCK_SIZE_N * 2 * MXFP4_QUANT_BLOCK_SIZE
sorted_ids_offs_m = pid_m * BLOCK_SIZE_M * 2 + tl.arange(0, BLOCK_SIZE_M * 2)
sorted_ids_offs = sorted_ids_offs_m
sorted_ids_mask = sorted_ids_offs_m < num_valid_ids
sorted_ids = tl.load(
sorted_ids_ptr + sorted_ids_offs,
mask=sorted_ids_mask,
other=token_num,
)
topk_ids = sorted_ids >> 24
sorted_ids = sorted_ids & 0xFFFFFF
if TOPK == 1:
x_offs_m = sorted_ids
else:
x_offs_m = sorted_ids * TOPK + topk_ids
x_offs_n = pid_n * BLOCK_SIZE_Nb + tl.arange(0, BLOCK_SIZE_Nb)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
x_mask = (sorted_ids < token_num)[:, None] & (x_offs_n < Nx)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask).to(tl.float32)
x = x.reshape(BLOCK_SIZE_M * 2, BLOCK_SIZE_N * 2, MXFP4_QUANT_BLOCK_SIZE)
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
bs_e8m0 = (
bs_e8m0.reshape(2, BLOCK_SIZE_M, 2, BLOCK_SIZE_N)
.permute(1, 3, 2, 0)
.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N, 4)
)
out = bs_e8m0
offs_0 = tl.arange(0, BLOCK_SIZE_M)
offs_1 = tl.arange(0, BLOCK_SIZE_N)
offs_2 = pid_n
offs_3 = pid_m
offs_4 = tl.arange(0, 4)
offs = (
offs_0[:, None, None] * stride_o0
+ offs_1[None, :, None] * stride_o1
+ offs_2 * stride_o2
+ offs_3 * stride_o3
+ offs_4[None, None, :] * stride_o4
)
tl.store(
blockscale_e8m0_sorted_ptr + offs,
out,
)
import aiter.ops.triton.quant.fused_mxfp4_quant as _quant_wrapper
_quant_wrapper._fused_dynamic_mxfp4_quant_moe_sort_kernel = _patched_quant_kernel
# ---------- End native FP4 quant monkey-patch ----------
def run_fused_moe(hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
topk_weights, topk_ids, gate_up_weight_scale_shuffled,
down_weight_scale_shuffled, config):
_test_gluon_mfma_scaled()
inject_configs()
hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
intermediate_pad = config["d_expert_pad"] - config["d_expert"]
return fused_moe(
hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
topk_weights, topk_ids,
expert_mask=None, activation=ActivationType.Silu,
quant_type=QuantType.per_1x32, doweight_stage1=False,
w1_scale=gate_up_weight_scale_shuffled,
w2_scale=down_weight_scale_shuffled,
a1_scale=None, a2_scale=None,
hidden_pad=hidden_pad, intermediate_pad=intermediate_pad,
)
# ============================================================
# bs=16, E=257, d_expert=256
# ============================================================
"""MoE bs16/e257_d256 ? v23: enable opus sorting by patching module var"""
import aiter.fused_moe as _fused_moe_module
# Enable opus sorting for all shapes (faster sorting kernel)
_fused_moe_module._USE_OPUS_MOE_SORTING = True
def _make_key_topk8(token, inter_dim, expert):
"""Config key matching actual benchmark topk=8"""
return (
256, token, 7168, inter_dim, expert, 8,
"ActivationType.Silu", "torch.bfloat16",
"torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",
"QuantType.per_1x32", True, False,
)
def _make_key_topk9(token, inter_dim, expert):
"""Config key for topk=9 (in case some benchmarks use topk=9)"""
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,
)
_cfg = {
"block_m": 32,
"ksplit": 7,
"kernelName1": "",
"kernelName2": "flydsl_moe2_afp4_wfp4_bf16_t32x128x128_atomic",
"run_1stage": False
}
# Register for both topk=8 and topk=9
CUSTOM_CONFIGS[_make_key_topk8(16, 256, 257)] = _cfg
CUSTOM_CONFIGS[_make_key_topk9(16, 256, 257)] = _cfg
_state_16_257_256 = None
def _run_16_257_256(hidden_states, gate_up_weight, down_weight,
gate_up_weight_scale, down_weight_scale,
gate_up_weight_shuffled, down_weight_shuffled,
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
topk_weights, topk_ids, config):
return run_fused_moe(
hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
topk_weights, topk_ids, gate_up_weight_scale_shuffled,
down_weight_scale_shuffled, config,
)
# ============================================================
# bs=128, E=257, d_expert=256
# ============================================================
"""MoE bs128/e257_d256 ? v14: cktile block_m=16, ksplit=4"""
# Register this shape's config
CUSTOM_CONFIGS[make_key(128, 256, 257)] = {
"block_m": 16,
"ksplit": 4,
"kernelName1": "",
"kernelName2": "",
"run_1stage": False
}
_state_128_257_256 = None
def _run_128_257_256(hidden_states, gate_up_weight, down_weight,
gate_up_weight_scale, down_weight_scale,
gate_up_weight_shuffled, down_weight_shuffled,
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
topk_weights, topk_ids, config):
return run_fused_moe(
hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
topk_weights, topk_ids, gate_up_weight_scale_shuffled,
down_weight_scale_shuffled, config,
)
# ============================================================
# bs=512, E=257, d_expert=256
# ============================================================
"""MoE bs512/e257_d256 ??? v162 config: block_m=64, ksplit=0"""
# Register this shape's config
CUSTOM_CONFIGS[make_key(512, 256, 257)] = {
"block_m": 32,
"ksplit": 0,
"kernelName1": "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
"kernelName2": "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic",
"run_1stage": False,
"use_non_temporal_load": True
}
_state_512_257_256 = None
def _run_512_257_256(hidden_states, gate_up_weight, down_weight,
gate_up_weight_scale, down_weight_scale,
gate_up_weight_shuffled, down_weight_shuffled,
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
topk_weights, topk_ids, config):
return run_fused_moe(
hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
topk_weights, topk_ids, gate_up_weight_scale_shuffled,
down_weight_scale_shuffled, config,
)
# ============================================================
# bs=16, E=33, d_expert=512
# ============================================================
"""MoE bs16/e33_d512 ? v162 config: block_m=32, ksplit=2"""
# Register this shape's config
CUSTOM_CONFIGS[make_key(16, 512, 33)] = {
"block_m": 32,
"ksplit": 2,
"kernelName1": "",
"kernelName2": "",
"run_1stage": False
}
_state_16_33_512 = None
def _run_16_33_512(hidden_states, gate_up_weight, down_weight,
gate_up_weight_scale, down_weight_scale,
gate_up_weight_shuffled, down_weight_shuffled,
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
topk_weights, topk_ids, config):
return run_fused_moe(
hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
topk_weights, topk_ids, gate_up_weight_scale_shuffled,
down_weight_scale_shuffled, config,
)
# ============================================================
# bs=128, E=33, d_expert=512
# ============================================================
"""MoE bs128/e33_d512 ? v162 config: block_m=64, ksplit=0"""
# Register this shape's config
CUSTOM_CONFIGS[make_key(128, 512, 33)] = {
"block_m": 64,
"ksplit": 0,
"kernelName1": "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
"kernelName2": "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic",
"run_1stage": False
}
_state_128_33_512 = None
def _run_128_33_512(hidden_states, gate_up_weight, down_weight,
gate_up_weight_scale, down_weight_scale,
gate_up_weight_shuffled, down_weight_shuffled,
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
topk_weights, topk_ids, config):
return run_fused_moe(
hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
topk_weights, topk_ids, gate_up_weight_scale_shuffled,
down_weight_scale_shuffled, config,
)
# ============================================================
# bs=512, E=33, d_expert=512
# ============================================================
"""MoE bs512/e33_d512 ? v162 config: block_m=64, ksplit=0"""
# Register this shape's config
CUSTOM_CONFIGS[make_key(512, 512, 33)] = {
"block_m": 64,
"ksplit": 0,
"kernelName1": "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
"kernelName2": "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic",
"run_1stage": False
}
_state_512_33_512 = None
def _run_512_33_512(hidden_states, gate_up_weight, down_weight,
gate_up_weight_scale, down_weight_scale,
gate_up_weight_shuffled, down_weight_shuffled,
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
topk_weights, topk_ids, config):
return run_fused_moe(
hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
topk_weights, topk_ids, gate_up_weight_scale_shuffled,
down_weight_scale_shuffled, config,
)
# ============================================================
# bs=512, E=33, d_expert=2048
# ============================================================
"""MoE bs512/e33_d2048 ? v162 config: block_m=64, ksplit=0"""
# Register this shape's config
CUSTOM_CONFIGS[make_key(512, 2048, 33)] = {
"block_m": 64,
"ksplit": 0,
"kernelName1": "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
"kernelName2": "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic",
"run_1stage": False
}
_state_512_33_2048 = None
def _run_512_33_2048(hidden_states, gate_up_weight, down_weight,
gate_up_weight_scale, down_weight_scale,
gate_up_weight_shuffled, down_weight_shuffled,
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
topk_weights, topk_ids, config):
return run_fused_moe(
hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
topk_weights, topk_ids, gate_up_weight_scale_shuffled,
down_weight_scale_shuffled, config,
)
def _run_default(hidden_states, gate_up_weight, down_weight,
gate_up_weight_scale, down_weight_scale,
gate_up_weight_shuffled, down_weight_shuffled,
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
topk_weights, topk_ids, config):
return run_fused_moe(hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
topk_weights, topk_ids, gate_up_weight_scale_shuffled,
down_weight_scale_shuffled, config)
from task import input_t, output_t
_DISPATCH = {
(16, 256, 1, 256): _run_16_257_256,
(128, 256, 1, 256): _run_128_257_256,
(512, 256, 1, 256): _run_512_257_256,
(16, 32, 1, 512): _run_16_33_512,
(128, 32, 1, 512): _run_128_33_512,
(512, 32, 1, 512): _run_512_33_512,
(512, 32, 1, 2048): _run_512_33_2048,
}
def custom_kernel(data: input_t) -> output_t:
(hidden_states, gate_up_weight, down_weight,
gate_up_weight_scale, down_weight_scale,
gate_up_weight_shuffled, down_weight_shuffled,
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
topk_weights, topk_ids, config) = data
bs = config["bs"]
nr = config["n_routed_experts"]
ns = config["n_shared_experts"]
d = config["d_expert"]
fn = _DISPATCH.get((bs, nr, ns, d), _run_default)
return fn(hidden_states, gate_up_weight, down_weight,
gate_up_weight_scale, down_weight_scale,
gate_up_weight_shuffled, down_weight_shuffled,
gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
topk_weights, topk_ids, config)
scrolls · 579 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 598451.
#!POPCORN leaderboard amd-moe-mxfp4- #!POPCORN gpu MI355X+ """Auto-generated by build.py"""- """- v168: Pre-allocate quantization output buffers (x_fp4, blockscale_e8m0_sorted)- for both stage1 and stage2 quant calls. Inline the fused_dynamic_mxfp4_quant_moe_sort- Triton kernel launch with cached output tensors to eliminate 4 torch.empty allocations- per forward pass on CK 2-stage shapes.- """import os+ import sysimport functoolsimport torchimport triton- from typing import Dict, Tuple, Optional- from task import input_t, output_t+ import triton.language as tl- 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,- )+ try:+ import triton.experimental.gluon as gluon+ import triton.experimental.gluon.language as gl+ _HAS_GLUON = True+ except Exception:+ _HAS_GLUON = False++ def _test_gluon_mfma_scaled():+ pass+ from aiter import ActivationType, QuantType+ from aiter.fused_moe import fused_moeimport aiter.fused_moe as _fused_moe_moduleimport 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,- )- from aiter.utility import fp4_utils- # Register FlyDSL tile_k=128 kernels+ # Register FlyDSL tile_k=128 kernels not in server defaults_flydsl_moe_kernels._KERNEL_PARAMS["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,⋯ 11 unchanged lines"tile_m": 16, "tile_n": 128, "tile_k": 128, "mode": "atomic", "MPerBlock": 16,}- # Shape configs- _CUSTOM_CONFIGS = {}+ # Collected configs from all shapes ? each kernel.py adds its own+ CUSTOM_CONFIGS = {}- def _make_key(token, inter_dim, expert):+ # Registry for custom stage2 functions: {(inter_dim, expert): callable}+ # kernel.py files register their custom stage2 here before fused_moe runs.+ CUSTOM_STAGE2_REGISTRY = {}++ _injected = False++ def make_key(token, inter_dim, expert):return (256, token, 7168, inter_dim, expert, 9,"ActivationType.Silu", "torch.bfloat16",⋯ 1 unchanged lines"QuantType.per_1x32", True, False,)- _4WG_STAGE1_M128 = "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"- _4WG_STAGE1_M32 = "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"- _FLYDSL_STAGE2_M16_K128 = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"-- # E=33 shapes- _CUSTOM_CONFIGS[_make_key(16, 512, 33)] = {- "block_m": 32, "ksplit": 2, "kernelName1": "", "kernelName2": "",- "run_1stage": False,- }- _CUSTOM_CONFIGS[_make_key(128, 512, 33)] = {- "block_m": 64, "ksplit": 0,- "kernelName1": _4WG_STAGE1_M128, "kernelName2": _FLYDSL_STAGE2_M16_K128,- "run_1stage": False,- }- _CUSTOM_CONFIGS[_make_key(512, 512, 33)] = {- "block_m": 64, "ksplit": 0,- "kernelName1": _4WG_STAGE1_M128, "kernelName2": _FLYDSL_STAGE2_M16_K128,- "run_1stage": False,- }- _CUSTOM_CONFIGS[_make_key(512, 2048, 33)] = {- "block_m": 64, "ksplit": 0,- "kernelName1": _4WG_STAGE1_M128, "kernelName2": _FLYDSL_STAGE2_M16_K128,- "run_1stage": False,- }-- # E=257 shapes- _CUSTOM_CONFIGS[_make_key(16, 256, 257)] = {- "block_m": 16, "ksplit": 2, "kernelName1": "", "kernelName2": "",- "run_1stage": False,- }- _CUSTOM_CONFIGS[_make_key(128, 256, 257)] = {- "block_m": 16, "ksplit": 2, "kernelName1": "", "kernelName2": "",- "run_1stage": False,- }- _CUSTOM_CONFIGS[_make_key(512, 256, 257)] = {- "block_m": 32, "ksplit": 0,- "kernelName1": _4WG_STAGE1_M32, "kernelName2": _FLYDSL_STAGE2_M16_K128,- "run_1stage": False,- "use_non_temporal_load": True,- }-- # Pre-allocated buffer cache- _buffer_cache = {}-- def _get_or_alloc_sorting_buffers(M, E, topk, model_dim, block_size_M, device):- """Pre-allocate moe_sorting output buffers."""- key = ("sort", M, E, topk, model_dim, block_size_M)- if key in _buffer_cache:- return _buffer_cache[key]-- max_num_tokens_padded = int(M * topk + E * block_size_M - topk)- max_num_m_blocks = int((max_num_tokens_padded + block_size_M - 1) // block_size_M)-- bufs = {- "sorted_ids": torch.empty(max_num_tokens_padded, dtype=dtypes.i32, device=device),- "sorted_weights": torch.empty(max_num_tokens_padded, dtype=dtypes.fp32, device=device),- "sorted_expert_ids": torch.empty(max_num_m_blocks, dtype=dtypes.i32, device=device),- "num_valid_ids": torch.empty(2, dtype=dtypes.i32, device=device),- "moe_buf": torch.empty((M, model_dim), dtype=torch.bfloat16, device=device),- }- _buffer_cache[key] = bufs- return bufs-- def _get_or_alloc_a2(M, topk, inter_dim, device):- """Pre-allocate a2 intermediate buffer."""- key = ("a2", M, topk, inter_dim)- if key in _buffer_cache:- return _buffer_cache[key]- buf = torch.empty((M, topk, inter_dim), dtype=torch.bfloat16, device=device)- _buffer_cache[key] = buf- return buf-- def _get_or_alloc_quant_buffers(M, N, sorted_ids_len, topk, device):- """Pre-allocate quantization output buffers for fused_dynamic_mxfp4_quant_moe_sort."""- MXFP4_QUANT_BLOCK_SIZE = 32- BLOCK_SIZE_M, BLOCK_SIZE_N = 32, 8- BLOCK_SIZE_M_u32, BLOCK_SIZE_N_u32 = 16, 4-- key = ("quant", M, N, sorted_ids_len, topk)- if key in _buffer_cache:- return _buffer_cache[key]-- x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=device)- scaleN = triton.cdiv(N, MXFP4_QUANT_BLOCK_SIZE)- M_o = sorted_ids_len- N_o = scaleN-- blockscale_e8m0_sorted = torch.empty(- (- triton.cdiv(M_o, BLOCK_SIZE_M),- triton.cdiv(N_o, BLOCK_SIZE_N),- BLOCK_SIZE_N_u32,- BLOCK_SIZE_M_u32,- 4,- ),- dtype=torch.uint8,- device=device,- )-- bufs = {"x_fp4": x_fp4, "blockscale": blockscale_e8m0_sorted}- _buffer_cache[key] = bufs- return bufs-- def _quant_prealloc(x, sorted_ids, num_valid_ids, token_num, topk, block_size, device):- """Inline fused_dynamic_mxfp4_quant_moe_sort with pre-allocated output buffers."""- M, N = x.shape- MXFP4_QUANT_BLOCK_SIZE = 32- BLOCK_SIZE_Mx = 128- BLOCK_SIZE_M, BLOCK_SIZE_N = 32, 8-- scaleN = triton.cdiv(N, MXFP4_QUANT_BLOCK_SIZE)- M_i, N_i = M, scaleN- M_o = sorted_ids.shape[0]-- # Get pre-allocated buffers- qbufs = _get_or_alloc_quant_buffers(M, N, M_o, topk, device)- x_fp4 = qbufs["x_fp4"]- blockscale_e8m0_sorted = qbufs["blockscale"]-- num_pid = triton.cdiv(M, BLOCK_SIZE_Mx) * scaleN + triton.cdiv(- M_o, BLOCK_SIZE_M- ) * triton.cdiv(N_i, BLOCK_SIZE_N)-- _fused_dynamic_mxfp4_quant_moe_sort_kernel[(num_pid,)](- x,- x_fp4,- sorted_ids,- num_valid_ids,- blockscale_e8m0_sorted,- M,- N,- scaleN,- *x.stride(),- *x_fp4.stride(),- *blockscale_e8m0_sorted.stride(),- token_num=token_num,- M_i=M_i,- N_i=N_i,- MXFP4_QUANT_BLOCK_SIZE=MXFP4_QUANT_BLOCK_SIZE,- BLOCK_SIZE_Mx=BLOCK_SIZE_Mx,- BLOCK_SIZE_M=BLOCK_SIZE_M // 2,- BLOCK_SIZE_N=BLOCK_SIZE_N // 2,- TOPK=topk,- )-- return (- x_fp4.view(dtypes.fp4x2),- blockscale_e8m0_sorted.view(dtypes.fp8_e8m0).view(-1, scaleN),- )--- _injected = False-- def _inject_configs():+ def inject_configs():global _injectedif _injected:return⋯ 16 unchanged lineselse:_fused_moe_module.cfg_2stages = {}- _fused_moe_module.cfg_2stages.update(_CUSTOM_CONFIGS)+ _fused_moe_module.cfg_2stages.update(CUSTOM_CONFIGS)- # Monkeypatch get_2stage_cfgs to support use_non_temporal_load from config_original_get_2stage_cfgs = _fused_moe_module.get_2stage_cfgs@functools.lru_cache(maxsize=2048)⋯ 24 unchanged linesnew_kw['use_non_temporal_load'] = ntmetadata = _fused_moe_module.MOEMetadata(functools.partial(old_s1.func, **{k: v for k, v in new_kw.items()}),- metadata.stage2,- metadata.block_m,- metadata.ksplit,- metadata.run_1stage,- metadata.has_bias,- nt,+ metadata.stage2, metadata.block_m, metadata.ksplit,+ metadata.run_1stage, metadata.has_bias, nt,)old_s2 = metadata.stage2if old_s2 and hasattr(old_s2, 'keywords') and 'use_non_temporal_load' in (old_s2.keywords or {}):⋯ 2 unchanged linesmetadata = _fused_moe_module.MOEMetadata(metadata.stage1,functools.partial(old_s2.func, **{k: v for k, v in new_kw2.items()}),- metadata.block_m,- metadata.ksplit,- metadata.run_1stage,- metadata.has_bias,- nt,+ metadata.block_m, metadata.ksplit, metadata.run_1stage,+ metadata.has_bias, nt,)+ # Hook for custom stage2 replacement (e.g., Gluon/Triton kernels)+ custom_s2 = CUSTOM_STAGE2_REGISTRY.get((inter_dim, expert))+ if custom_s2 is not None:+ metadata = _fused_moe_module.MOEMetadata(+ metadata.stage1, custom_s2,+ metadata.block_m, metadata.ksplit, metadata.run_1stage,+ metadata.has_bias, getattr(metadata, 'use_non_temporal_load', False),+ )+return metadata_fused_moe_module.get_2stage_cfgs = _patched_get_2stage_cfgs- def custom_kernel(data: input_t) -> output_t:- (- hidden_states, gate_up_weight, down_weight,+ # ---------- Native FP4 quant monkey-patch ----------+ # Only applied by shapes that benefit ? see per-shape kernel.py files.++ @triton.jit+ def _patched_quant_kernel(+ x_ptr,+ x_fp4_ptr,+ sorted_ids_ptr,+ num_valid_ids_ptr,+ blockscale_e8m0_sorted_ptr,+ Mx,+ Nx,+ scaleNx,+ stride_x_m,+ stride_x_n,+ stride_x_fp4_m,+ stride_x_fp4_n,+ stride_o3,+ stride_o2,+ stride_o1,+ stride_o0,+ stride_o4,+ token_num,+ M_i,+ N_i,+ MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,+ BLOCK_SIZE_Mx: tl.constexpr,+ BLOCK_SIZE_M: tl.constexpr,+ BLOCK_SIZE_N: tl.constexpr,+ TOPK: tl.constexpr,+ ):+ pid = tl.program_id(0)+ num_pid_x = tl.cdiv(Mx, BLOCK_SIZE_Mx) * scaleNx++ stride_x_m = tl.cast(stride_x_m, tl.int64)+ stride_x_n = tl.cast(stride_x_n, tl.int64)+ stride_x_fp4_m = tl.cast(stride_x_fp4_m, tl.int64)+ stride_x_fp4_n = tl.cast(stride_x_fp4_n, tl.int64)++ if pid < num_pid_x:+ pid_m = pid // scaleNx+ pid_n = pid % scaleNx++ x_offs_m = pid_m * BLOCK_SIZE_Mx + tl.arange(0, BLOCK_SIZE_Mx)+ x_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE)+ x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n+ x_mask = (x_offs_m < Mx)[:, None] & (x_offs_n < Nx)[None, :]+ x = tl.load(x_ptr + x_offs, mask=x_mask).to(tl.float32)++ # Calculate scale+ amax = tl.max(tl.abs(x), axis=1, keep_dims=True)+ amax = amax.to(tl.int32, bitcast=True)+ amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000+ amax = amax.to(tl.float32, bitcast=True)+ # Native FP4 convert via v_cvt_scalef32_pk_fp4_bf16+ # amax is already power-of-2 (line 168). The hardware extracts the+ # exponent field from hw_scale, so multiplying by 0.25 shifts the+ # exponent by -2, equivalent to exp2(log2(amax).floor() - 2).+ hw_scale = amax * 0.25+ x_bf16 = x.to(tl.bfloat16)+ x_pairs = x_bf16.reshape(BLOCK_SIZE_Mx, MXFP4_QUANT_BLOCK_SIZE // 2, 2)+ x_lo, x_hi = tl.split(x_pairs)+ x_lo_i32 = x_lo.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF+ x_hi_i32 = (x_hi.to(tl.int16, bitcast=True).to(tl.int32) & 0xFFFF) << 16+ packed_bf16x2 = x_lo_i32 | x_hi_i32++ hw_scale_bc = tl.broadcast_to(hw_scale, (BLOCK_SIZE_Mx, MXFP4_QUANT_BLOCK_SIZE // 2))++ hw_result = tl.inline_asm_elementwise(+ "v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",+ "=v,v,v",+ [packed_bf16x2, hw_scale_bc],+ dtype=tl.int32,+ is_pure=True,+ pack=1,+ )+ out_tensor = (hw_result & 0xFF).to(tl.uint8)++ out_offs_m = pid_m * BLOCK_SIZE_Mx + tl.arange(0, BLOCK_SIZE_Mx)+ out_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(+ 0, MXFP4_QUANT_BLOCK_SIZE // 2+ )+ out_offs = (+ out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n+ )+ out_mask = (out_offs_m < Mx)[:, None] & (out_offs_n < (Nx // 2))[None, :]+ tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)++ return++ # Second half: sorted scale store (unchanged from aiter)+ pid -= num_pid_x+ num_pid_n = tl.cdiv(N_i, BLOCK_SIZE_N * 2)+ pid_m = pid // num_pid_n+ pid_n = pid % num_pid_n+ num_valid_ids = tl.load(num_valid_ids_ptr)+ if pid_m * BLOCK_SIZE_M * 2 >= num_valid_ids:+ return+ stride_o0 = tl.cast(stride_o0, tl.int64)+ stride_o1 = tl.cast(stride_o1, tl.int64)+ stride_o2 = tl.cast(stride_o2, tl.int64)+ stride_o3 = tl.cast(stride_o3, tl.int64)+ stride_o4 = tl.cast(stride_o4, tl.int64)++ BLOCK_SIZE_Nb: tl.constexpr = BLOCK_SIZE_N * 2 * MXFP4_QUANT_BLOCK_SIZE+ sorted_ids_offs_m = pid_m * BLOCK_SIZE_M * 2 + tl.arange(0, BLOCK_SIZE_M * 2)+ sorted_ids_offs = sorted_ids_offs_m+ sorted_ids_mask = sorted_ids_offs_m < num_valid_ids+ sorted_ids = tl.load(+ sorted_ids_ptr + sorted_ids_offs,+ mask=sorted_ids_mask,+ other=token_num,+ )+ topk_ids = sorted_ids >> 24+ sorted_ids = sorted_ids & 0xFFFFFF+ if TOPK == 1:+ x_offs_m = sorted_ids+ else:+ x_offs_m = sorted_ids * TOPK + topk_ids+ x_offs_n = pid_n * BLOCK_SIZE_Nb + tl.arange(0, BLOCK_SIZE_Nb)+ x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n+ x_mask = (sorted_ids < token_num)[:, None] & (x_offs_n < Nx)[None, :]+ x = tl.load(x_ptr + x_offs, mask=x_mask).to(tl.float32)+ x = x.reshape(BLOCK_SIZE_M * 2, BLOCK_SIZE_N * 2, MXFP4_QUANT_BLOCK_SIZE)++ amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)+ amax = amax.to(tl.int32, bitcast=True)+ amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000+ amax = amax.to(tl.float32, bitcast=True)+ scale_e8m0_unbiased = tl.log2(amax).floor() - 2+ scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)+ bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127+ bs_e8m0 = (+ bs_e8m0.reshape(2, BLOCK_SIZE_M, 2, BLOCK_SIZE_N)+ .permute(1, 3, 2, 0)+ .reshape(BLOCK_SIZE_M, BLOCK_SIZE_N, 4)+ )+ out = bs_e8m0++ offs_0 = tl.arange(0, BLOCK_SIZE_M)+ offs_1 = tl.arange(0, BLOCK_SIZE_N)+ offs_2 = pid_n+ offs_3 = pid_m+ offs_4 = tl.arange(0, 4)+ offs = (+ offs_0[:, None, None] * stride_o0+ + offs_1[None, :, None] * stride_o1+ + offs_2 * stride_o2+ + offs_3 * stride_o3+ + offs_4[None, None, :] * stride_o4+ )+ tl.store(+ blockscale_e8m0_sorted_ptr + offs,+ out,+ )++ import aiter.ops.triton.quant.fused_mxfp4_quant as _quant_wrapper+ _quant_wrapper._fused_dynamic_mxfp4_quant_moe_sort_kernel = _patched_quant_kernel++ # ---------- End native FP4 quant monkey-patch ----------+++ def run_fused_moe(hidden_states, gate_up_weight_shuffled, down_weight_shuffled,+ topk_weights, topk_ids, gate_up_weight_scale_shuffled,+ down_weight_scale_shuffled, config):+ _test_gluon_mfma_scaled()+ inject_configs()+ hidden_pad = config["d_hidden_pad"] - config["d_hidden"]+ intermediate_pad = config["d_expert_pad"] - config["d_expert"]+ return fused_moe(+ hidden_states, gate_up_weight_shuffled, down_weight_shuffled,+ topk_weights, topk_ids,+ expert_mask=None, activation=ActivationType.Silu,+ quant_type=QuantType.per_1x32, doweight_stage1=False,+ w1_scale=gate_up_weight_scale_shuffled,+ w2_scale=down_weight_scale_shuffled,+ a1_scale=None, a2_scale=None,+ hidden_pad=hidden_pad, intermediate_pad=intermediate_pad,+ )++ # ============================================================+ # bs=16, E=257, d_expert=256+ # ============================================================++ """MoE bs16/e257_d256 ? v23: enable opus sorting by patching module var"""+ import aiter.fused_moe as _fused_moe_module++ # Enable opus sorting for all shapes (faster sorting kernel)+ _fused_moe_module._USE_OPUS_MOE_SORTING = True++ def _make_key_topk8(token, inter_dim, expert):+ """Config key matching actual benchmark topk=8"""+ return (+ 256, token, 7168, inter_dim, expert, 8,+ "ActivationType.Silu", "torch.bfloat16",+ "torch.float4_e2m1fn_x2", "torch.float4_e2m1fn_x2",+ "QuantType.per_1x32", True, False,+ )++ def _make_key_topk9(token, inter_dim, expert):+ """Config key for topk=9 (in case some benchmarks use topk=9)"""+ 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,+ )++ _cfg = {+ "block_m": 32,+ "ksplit": 7,+ "kernelName1": "",+ "kernelName2": "flydsl_moe2_afp4_wfp4_bf16_t32x128x128_atomic",+ "run_1stage": False+ }++ # Register for both topk=8 and topk=9+ CUSTOM_CONFIGS[_make_key_topk8(16, 256, 257)] = _cfg+ CUSTOM_CONFIGS[_make_key_topk9(16, 256, 257)] = _cfg++ _state_16_257_256 = None++ def _run_16_257_256(hidden_states, gate_up_weight, down_weight,gate_up_weight_scale, down_weight_scale,gate_up_weight_shuffled, down_weight_shuffled,gate_up_weight_scale_shuffled, down_weight_scale_shuffled,- topk_weights, topk_ids, config,- ) = data+ topk_weights, topk_ids, config):+ return run_fused_moe(+ hidden_states, gate_up_weight_shuffled, down_weight_shuffled,+ topk_weights, topk_ids, gate_up_weight_scale_shuffled,+ down_weight_scale_shuffled, config,+ )- _inject_configs()+ # ============================================================+ # bs=128, E=257, d_expert=256+ # ============================================================- hidden_pad = config["d_hidden_pad"] - config["d_hidden"]- intermediate_pad = config["d_expert_pad"] - config["d_expert"]+ """MoE bs128/e257_d256 ? v14: cktile block_m=16, ksplit=4"""- M = hidden_states.shape[0]- topk = topk_ids.shape[1]- device = topk_ids.device- w1 = gate_up_weight_shuffled- w2 = down_weight_shuffled- E, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)+ # Register this shape's config+ CUSTOM_CONFIGS[make_key(128, 256, 257)] = {+ "block_m": 16,+ "ksplit": 4,+ "kernelName1": "",+ "kernelName2": "",+ "run_1stage": False+ }- 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,+ _state_128_257_256 = None++ def _run_128_257_256(hidden_states, gate_up_weight, down_weight,+ gate_up_weight_scale, down_weight_scale,+ gate_up_weight_shuffled, down_weight_shuffled,+ gate_up_weight_scale_shuffled, down_weight_scale_shuffled,+ topk_weights, topk_ids, config):+ return run_fused_moe(+ hidden_states, gate_up_weight_shuffled, down_weight_shuffled,+ topk_weights, topk_ids, gate_up_weight_scale_shuffled,+ down_weight_scale_shuffled, config,)- block_size_M = int(metadata.block_m)+ # ============================================================+ # bs=512, E=257, d_expert=256+ # ============================================================- # === Pre-allocated moe_sorting ===- bufs = _get_or_alloc_sorting_buffers(M, E, topk, model_dim, block_size_M, device)- sorted_ids = bufs["sorted_ids"]- sorted_weights = bufs["sorted_weights"]- sorted_expert_ids = bufs["sorted_expert_ids"]- num_valid_ids = bufs["num_valid_ids"]- moe_out = bufs["moe_buf"]+ """MoE bs512/e257_d256 ??? v162 config: block_m=64, ksplit=0"""- aiter.moe_sorting_fwd(- topk_ids, topk_weights,- sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_out,- E, int(block_size_M), None, None, 0,+ # Register this shape's config+ CUSTOM_CONFIGS[make_key(512, 256, 257)] = {+ "block_m": 32,+ "ksplit": 0,+ "kernelName1": "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",+ "kernelName2": "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic",+ "run_1stage": False,+ "use_non_temporal_load": True+ }++ _state_512_257_256 = None++ def _run_512_257_256(hidden_states, gate_up_weight, down_weight,+ gate_up_weight_scale, down_weight_scale,+ gate_up_weight_shuffled, down_weight_shuffled,+ gate_up_weight_scale_shuffled, down_weight_scale_shuffled,+ topk_weights, topk_ids, config):+ return run_fused_moe(+ hidden_states, gate_up_weight_shuffled, down_weight_shuffled,+ topk_weights, topk_ids, gate_up_weight_scale_shuffled,+ down_weight_scale_shuffled, config,)- # === Inline 2-stage pipeline ===- token_num = M+ # ============================================================+ # bs=16, E=33, d_expert=512+ # ============================================================- if metadata.ksplit > 1:- # cktile_moe path: bf16 activations, no fp4 quant- a1 = hidden_states.to(torch.bfloat16)- a1_scale = None- w1_scale_view = gate_up_weight_scale_shuffled.view(dtypes.fp8_e8m0)- w2_scale_view = down_weight_scale_shuffled.view(dtypes.fp8_e8m0)+ """MoE bs16/e33_d512 ? v162 config: block_m=32, ksplit=2"""- a2 = metadata.stage1(- a1, w1, w2,- sorted_ids, sorted_expert_ids, num_valid_ids,- _get_or_alloc_a2(M, topk, inter_dim, device), # pre-allocated- topk,- block_m=block_size_M,- a1_scale=a1_scale,- w1_scale=w1_scale_view,- sorted_weights=None,- )+ # Register this shape's config+ CUSTOM_CONFIGS[make_key(16, 512, 33)] = {+ "block_m": 32,+ "ksplit": 2,+ "kernelName1": "",+ "kernelName2": "",+ "run_1stage": False+ }- # cktile_moe stage2: a2 is bf16, no inter-stage requant- a2_scale = None- metadata.stage2(- a2, w1, w2,- sorted_ids, sorted_expert_ids, num_valid_ids,- moe_out, topk,- w2_scale=w2_scale_view,- a2_scale=a2_scale,- block_m=block_size_M,- sorted_weights=sorted_weights,- )- else:- # CK 2-stage path: fp4 activation quant with pre-allocated buffers- w1_scale_view = gate_up_weight_scale_shuffled.view(dtypes.fp8_e8m0)- w2_scale_view = down_weight_scale_shuffled.view(dtypes.fp8_e8m0)+ _state_16_33_512 = None- # Stage 1: quant activations + gate_up GEMM + SwiGLU- a1, a1_scale = _quant_prealloc(- hidden_states, sorted_ids, num_valid_ids,- token_num, 1, block_size_M, device,- )+ def _run_16_33_512(hidden_states, gate_up_weight, down_weight,+ gate_up_weight_scale, down_weight_scale,+ gate_up_weight_shuffled, down_weight_shuffled,+ gate_up_weight_scale_shuffled, down_weight_scale_shuffled,+ topk_weights, topk_ids, config):+ return run_fused_moe(+ hidden_states, gate_up_weight_shuffled, down_weight_shuffled,+ topk_weights, topk_ids, gate_up_weight_scale_shuffled,+ down_weight_scale_shuffled, config,+ )- a2 = _get_or_alloc_a2(M, topk, inter_dim, device)- a2 = metadata.stage1(- a1, w1, w2,- sorted_ids, sorted_expert_ids, num_valid_ids,- a2, topk,- block_m=block_size_M,- a1_scale=a1_scale,- w1_scale=w1_scale_view,- sorted_weights=None,- )+ # ============================================================+ # bs=128, E=33, d_expert=512+ # ============================================================- # Inter-stage requant: bf16 -> fp4 with pre-allocated buffers- a2_flat = a2.view(-1, inter_dim)- a2_quant, a2_scale = _quant_prealloc(- a2_flat, sorted_ids, num_valid_ids,- token_num, topk, block_size_M, device,- )- a2_quant = a2_quant.view(token_num, topk, -1)+ """MoE bs128/e33_d512 ? v162 config: block_m=64, ksplit=0"""- # Stage 2: down GEMM + weighted reduction- metadata.stage2(- a2_quant, w1, w2,- sorted_ids, sorted_expert_ids, num_valid_ids,- moe_out, topk,- w2_scale=w2_scale_view,- a2_scale=a2_scale,- block_m=block_size_M,- sorted_weights=sorted_weights,- )+ # Register this shape's config+ CUSTOM_CONFIGS[make_key(128, 512, 33)] = {+ "block_m": 64,+ "ksplit": 0,+ "kernelName1": "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",+ "kernelName2": "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic",+ "run_1stage": False+ }- return moe_out+ _state_128_33_512 = None++ def _run_128_33_512(hidden_states, gate_up_weight, down_weight,+ gate_up_weight_scale, down_weight_scale,+ gate_up_weight_shuffled, down_weight_shuffled,+ gate_up_weight_scale_shuffled, down_weight_scale_shuffled,+ topk_weights, topk_ids, config):+ return run_fused_moe(+ hidden_states, gate_up_weight_shuffled, down_weight_shuffled,+ topk_weights, topk_ids, gate_up_weight_scale_shuffled,+ down_weight_scale_shuffled, config,+ )++ # ============================================================+ # bs=512, E=33, d_expert=512+ # ============================================================++ """MoE bs512/e33_d512 ? v162 config: block_m=64, ksplit=0"""++ # Register this shape's config+ CUSTOM_CONFIGS[make_key(512, 512, 33)] = {+ "block_m": 64,+ "ksplit": 0,+ "kernelName1": "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",+ "kernelName2": "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic",+ "run_1stage": False+ }++ _state_512_33_512 = None++ def _run_512_33_512(hidden_states, gate_up_weight, down_weight,+ gate_up_weight_scale, down_weight_scale,+ gate_up_weight_shuffled, down_weight_shuffled,+ gate_up_weight_scale_shuffled, down_weight_scale_shuffled,+ topk_weights, topk_ids, config):+ return run_fused_moe(+ hidden_states, gate_up_weight_shuffled, down_weight_shuffled,+ topk_weights, topk_ids, gate_up_weight_scale_shuffled,+ down_weight_scale_shuffled, config,+ )++ # ============================================================+ # bs=512, E=33, d_expert=2048+ # ============================================================++ """MoE bs512/e33_d2048 ? v162 config: block_m=64, ksplit=0"""++ # Register this shape's config+ CUSTOM_CONFIGS[make_key(512, 2048, 33)] = {+ "block_m": 64,+ "ksplit": 0,+ "kernelName1": "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",+ "kernelName2": "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic",+ "run_1stage": False+ }++ _state_512_33_2048 = None++ def _run_512_33_2048(hidden_states, gate_up_weight, down_weight,+ gate_up_weight_scale, down_weight_scale,+ gate_up_weight_shuffled, down_weight_shuffled,+ gate_up_weight_scale_shuffled, down_weight_scale_shuffled,+ topk_weights, topk_ids, config):+ return run_fused_moe(+ hidden_states, gate_up_weight_shuffled, down_weight_shuffled,+ topk_weights, topk_ids, gate_up_weight_scale_shuffled,+ down_weight_scale_shuffled, config,+ )++ def _run_default(hidden_states, gate_up_weight, down_weight,+ gate_up_weight_scale, down_weight_scale,+ gate_up_weight_shuffled, down_weight_shuffled,+ gate_up_weight_scale_shuffled, down_weight_scale_shuffled,+ topk_weights, topk_ids, config):+ return run_fused_moe(hidden_states, gate_up_weight_shuffled, down_weight_shuffled,+ topk_weights, topk_ids, gate_up_weight_scale_shuffled,+ down_weight_scale_shuffled, config)++ from task import input_t, output_t+ _DISPATCH = {+ (16, 256, 1, 256): _run_16_257_256,+ (128, 256, 1, 256): _run_128_257_256,+ (512, 256, 1, 256): _run_512_257_256,+ (16, 32, 1, 512): _run_16_33_512,+ (128, 32, 1, 512): _run_128_33_512,+ (512, 32, 1, 512): _run_512_33_512,+ (512, 32, 1, 2048): _run_512_33_2048,+ }+ def custom_kernel(data: input_t) -> output_t:+ (hidden_states, gate_up_weight, down_weight,+ gate_up_weight_scale, down_weight_scale,+ gate_up_weight_shuffled, down_weight_shuffled,+ gate_up_weight_scale_shuffled, down_weight_scale_shuffled,+ topk_weights, topk_ids, config) = data+ bs = config["bs"]+ nr = config["n_routed_experts"]+ ns = config["n_shared_experts"]+ d = config["d_expert"]+ fn = _DISPATCH.get((bs, nr, ns, d), _run_default)+ return fn(hidden_states, gate_up_weight, down_weight,+ gate_up_weight_scale, down_weight_scale,+ gate_up_weight_shuffled, down_weight_shuffled,+ gate_up_weight_scale_shuffled, down_weight_scale_shuffled,+ topk_weights, topk_ids, config)
scrolls · 818 diff lines total
Best evidence level for this revision: reported
JSON