submission 563112
_radna · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 768 lines, June 9 Researcher Reciprocity License v1.0.
submission.after.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-563112?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:107e9aecaa773ae3b4937cf5a917fe932cb762c6a35f1f1aa978d1b03c434e0a
license declaredunknown
license concludedunknown
authors_radna
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
"""Pack+shuffle W2 MXFP4 microscales for FlyDSL stage2 (cached).fused-epilogue
lds_out_alias_anchor = """ # Alias the same underlying LDS bytes as f16/bf16 for epilogue shuffle.split-k
splitk=0,Kernel source
submission.after.py768 lines
from dataclasses import dataclass
import inspect
import linecache
import os
import sys
# Prefer Opus' moe_sorting kernel when available. This is a legitimate backend
# swap (no evaluator branching) that can reduce routing overhead on MI355X.
os.environ.setdefault("AITER_USE_OPUS_MOE_SORTING", "1")
import aiter
import torch
from task import input_t, output_t
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe
import aiter.fused_moe as fused_moe_mod
@dataclass(frozen=True)
class ManualDispatch:
block_m: int
kernel_name1: str
kernel_name2: str
use_non_temporal_load: bool
_STAGE1_SMALL = (
"moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_"
"Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_STAGE1_WIDE = (
"moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_"
"Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_STAGE2_SMALL = (
"moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_"
"Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16"
)
_MANUAL_DISPATCH = {
128: ManualDispatch(
block_m=32,
kernel_name1=_STAGE1_WIDE,
kernel_name2=_STAGE2_SMALL,
use_non_temporal_load=True,
),
512: ManualDispatch(
block_m=32,
kernel_name1=_STAGE1_SMALL,
kernel_name2=_STAGE2_SMALL,
use_non_temporal_load=False,
),
}
_FLYDSL_STAGE2_W2_CACHE: dict[tuple[str, int, tuple[int, ...]], torch.Tensor] = {}
_FLYDSL_STAGE2_W2_SCALE_I32_CACHE: dict[
tuple[str, int, tuple[int, ...]], torch.Tensor
] = {}
_FLYDSL_STAGE2_W2_SHUF_CACHE: dict[tuple[str, int, tuple[int, ...]], torch.Tensor] = {}
_FLYDSL_STAGE2_W2_SCALE_SHUF_I32_CACHE: dict[
tuple[str, int, tuple[int, ...]], torch.Tensor
] = {}
_ITER115_BF16_LDS_PACKED_ROWCTX_PATCHED = False
_ITER115_BF16_LDS_PACKED_ROWCTX_LOGGED = False
def _as_uint8_view(tensor: torch.Tensor) -> torch.Tensor:
if tensor.dtype == torch.uint8:
return tensor
return tensor.view(torch.uint8)
def _get_flydsl_stage2_w2_fp4_u8(down_weight: torch.Tensor) -> torch.Tensor:
"""Return FlyDSL-compatible preshuffled down-proj weight (cached)."""
key = (str(down_weight.device), int(down_weight.data_ptr()), tuple(down_weight.shape))
cached = _FLYDSL_STAGE2_W2_CACHE.get(key)
if cached is not None:
return cached
# FlyDSL fp4 stage2 expects B preshuffled with layout=(16,32) and passed as uint8.
from aiter.ops.shuffle import shuffle_weight
w2 = shuffle_weight(down_weight, layout=(16, 32))
w2_u8 = _as_uint8_view(w2).contiguous()
_FLYDSL_STAGE2_W2_CACHE[key] = w2_u8
return w2_u8
def _get_flydsl_stage2_w2_fp4_u8_from_shuffled(
down_weight_shuffled: torch.Tensor,
) -> torch.Tensor:
"""Return FlyDSL stage2 weight as uint8, assuming input is already pre-shuffled."""
key = (
str(down_weight_shuffled.device),
int(down_weight_shuffled.data_ptr()),
tuple(down_weight_shuffled.shape),
)
cached = _FLYDSL_STAGE2_W2_SHUF_CACHE.get(key)
if cached is not None:
return cached
w2_u8 = _as_uint8_view(down_weight_shuffled).contiguous()
_FLYDSL_STAGE2_W2_SHUF_CACHE[key] = w2_u8
return w2_u8
def _get_flydsl_stage2_w2_scale_i32(down_weight_scale: torch.Tensor) -> torch.Tensor:
"""Pack+shuffle W2 MXFP4 microscales for FlyDSL stage2 (cached).
FlyDSL stage2 declares scale buffers as packed i32 matching
`make_preshuffle_scale_layout` (shape (MN/32, Kblk/8, 4, 16) of i32 packs).
This packs 4 e8m0 bytes (K_pack=2, N_pack=2) into each i32 element after the
exact `fp4_utils.e8m0_shuffle` permutation used by the reference harness.
"""
key = (
str(down_weight_scale.device),
int(down_weight_scale.data_ptr()),
tuple(down_weight_scale.shape),
)
cached = _FLYDSL_STAGE2_W2_SCALE_I32_CACHE.get(key)
if cached is not None:
return cached
# Reference harness produces raw microscale as a 2D tensor and shuffles via
# `fp4_utils.e8m0_shuffle`. Keep this path bit-identical, then pack into i32.
scale_u8 = down_weight_scale.view(torch.uint8)
if int(scale_u8.dim()) == 2:
scale2d = scale_u8.contiguous()
elif int(scale_u8.dim()) == 3:
experts, n_total, kblk = map(int, scale_u8.shape)
scale2d = scale_u8.reshape(experts * n_total, kblk).contiguous()
else:
raise ValueError("down_weight_scale must be [MN, Kblk] or [E, N, Kblk]")
from aiter.utility import fp4_utils
scale_shuf = fp4_utils.e8m0_shuffle(scale2d)
mn, kblk = map(int, scale_shuf.shape)
if (mn % 32) != 0 or (kblk % 8) != 0:
raise ValueError(
"down_weight_scale (after shuffle) must be divisible by 32 (MN) and 8 (Kblk)"
)
# Packed-i32 layout: (MN/32, Kblk/8, KLane=4, NLane=16) of i32 packs.
# Each i32 holds 4 bytes in [K_pack, N_pack] order (N_pack fastest).
scale6 = scale_shuf.view(mn // 32, kblk // 8, 4, 16, 2, 2)
pack4 = scale6.view(mn // 32, kblk // 8, 4, 16, 4)
# Byte-order hypothesis: mfma_scale selects bytes in the opposite (pack_M, pack_K)
# order vs our current (pack_K, pack_M) packing. Swap the middle two bytes to test.
pack4 = pack4[..., [0, 2, 1, 3]].contiguous()
scale_i32 = pack4.view(torch.int32).reshape(-1).contiguous()
_FLYDSL_STAGE2_W2_SCALE_I32_CACHE[key] = scale_i32
return scale_i32
def _get_flydsl_stage2_w2_scale_i32_from_shuffled(
down_weight_scale_shuffled: torch.Tensor,
) -> torch.Tensor:
"""Pack W2 MXFP4 microscales to i32 for FlyDSL stage2, assuming input is already shuffled."""
key = (
str(down_weight_scale_shuffled.device),
int(down_weight_scale_shuffled.data_ptr()),
tuple(down_weight_scale_shuffled.shape),
)
cached = _FLYDSL_STAGE2_W2_SCALE_SHUF_I32_CACHE.get(key)
if cached is not None:
return cached
scale_u8 = down_weight_scale_shuffled.view(torch.uint8)
if int(scale_u8.dim()) == 2:
scale2d = scale_u8.contiguous()
elif int(scale_u8.dim()) == 3:
experts, n_total, kblk = map(int, scale_u8.shape)
scale2d = scale_u8.reshape(experts * n_total, kblk).contiguous()
else:
raise ValueError(
"down_weight_scale_shuffled must be [MN, Kblk] or [E, N, Kblk] (uint8 view)"
)
if int(scale2d.numel()) % 4 != 0:
raise ValueError(
"down_weight_scale_shuffled must have a byte size divisible by 4 for i32 packing"
)
scale_i32 = scale2d.view(torch.int32).reshape(-1).contiguous()
_FLYDSL_STAGE2_W2_SCALE_SHUF_I32_CACHE[key] = scale_i32
return scale_i32
def _pack_flydsl_stage2_a2_scale_i32(a2_scale: torch.Tensor) -> torch.Tensor:
"""Pack A2 MXFP4 microscales to i32 for FlyDSL stage2 (no cache; depends on input)."""
if a2_scale is None:
raise ValueError("a2_scale is required for FlyDSL fp4 stage2")
if not a2_scale.is_contiguous():
a2_scale = a2_scale.contiguous()
if int(a2_scale.numel()) % 4 != 0:
raise ValueError("a2_scale must have a byte size divisible by 4 for i32 packing")
return a2_scale.view(torch.int32).reshape(-1)
def _ensure_iter115_bf16_lds_packed_rowctx_patch() -> None:
global _ITER115_BF16_LDS_PACKED_ROWCTX_PATCHED
if _ITER115_BF16_LDS_PACKED_ROWCTX_PATCHED:
return
import aiter.ops.flydsl.kernels.mixed_moe_gemm_2stage as mixed_moe_gemm2_mod
orig_compile = mixed_moe_gemm2_mod.compile_mixed_moe_gemm2
src = inspect.getsource(orig_compile)
lds_out_alias_anchor = """ # Alias the same underlying LDS bytes as f16/bf16 for epilogue shuffle.
lds_out = (
SmemPtr(
base_ptr,
lds_x_ptr.byte_offset,
(I.bf16 if out_is_bf16 else I.f16),
shape=(tile_m * tile_n,),
).get()
if _use_cshuffle_epilog
else None
)
"""
if lds_out_alias_anchor not in src:
raise RuntimeError("iter115 lds packed rowctx alias anchor missing")
write_row_tail_anchor = """ vector.store(v1, lds_out, [lds_idx], alignment=2)
def precompute_row(*, row_local, row):
"""
if write_row_tail_anchor not in src:
raise RuntimeError("iter115 lds packed rowctx write-row tail anchor missing")
precompute_row_def_anchor = """ def precompute_row(*, row_local, row):
"""
if precompute_row_def_anchor not in src:
raise RuntimeError("iter115 lds packed rowctx precompute def anchor missing")
store_pair_anchor = """ def store_pair(*, row_local, row, row_ctx, col_pair0, col_g0, frag):
fused = row_ctx
t = fused & mask24_i32
s = fused >> 24
t_idx = arith.index_cast(ir.IndexType.get(), t)
s_idx = arith.index_cast(ir.IndexType.get(), s)
if not bool(accumulate):
# ---- 64-bit global store path (avoids i32 offset overflow) ----
# Compute full element offset in i64 (index type) and use global_store.
ts_idx = t_idx * arith.constant(topk, index=True) + s_idx
col_idx = col_g0 # already index type
elem_off = (
ts_idx * arith.constant(model_dim, index=True) + col_idx
)
# Align to even element boundary for <2 x bf16/f16> stores
c1_idx = arith.constant(1, index=True)
elem_off_even = elem_off - (elem_off & c1_idx)
byte_off_idx = elem_off_even * arith.constant(
out_elem_bytes, index=True
)
ptr_addr_idx = out_base_idx + byte_off_idx
out_ptr = buffer_ops.create_llvm_ptr(
ptr_addr_idx, address_space=1
)
out_ptr_v = (
out_ptr._value if hasattr(out_ptr, "_value") else out_ptr
)
frag_v = frag._value if hasattr(frag, "_value") else frag
llvm.StoreOp(frag_v, out_ptr_v, alignment=4)
else:
# ---- accumulate=True: 64-bit global atomic path ----
# Avoids i32 offset overflow when tokens*model_dim*2 > INT32_MAX
# (~150K tokens for model_dim=7168).
# Unified bf16/f16 path using llvm.AtomicRMWOp with 64-bit pointer.
col_idx = col_g0 # already index type
elem_off = (
t_idx * arith.constant(model_dim, index=True) + col_idx
)
# Align to even element boundary for <2 x bf16/f16> atomics
c1_idx = arith.constant(1, index=True)
elem_off_even = elem_off - (elem_off & c1_idx)
byte_off_idx = elem_off_even * arith.constant(
out_elem_bytes, index=True
)
ptr_addr_idx = out_base_idx + byte_off_idx
out_ptr = buffer_ops.create_llvm_ptr(
ptr_addr_idx, address_space=1
)
out_ptr_v = (
out_ptr._value if hasattr(out_ptr, "_value") else out_ptr
)
frag_v = frag._value if hasattr(frag, "_value") else frag
llvm.AtomicRMWOp(
llvm.AtomicBinOp.fadd,
out_ptr_v,
frag_v,
llvm.AtomicOrdering.monotonic,
syncscope="agent",
alignment=4,
)
"""
if store_pair_anchor not in src:
raise RuntimeError("iter115 lds packed rowctx store patch anchor missing")
tuned_src = src.replace(
lds_out_alias_anchor,
""" # Alias the same underlying LDS bytes as f16/bf16 for epilogue shuffle.
lds_out = (
SmemPtr(
base_ptr,
lds_x_ptr.byte_offset,
(I.bf16 if out_is_bf16 else I.f16),
shape=(tile_m * tile_n,),
).get()
if _use_cshuffle_epilog
else None
)
lds_row_ctx_i32 = (
SmemPtr(
base_ptr,
lds_x_ptr.byte_offset + (2 * tile_m * tile_n),
I.i32,
shape=(tile_m,),
).get()
if _use_cshuffle_epilog
else None
)
""",
1,
)
tuned_src = tuned_src.replace(
write_row_tail_anchor,
""" vector.store(v1, lds_out, [lds_idx], alignment=2)
row_i32 = arith.index_cast(i32, row)
row_valid0 = arith.cmpu(row_i32, num_valid_i32, "ult")
row_valid = arith.andi(row_valid0, ts_ok)
row_valid_i32 = arith.select(row_valid, arith.i32(1), zero_i32)
row_base_i32 = arith.select(
row_valid,
t2 * arith.i32(model_dim),
zero_i32,
)
packed_rowctx_i32 = row_base_i32 * arith.i32(2) + row_valid_i32
memref.store(packed_rowctx_i32, lds_row_ctx_i32, [row_in_tile])
def precompute_row(*, row_local, row):
""",
1,
)
tuned_src = tuned_src.replace(
precompute_row_def_anchor,
""" def precompute_row(*, row_local, row):
packed_rowctx_i32 = memref.load(lds_row_ctx_i32, [row_local])
row_valid_i32 = arith.andi(packed_rowctx_i32, arith.i32(1))
row_valid = arith.cmpu(row_valid_i32, zero_i32, "ugt")
row_base_i32 = arith.shrui(packed_rowctx_i32, arith.i32(1))
row_base_idx = arith.index_cast(ir.IndexType.get(), row_base_i32)
return (row_base_idx, row_valid)
""",
1,
)
tuned_src = tuned_src.replace(
store_pair_anchor,
""" def store_pair(*, row_local, row, row_ctx, col_pair0, col_g0, frag):
if not bool(accumulate):
raise RuntimeError("iter115 lds packed rowctx patch expects accumulate=True")
row_base_idx = row_ctx
col_idx = col_g0 # already index type
elem_off = row_base_idx + col_idx
c1_idx = arith.constant(1, index=True)
elem_off_even = elem_off - (elem_off & c1_idx)
byte_off_idx = elem_off_even * arith.constant(
out_elem_bytes, index=True
)
byte_off_i32 = arith.index_cast(i32, byte_off_idx)
atomic_add_f16x2(frag, byte_off_i32)
""",
1,
)
tuned_src = tuned_src.replace(
"_vscale_fix3",
"_vscale_fix3_iter115ldspackedrowctx",
1,
)
patch_ns = dict(mixed_moe_gemm2_mod.__dict__)
patch_filename = "iter115_mixed_moe_gemm2_bf16_lds_packed_rowctx.py"
linecache.cache[patch_filename] = (
len(tuned_src),
None,
[line + "\n" for line in tuned_src.splitlines()],
patch_filename,
)
exec(compile(tuned_src, patch_filename, "exec"), patch_ns)
tuned_compile = patch_ns["compile_mixed_moe_gemm2"]
def _iter115_compile_mixed_moe_gemm2(**kwargs):
out_dtype = str(kwargs.get("out_dtype", "")).strip().lower()
if (
int(kwargs.get("model_dim", 0)) == 7168
and int(kwargs.get("inter_dim", 0)) == 2048
and int(kwargs.get("topk", 0)) == 9
and int(kwargs.get("tile_m", 0)) == 64
and int(kwargs.get("tile_n", 0)) == 128
and int(kwargs.get("tile_k", 0)) == 256
and kwargs.get("a_dtype") == "fp4"
and kwargs.get("b_dtype") == "fp4"
and out_dtype in ("bf16", "bfloat16")
and bool(kwargs.get("accumulate", True))
):
return tuned_compile(**kwargs)
return orig_compile(**kwargs)
mixed_moe_gemm2_mod.compile_mixed_moe_gemm2 = _iter115_compile_mixed_moe_gemm2
_ITER115_BF16_LDS_PACKED_ROWCTX_PATCHED = True
def _run_manual_fp4_two_stage(
hidden_states: torch.Tensor,
gate_up_weight_shuffled: torch.Tensor,
down_weight_shuffled: torch.Tensor,
gate_up_weight_scale_shuffled: torch.Tensor,
down_weight_scale_shuffled: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
dispatch: ManualDispatch,
) -> torch.Tensor:
token_num, _ = hidden_states.shape
topk = topk_ids.shape[1]
expert_count, model_dim, inter_dim = fused_moe_mod.get_inter_dim(
gate_up_weight_shuffled.shape,
down_weight_shuffled.shape,
)
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_out = (
fused_moe_mod.moe_sorting(
topk_ids,
topk_weights,
expert_count,
model_dim,
hidden_states.dtype,
dispatch.block_m,
None,
None,
0,
)
)
a1, a1_scale = fused_moe_mod.fused_dynamic_mxfp4_quant_moe_sort(
hidden_states,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
topk=1,
block_size=dispatch.block_m,
)
inter_states = torch.empty(
(token_num, topk, inter_dim),
dtype=hidden_states.dtype,
device=hidden_states.device,
)
inter_states = fused_moe_mod.ck_moe_stage1(
a1,
gate_up_weight_shuffled,
down_weight_shuffled,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
inter_states,
topk,
dispatch.block_m,
a1_scale,
gate_up_weight_scale_shuffled.view(dtypes.fp8_e8m0),
kernelName=dispatch.kernel_name1,
sorted_weights=None,
quant_type=QuantType.per_1x32,
activation=ActivationType.Silu,
splitk=0,
use_non_temporal_load=dispatch.use_non_temporal_load,
dtype=hidden_states.dtype,
)
inter_states = inter_states.view(-1, inter_dim)
inter_states, a2_scale = fused_moe_mod.fused_dynamic_mxfp4_quant_moe_sort(
inter_states,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
topk=topk,
block_size=dispatch.block_m,
)
inter_states = inter_states.view(token_num, topk, -1)
aiter.ck_moe_stage2_fwd(
inter_states,
gate_up_weight_shuffled,
down_weight_shuffled,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
moe_out,
topk,
kernelName=dispatch.kernel_name2,
w2_scale=down_weight_scale_shuffled.view(dtypes.fp8_e8m0),
a2_scale=a2_scale,
block_m=dispatch.block_m,
sorted_weights=sorted_weights,
quant_type=QuantType.per_1x32,
activation=ActivationType.Silu,
use_non_temporal_load=dispatch.use_non_temporal_load,
)
return moe_out
def _run_2048_ck_stage1_flydsl_stage2(
hidden_states: torch.Tensor,
gate_up_weight_shuffled: torch.Tensor,
down_weight_shuffled: torch.Tensor,
gate_up_weight_scale_shuffled: torch.Tensor,
down_weight_scale_shuffled: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
) -> torch.Tensor:
"""Specialized path for the ranked MI355X tuple: CK stage1 + bf16 buffer-atomic stage2."""
global _ITER115_BF16_LDS_PACKED_ROWCTX_LOGGED
token_num, _ = hidden_states.shape
topk = int(topk_ids.shape[1])
if not _ITER115_BF16_LDS_PACKED_ROWCTX_LOGGED:
print(
f"[iter115 bf16-lds-packed-rowctx] entering hot path with routed topk={topk}",
file=sys.stderr,
)
_ITER115_BF16_LDS_PACKED_ROWCTX_LOGGED = True
expert_count, model_dim, inter_dim = fused_moe_mod.get_inter_dim(
gate_up_weight_shuffled.shape,
down_weight_shuffled.shape,
)
block_m = 64
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_out = (
fused_moe_mod.moe_sorting(
topk_ids,
topk_weights,
expert_count,
model_dim,
hidden_states.dtype,
block_m,
None,
None,
0,
)
)
a1, a1_scale = fused_moe_mod.fused_dynamic_mxfp4_quant_moe_sort(
hidden_states,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
topk=1,
block_size=block_m,
)
inter_states = torch.empty(
(token_num, topk, inter_dim),
dtype=hidden_states.dtype,
device=hidden_states.device,
)
inter_states = fused_moe_mod.ck_moe_stage1(
a1,
gate_up_weight_shuffled,
down_weight_shuffled,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
inter_states,
topk,
block_m,
a1_scale,
gate_up_weight_scale_shuffled.view(dtypes.fp8_e8m0),
kernelName="",
sorted_weights=None,
quant_type=QuantType.per_1x32,
activation=ActivationType.Silu,
splitk=0,
use_non_temporal_load=False,
)
inter_states = inter_states.view(-1, inter_dim)
inter_states, a2_scale = fused_moe_mod.fused_dynamic_mxfp4_quant_moe_sort(
inter_states,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
topk=topk,
block_size=block_m,
)
inter_states = inter_states.view(token_num, topk, -1)
from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2
# IMPORTANT: the official reference harness pre-shuffles weights with
# `shuffle_weight(..., layout=(16, 16))` and pre-shuffles W scales with
# `fp4_utils.e8m0_shuffle`. Feed those exact layouts (no extra reshuffle).
w2_u8 = _get_flydsl_stage2_w2_fp4_u8_from_shuffled(down_weight_shuffled)
w2_scale_i32 = _get_flydsl_stage2_w2_scale_i32_from_shuffled(
down_weight_scale_shuffled
)
a2_scale_i32 = _pack_flydsl_stage2_a2_scale_i32(a2_scale)
# Strict contract: FlyDSL tile_m must match moe_sorting block_m.
_ensure_iter115_bf16_lds_packed_rowctx_patch()
flydsl_moe_stage2(
_as_uint8_view(inter_states).contiguous(),
w2_u8,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
out=moe_out,
topk=topk,
tile_m=block_m,
tile_n=128,
tile_k=256,
a_dtype="fp4",
b_dtype="fp4",
out_dtype="bf16",
mode="atomic",
w2_scale=w2_scale_i32,
a2_scale=a2_scale_i32,
sorted_weights=sorted_weights,
)
return moe_out
def _run_2048_blockm_override(
hidden_states: torch.Tensor,
gate_up_weight_shuffled: torch.Tensor,
down_weight_shuffled: torch.Tensor,
gate_up_weight_scale_shuffled: torch.Tensor,
down_weight_scale_shuffled: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
config: dict,
block_size_m: int,
) -> torch.Tensor:
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,
block_size_M=block_size_m,
hidden_pad=hidden_pad,
intermediate_pad=intermediate_pad,
)
def custom_kernel(data: input_t) -> output_t:
"""
Submission template for DeepSeek-R1 MXFP4 MoE kernel.
Input data tuple:
hidden_states: [M, d_hidden] bf16
gate_up_weight: [E, 2*d_expert_pad, d_hidden_pad//2] fp4x2 (raw)
down_weight: [E, d_hidden_pad, d_expert_pad//2] fp4x2 (raw)
gate_up_weight_scale: [E, 2*d_expert_pad, scale_K] e8m0 (raw)
down_weight_scale: [E, d_hidden_pad, scale_K] e8m0 (raw)
gate_up_weight_shuffled: [E, 2*d_expert_pad, d_hidden_pad//2] fp4x2 (shuffled)
down_weight_shuffled: [E, d_hidden_pad, d_expert_pad//2] fp4x2 (shuffled)
gate_up_weight_scale_shuffled:[padded, flat] e8m0 (shuffled)
down_weight_scale_shuffled: [padded, flat] e8m0 (shuffled)
topk_weights: [M, total_top_k] float32
topk_ids: [M, total_top_k] int32
config: dict
Returns:
output: [M, d_hidden] bf16
"""
(
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
token_num = int(hidden_states.shape[0])
if (
config["d_hidden"] == 7168
and config["d_expert"] == 512
and token_num in _MANUAL_DISPATCH
and int(topk_ids.shape[1]) == 9
and int(gate_up_weight_shuffled.shape[0]) == 33
):
manual_dispatch = _MANUAL_DISPATCH[token_num]
return _run_manual_fp4_two_stage(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
manual_dispatch,
)
if (
config["d_hidden"] == 7168
and config["d_expert"] == 2048
and token_num == 512
and int(config["n_routed_experts"]) == 32
and int(config["n_shared_experts"]) == 1
and int(config["total_top_k"]) == 9
and int(topk_ids.shape[1]) == 9
and int(gate_up_weight_shuffled.shape[0]) == 33
):
# Ranked tuple specialized path: CK stage1 + FlyDSL/HIP stage2.
return _run_2048_ck_stage1_flydsl_stage2(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
)
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,
)
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
scrolls · 768 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