submission 516786
parcadei · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 893 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-516786?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:68064e2614149c2d6dfb80ad4ac20b75523e31ee90cae00bc66a3a4bfcaef2fc
license declaredunknown
license concludedunknown
authorsparcadei
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Reinterpret byte-packed FP4 / E8M0 tensors as raw uint8."""num-warps = 4
num_warps=4,stages = 2
num_stages=2,tile-k = 256
BLOCK_K = 256 # bf16 elements per K iterationKernel source
submission.py893 lines
import os
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")
import torch
import triton
import triton.language as tl
import aiter
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe, get_2stage_cfgs, get_inter_dim, get_padded_M
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
from aiter.utility import fp4_utils
from task import input_t, output_t
_TOKEN_SORT_FUSE_THRESHOLD = 1024
_MANUAL_PATH_SHAPES = {
(16, 257, 7168, 256, 9),
(128, 257, 7168, 256, 9),
(512, 257, 7168, 256, 9),
(512, 33, 7168, 512, 9), # shape 6: large M, ck2stages better
(512, 33, 7168, 2048, 9), # shape 7: large M, ck2stages better
}
# E=32 shapes with small M: route through fused_moe with ksplit=2 (cktile kernel)
_KSPLIT_SHAPES = {
(16, 33, 7168, 512, 9), # shape 4: m_per_expert≈4
(128, 33, 7168, 512, 9), # shape 5: m_per_expert≈34
}
_BLOCK_M_EXACT: dict[tuple[int, int, int, int, int], int] = {
# shape key = (token_num, expert_count, model_dim, inter_dim, topk)
(128, 33, 7168, 512, 9): 32, # was 64, m_per_expert≈34 → 32 saves 5µs
(512, 33, 7168, 2048, 9): 64, # was 128, try 64 (128→32 regressed badly)
}
_SORT_CACHE: dict[
tuple[int, int, int, int, int, torch.dtype, int],
tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],
] = {}
_WORKSPACE_CACHE: dict[
tuple[str, int, int, int, int, torch.dtype],
torch.Tensor,
] = {}
_DISPATCH_CACHE: dict[tuple, dict] = {}
_SCALE_VIEW_CACHE: dict[int, torch.Tensor] = {}
# ---------------------------------------------------------------------------
# Inline MXFP4 quantization (bf16 tile -> fp4x2 + E8M0 scales)
# Ported from mxfp4-mm/submission.py
# ---------------------------------------------------------------------------
@triton.jit
def mxfp4_quant_tile(
x, # [BLOCK_M, BLOCK_K] fp32
BLOCK_M: tl.constexpr,
BLOCK_K: tl.constexpr,
SCALE_GROUP_SIZE: tl.constexpr,
):
EXP_BIAS_FP32: tl.constexpr = 127
EXP_BIAS_FP4: tl.constexpr = 1
MBITS_F32: tl.constexpr = 23
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_K // SCALE_GROUP_SIZE
x = x.reshape(BLOCK_M, NUM_QUANT_BLOCKS, SCALE_GROUP_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
quant_scale = tl.exp2(-scale_e8m0_unbiased)
qx = x * quant_scale
qx = qx.to(tl.uint32, bitcast=True)
s = qx & 0x80000000
e = (qx >> MBITS_F32) & 0xFF
m = qx & 0x7FFFFF
E8_BIAS: tl.constexpr = 127
E2_BIAS: tl.constexpr = 1
adjusted_exponents = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)
m = tl.where(e < E8_BIAS, (0x400000 | (m >> 1)) >> adjusted_exponents, m)
e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
e2m1_value = ((s >> 28) | e2m1_tmp).to(tl.uint8)
e2m1_value = tl.reshape(
e2m1_value, [BLOCK_M, NUM_QUANT_BLOCKS, SCALE_GROUP_SIZE // 2, 2]
)
evens, odds = tl.split(e2m1_value)
x_fp4 = evens | (odds << 4)
x_fp4 = x_fp4.reshape(BLOCK_M, BLOCK_K // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_M, NUM_QUANT_BLOCKS)
# ---------------------------------------------------------------------------
# Fused MoE Stage1: QUANT1 + GEMM1(gate) + GEMM1(up) + SiLU + QUANT2
# Option C: dual accumulators for gate and up halves
#
# Each program computes [BLOCK_M, BLOCK_N_HALF] of the SiLU(gate)*up output
# for one expert and one column block, then quantizes to fp4 in registers.
#
# Uses shuffled weights (same layout as CK stage1 kernels).
# Grid: 1D, mapped to (token_block_global, col_block). Expert via sorted_expert_ids[block].
# ---------------------------------------------------------------------------
@triton.jit
def _fused_moe_stage1_kernel(
# Inputs
hidden_states_ptr, # [M, d_hidden] bf16
w1_ptr, # [E, 2*d_expert_pad//16, d_hidden_pad//2*16] fp4x2 (shuffled)
w1_scale_ptr, # [E, 2*d_expert_pad//32, d_hidden_pad] e8m0 (shuffled)
sorted_ids_ptr, # [max_num_tokens_padded] int32
sorted_expert_ids_ptr, # [max_num_m_blocks] int32 — expert_id per block
num_valid_ids_ptr, # [2] int32 — total_tokens on GPU (avoids CPU sync)
# Outputs
a2_ptr, # [token_num*topk, d_expert_pad//2] fp4x2
a2_scale_ptr, # [token_num*topk, d_expert_pad//32] e8m0
# Dimensions
d_hidden: int,
d_expert_pad: int,
d_hidden_pad: int,
stride_hs_m: int, # hidden_states stride dim 0
stride_hs_k: int, # hidden_states stride dim 1
stride_w1_e: int, # w1 stride dim 0 (expert)
stride_w1_n: int, # w1 stride dim 1 (N//16)
stride_w1_k: int, # w1 stride dim 2 (K_packed*16)
stride_w1s_e: int, # w1_scale stride dim 0 (expert)
stride_w1s_n: int, # w1_scale stride dim 1 (N//32)
stride_w1s_k: int, # w1_scale stride dim 2 (K)
stride_a2_row: int, # a2 stride dim 0
stride_a2s_row: int, # a2_scale stride dim 0
M_output: int, # token_num * topk — bounds for scatter writes
token_num: int, # number of real tokens (bounds for hidden_states reads)
# Meta-parameters
BLOCK_M: tl.constexpr,
BLOCK_K: tl.constexpr, # element-space K block (bf16 elements)
BLOCK_N_HALF: tl.constexpr, # columns of d_expert per program
NUM_K_ITERS: tl.constexpr, # d_hidden_pad // BLOCK_K
TOPK: tl.constexpr, # experts per token (for decoding sorted_ids)
):
SCALE_GROUP_SIZE: tl.constexpr = 32
pid = tl.program_id(0)
# Read total_tokens from GPU (no CPU sync needed)
total_tokens = tl.load(num_valid_ids_ptr).to(tl.int32)
# Map pid -> (token_block_global, col_block_id)
num_col_blocks = d_expert_pad // BLOCK_N_HALF
col_block_id = pid % num_col_blocks
token_block_global = pid // num_col_blocks
# Early exit for programs beyond actual token blocks
num_token_blocks = tl.cdiv(total_tokens, BLOCK_M)
if token_block_global >= num_token_blocks:
return
# sorted_expert_ids[block_idx] = expert_id for that block (from moe_sorting_fwd).
# Direct lookup — one load per program.
expert_id = tl.load(sorted_expert_ids_ptr + token_block_global).to(tl.int32)
token_offset = token_block_global * BLOCK_M
offs_m = tl.arange(0, BLOCK_M)
valid_mask = (token_offset + offs_m) < total_tokens
# Load and decode packed sorted_ids for this block.
# sorted_ids encoding: (token_id << 24) | topk_id
# token_id indexes into hidden_states [M, d_hidden]
# original_m_idx = token_id * topk + topk_id indexes into a2 [M*topk, ...]
packed_ids = tl.load(
sorted_ids_ptr + token_offset + offs_m,
mask=valid_mask,
other=0,
)
token_ids = (packed_ids & 0xFFFFFF).to(tl.int32)
topk_ids = (packed_ids >> 24).to(tl.int32)
# Bounds masks: padding entries in sorted_ids may have out-of-range indices.
# token_ids must be < token_num (hidden_states rows).
# original_m_idx must be < M_output (a2 rows = token_num * topk).
original_m_idx = token_ids * TOPK + topk_ids
hs_valid = valid_mask & (token_ids < token_num)
out_valid = valid_mask & (original_m_idx < M_output)
# Dual accumulators for gate and up
acc_gate = tl.zeros((BLOCK_M, BLOCK_N_HALF), dtype=tl.float32)
acc_up = tl.zeros((BLOCK_M, BLOCK_N_HALF), dtype=tl.float32)
# Column offsets in the [2*d_expert_pad] weight dimension
gate_col_offset = col_block_id * BLOCK_N_HALF
up_col_offset = d_expert_pad + gate_col_offset
# Pre-compute shuffled N-dim offsets (constant across K iterations)
# Weight super-row offsets (16 rows per super-row, gate and up in separate halves)
offs_bn_gate = gate_col_offset // 16 + tl.arange(0, BLOCK_N_HALF // 16)
offs_bn_up = up_col_offset // 16 + tl.arange(0, BLOCK_N_HALF // 16)
# Scale N1-block offsets: each N1 block interleaves 16 gate + 16 up rows.
# Need BLOCK_N_HALF // 16 N1 blocks (not //32) to get BLOCK_N_HALF gate-only rows.
# Gate and up share the SAME N1 blocks — differ only in N_Pack byte position.
offs_bsn = gate_col_offset // 16 + tl.arange(0, BLOCK_N_HALF // 16)
# Gate-only and up-only scale K offsets within each K1 block (256 bytes).
# Shuffled layout per N1 per K1: [K_Lane=4, N_Lane=16, K_Pack=2, N_Pack=2].
# Gate = N_Pack=0 (even bytes), Up = N_Pack=1 (odd bytes).
_sidx = tl.arange(0, 128) # 128 = K_Lane(4) * N_Lane(16) * K_Pack(2)
_skl = _sidx // 32 # K_Lane index
_snl = (_sidx // 2) % 16 # N_Lane index
_skp = _sidx % 2 # K_Pack index
_offs_k_gate = _skl * 64 + _snl * 4 + _skp * 2 # N_Pack=0 byte positions
_offs_k_up = _offs_k_gate + 1 # N_Pack=1 byte positions
# Base pointers for this expert's weights
w1_base = w1_ptr + expert_id * stride_w1_e
w1s_base = w1_scale_ptr + expert_id * stride_w1s_e
for k_iter in range(NUM_K_ITERS):
k_start = k_iter * BLOCK_K
# --- Load hidden_states [BLOCK_M, BLOCK_K] via gathered token_ids ---
offs_k = tl.arange(0, BLOCK_K)
hs_ptrs = hidden_states_ptr + token_ids[:, None] * stride_hs_m + (k_start + offs_k[None, :]) * stride_hs_k
hs_mask = hs_valid[:, None] & ((k_start + offs_k[None, :]) < d_hidden)
hs = tl.load(hs_ptrs, mask=hs_mask, other=0.0)
# Quantize hidden_states to MXFP4 inline
a_fp4, a_scales = mxfp4_quant_tile(
hs.to(tl.float32), BLOCK_M=BLOCK_M, BLOCK_K=BLOCK_K,
SCALE_GROUP_SIZE=SCALE_GROUP_SIZE,
)
# --- Shuffled weight K-dim offsets ---
offs_k_shuffle = (k_start // 2) * 16 + tl.arange(0, (BLOCK_K // 2) * 16)
# --- Load + unshuffle GATE weight tile (unchanged) ---
w1_gate = tl.load(
w1_base + offs_bn_gate[:, None] * stride_w1_n + offs_k_shuffle[None, :] * stride_w1_k,
cache_modifier=".cg",
)
w1_gate = (
w1_gate.reshape(1, BLOCK_N_HALF // 16, BLOCK_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_N_HALF, BLOCK_K // 2)
.trans(1, 0)
)
# Gate scales: load gate-only bytes (N_Pack=0) + 4D unshuffle
offs_bsk_gate = k_start + _offs_k_gate
w1_gate_scales = tl.load(
w1s_base + offs_bsn[:, None] * stride_w1s_n + offs_bsk_gate[None, :] * stride_w1s_k,
cache_modifier=".cg",
)
w1_gate_scales = (
w1_gate_scales
.reshape(BLOCK_N_HALF // 16, 4, 16, 2) # [N1, K_Lane, N_Lane, K_Pack]
.permute(0, 2, 3, 1) # [N1, N_Lane, K_Pack, K_Lane]
.reshape(BLOCK_N_HALF, BLOCK_K // SCALE_GROUP_SIZE)
)
acc_gate = tl.dot_scaled(a_fp4, a_scales, "e2m1", w1_gate, w1_gate_scales, "e2m1", acc_gate)
# --- Load + unshuffle UP weight tile ---
w1_up = tl.load(
w1_base + offs_bn_up[:, None] * stride_w1_n + offs_k_shuffle[None, :] * stride_w1_k,
cache_modifier=".cg",
)
w1_up = (
w1_up.reshape(1, BLOCK_N_HALF // 16, BLOCK_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_N_HALF, BLOCK_K // 2)
.trans(1, 0)
)
# Up scales: load up-only bytes (N_Pack=1) + 4D unshuffle
offs_bsk_up = k_start + _offs_k_up
w1_up_scales = tl.load(
w1s_base + offs_bsn[:, None] * stride_w1s_n + offs_bsk_up[None, :] * stride_w1s_k,
cache_modifier=".cg",
)
w1_up_scales = (
w1_up_scales
.reshape(BLOCK_N_HALF // 16, 4, 16, 2) # [N1, K_Lane, N_Lane, K_Pack]
.permute(0, 2, 3, 1) # [N1, N_Lane, K_Pack, K_Lane]
.reshape(BLOCK_N_HALF, BLOCK_K // SCALE_GROUP_SIZE)
)
acc_up = tl.dot_scaled(a_fp4, a_scales, "e2m1", w1_up, w1_up_scales, "e2m1", acc_up)
# --- SiLU(gate) * up in registers ---
intermediate = tl.sigmoid(acc_gate) * acc_gate * acc_up # [BLOCK_M, BLOCK_N_HALF] fp32
# --- Quantize intermediate to MXFP4 ---
inter_fp4, inter_scales = mxfp4_quant_tile(
intermediate, BLOCK_M=BLOCK_M, BLOCK_K=BLOCK_N_HALF,
SCALE_GROUP_SIZE=SCALE_GROUP_SIZE,
)
# --- Write fp4 output in ORIGINAL token order (scatter via original_m_idx) ---
# original_m_idx computed above: token_ids * TOPK + topk_ids
# out_valid masks out padding entries where original_m_idx >= M_output.
offs_col_fp4 = (col_block_id * BLOCK_N_HALF // 2) + tl.arange(0, BLOCK_N_HALF // 2)
a2_ptrs = a2_ptr + original_m_idx[:, None] * stride_a2_row + offs_col_fp4[None, :]
tl.store(a2_ptrs, inter_fp4, mask=out_valid[:, None])
# Write scales in ORIGINAL token order (launcher will sort via moe_mxfp4_sort)
NUM_SCALE_COLS: tl.constexpr = BLOCK_N_HALF // SCALE_GROUP_SIZE
offs_scale_col = (col_block_id * NUM_SCALE_COLS) + tl.arange(0, NUM_SCALE_COLS)
a2s_ptrs = a2_scale_ptr + original_m_idx[:, None] * stride_a2s_row + offs_scale_col[None, :]
tl.store(a2s_ptrs, inter_scales, mask=out_valid[:, None])
def _run_fused_stage1(
hidden_states: torch.Tensor,
gate_up_weight_shuffled: torch.Tensor,
gate_up_weight_scale_shuffled: torch.Tensor,
sorted_ids: torch.Tensor,
sorted_expert_ids: torch.Tensor,
num_valid_ids: torch.Tensor,
token_num: int,
topk: int,
d_expert_pad: int,
d_hidden_pad: int,
block_m: int,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Launch the fused stage1 kernel returning (a2_fp4, a2_scale)."""
d_hidden = hidden_states.shape[1]
device = hidden_states.device
# Use upper-bound grid from CPU-known values to avoid GPU sync.
# total_tokens <= token_num * topk + num_experts * block_m (padding from moe_sorting).
# The kernel reads actual total_tokens from num_valid_ids on GPU and bounds-checks.
# sorted_expert_ids has max_num_m_blocks entries (block_idx -> expert_id).
# Use its length directly as the upper bound on token blocks.
max_token_blocks = sorted_expert_ids.shape[0]
# Output buffers in ORIGINAL token order: [token_num * topk, ...]
# The kernel scatter-writes via sorted_ids (token_ids) back to original positions.
orig_rows = token_num * topk
a2 = torch.empty((orig_rows, d_expert_pad // 2), dtype=torch.uint8, device=device)
a2_scale = torch.empty((orig_rows, d_expert_pad // 32), dtype=torch.uint8, device=device)
# Tuning constants
# BLOCK_K >= 256 required: the 7D scale unshuffle needs K_scales >= 8
# (K_scales = BLOCK_K // 32, and the reshape has a //8 factor)
BLOCK_K = 256 # bf16 elements per K iteration
BLOCK_N_HALF = min(d_expert_pad, 128) # columns of d_expert per program
num_col_blocks = d_expert_pad // BLOCK_N_HALF
NUM_K_ITERS = d_hidden_pad // BLOCK_K
# Upper-bound grid: extra programs early-exit via GPU-side bounds check
grid = (max_token_blocks * num_col_blocks,)
w1_scale = _as_e8m0_scale(gate_up_weight_scale_shuffled)
# The shuffled scale tensor comes as 2D [padded, flat] from the task harness.
# Reshape to 3D [E, 2*d_expert_pad//32, d_hidden_pad] so we can extract
# per-expert strides for the kernel's [expert, N//32, K_shuffled] indexing.
num_experts_w = gate_up_weight_shuffled.shape[0]
scale_n_dim = 2 * d_expert_pad // 32
w1_scale_3d = w1_scale.reshape(num_experts_w, scale_n_dim, d_hidden_pad)
# Triton on the runner doesn't recognise float4_e2m1fn_x2 or fp8_e8m0.
# View as raw uint8 at the kernel boundary (same fix as mxfp4-mm).
w1_u8 = _as_u8_storage(gate_up_weight_shuffled)
# Reshape to [E, N//16, K_packed*16] — kernel expects super-rows of 16 concatenated N-rows.
# Same reshape as MXFP4-MM (see mxfp4-mm/submission.py line 1893).
w1_u8 = w1_u8.view(w1_u8.shape[0], w1_u8.shape[1] // 16, w1_u8.shape[2] * 16)
w1_scale_u8 = _as_u8_storage(w1_scale_3d)
_fused_moe_stage1_kernel[grid](
hidden_states,
w1_u8,
w1_scale_u8,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
a2,
a2_scale,
d_hidden,
d_expert_pad,
d_hidden_pad,
hidden_states.stride(0),
hidden_states.stride(1),
w1_u8.stride(0),
w1_u8.stride(1),
w1_u8.stride(2),
w1_scale_u8.stride(0),
w1_scale_u8.stride(1),
w1_scale_u8.stride(2),
a2.stride(0),
a2_scale.stride(0),
orig_rows, # M_output: bounds for scatter writes
token_num, # token_num: bounds for hidden_states reads
BLOCK_M=block_m,
BLOCK_K=BLOCK_K,
BLOCK_N_HALF=BLOCK_N_HALF,
NUM_K_ITERS=NUM_K_ITERS,
TOPK=topk,
num_warps=4,
num_stages=2,
)
# Reshape to match expected layout: [token_num, topk, d_expert_pad//2]
# View as fp4x2 / fp8_e8m0 so CK stage2 gets the dtypes it expects.
a2 = a2.view(dtypes.fp4x2).view(token_num, topk, d_expert_pad // 2)
a2_scale = a2_scale.view(dtypes.fp8_e8m0).view(token_num, topk, d_expert_pad // 32)
# Apply moe_mxfp4_sort for CK stage2 scale layout compatibility.
# a2_scale is in original token order; moe_mxfp4_sort reorders into
# block-aligned layout that CK stage2 expects.
a2_scale = fp4_utils.moe_mxfp4_sort(
a2_scale,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
block_size=block_m,
)
return a2, a2_scale
def _as_u8_storage(x: torch.Tensor) -> torch.Tensor:
"""Triton on the runner doesn't recognise float4_e2m1fn_x2.
Reinterpret byte-packed FP4 / E8M0 tensors as raw uint8."""
if x.dtype == torch.uint8:
return x
if x.element_size() != 1:
return x
return x.view(torch.uint8)
def _as_e8m0_scale(x: torch.Tensor) -> torch.Tensor:
ptr = x.data_ptr()
cached = _SCALE_VIEW_CACHE.get(ptr)
if cached is not None:
return cached
result = x.view(dtypes.fp8_e8m0)
_SCALE_VIEW_CACHE[ptr] = result
return result
def _as_e8m0_scale_uncached(x: torch.Tensor) -> torch.Tensor:
return x.view(dtypes.fp8_e8m0)
def _shape_key(
token_num: int,
expert_count: int,
model_dim: int,
inter_dim: int,
topk: int,
) -> tuple[int, int, int, int, int]:
return (token_num, expert_count, model_dim, inter_dim, topk)
def _get_sort_buffers(
topk_ids: torch.Tensor,
num_experts: int,
model_dim: int,
moebuf_dtype: torch.dtype,
block_size: int,
):
device_index = topk_ids.device.index or 0
m, topk = topk_ids.shape
key = (m, topk, num_experts, model_dim, block_size, moebuf_dtype, device_index)
cached = _SORT_CACHE.get(key)
if cached is not None:
return cached
max_num_tokens_padded = int(topk_ids.numel() + num_experts * block_size - topk)
max_num_m_blocks = int((max_num_tokens_padded + block_size - 1) // block_size)
device = topk_ids.device
buffers = (
torch.empty(max_num_tokens_padded, dtype=dtypes.i32, device=device),
torch.empty(max_num_tokens_padded, dtype=dtypes.fp32, device=device),
torch.empty(max_num_m_blocks, dtype=dtypes.i32, device=device),
torch.empty(2, dtype=dtypes.i32, device=device),
torch.empty((m, model_dim), dtype=moebuf_dtype, device=device),
)
_SORT_CACHE[key] = buffers
return buffers
def _moe_sorting_cached(
topk_ids: torch.Tensor,
topk_weights: torch.Tensor,
num_experts: int,
model_dim: int,
moebuf_dtype: torch.dtype,
block_size: int,
):
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = _get_sort_buffers(
topk_ids, num_experts, model_dim, moebuf_dtype, block_size
)
aiter.moe_sorting_fwd(
topk_ids,
topk_weights,
sorted_ids,
sorted_weights,
sorted_expert_ids,
num_valid_ids,
moe_buf,
num_experts,
int(block_size),
None,
None,
0,
)
return sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids
def _get_workspace(
tag: str,
rows: int,
cols0: int,
cols1: int,
device: torch.device,
dtype: torch.dtype,
) -> torch.Tensor:
device_index = device.index or 0
key = (tag, device_index, rows, cols0, cols1, dtype)
cached = _WORKSPACE_CACHE.get(key)
if cached is None:
shape = (rows, cols0) if cols1 == 0 else (rows, cols0, cols1)
cached = torch.empty(shape, dtype=dtype, device=device)
_WORKSPACE_CACHE[key] = cached
return cached
def _quantize_hidden_states(
hidden_states: torch.Tensor,
sorted_ids: torch.Tensor,
num_valid_ids: torch.Tensor,
token_num: int,
block_m: int,
):
if token_num <= _TOKEN_SORT_FUSE_THRESHOLD:
return 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,
)
quant_func = aiter.get_hip_quant(QuantType.per_1x32)
a1, a1_scale = quant_func(hidden_states, quant_dtype=dtypes.fp4x2)
a1_scale = fp4_utils.moe_mxfp4_sort(
a1_scale,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
block_size=block_m,
)
return a1, a1_scale
def _quantize_intermediate(
intermediate: torch.Tensor,
sorted_ids: torch.Tensor,
num_valid_ids: torch.Tensor,
token_num: int,
topk: int,
block_m: int,
):
inter_dim = intermediate.shape[-1]
flat = intermediate.view(-1, inter_dim)
if token_num <= _TOKEN_SORT_FUSE_THRESHOLD:
a2, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
flat,
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
topk=topk,
block_size=block_m,
)
return a2.view(token_num, topk, -1), a2_scale
quant_func = aiter.get_hip_quant(QuantType.per_1x32)
a2, a2_scale = quant_func(flat, quant_dtype=dtypes.fp4x2, num_rows_factor=topk)
a2_scale = fp4_utils.moe_mxfp4_sort(
a2_scale[: token_num * topk, :].view(token_num, topk, -1),
sorted_ids=sorted_ids,
num_valid_ids=num_valid_ids,
token_num=token_num,
block_size=block_m,
)
return a2.view(token_num, topk, -1), a2_scale
def _build_dispatch_state(
hidden_states: torch.Tensor,
gate_up_weight_shuffled: torch.Tensor,
down_weight_shuffled: torch.Tensor,
topk_ids: torch.Tensor,
hidden_pad: int,
intermediate_pad: int,
):
token_num = hidden_states.shape[0]
topk = topk_ids.shape[1]
num_experts = gate_up_weight_shuffled.shape[0]
cache_key = (token_num, num_experts, topk, hidden_pad, intermediate_pad,
hidden_states.dtype, gate_up_weight_shuffled.dtype,
gate_up_weight_shuffled.shape[1], down_weight_shuffled.shape[1])
cached = _DISPATCH_CACHE.get(cache_key)
if cached is not None:
return cached
_, model_dim, inter_dim = get_inter_dim(
gate_up_weight_shuffled.shape, down_weight_shuffled.shape
)
is_g1u1 = inter_dim != gate_up_weight_shuffled.shape[1]
metadata = get_2stage_cfgs(
get_padded_M(token_num),
model_dim,
inter_dim,
num_experts,
topk,
hidden_states.dtype,
dtypes.fp4x2,
gate_up_weight_shuffled.dtype,
QuantType.per_1x32,
is_g1u1,
ActivationType.Silu,
False,
hidden_pad,
intermediate_pad,
True,
)
shape = _shape_key(token_num, num_experts, model_dim, inter_dim, topk)
block_m = int(_BLOCK_M_EXACT.get(shape, metadata.block_m))
result = {
"token_num": token_num,
"topk": topk,
"model_dim": model_dim,
"inter_dim": inter_dim,
"metadata": metadata,
"shape": shape,
"block_m": block_m,
}
_DISPATCH_CACHE[cache_key] = result
return result
def _run_manual_2stage(
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,
hidden_pad: int,
intermediate_pad: int,
) -> torch.Tensor:
state = _build_dispatch_state(
hidden_states,
gate_up_weight_shuffled,
down_weight_shuffled,
topk_ids,
hidden_pad,
intermediate_pad,
)
if state["shape"] not in _MANUAL_PATH_SHAPES or state["metadata"].run_1stage:
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,
)
sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids = _moe_sorting_cached(
topk_ids,
topk_weights,
gate_up_weight_shuffled.shape[0],
state["model_dim"],
hidden_states.dtype,
state["block_m"],
)
# Fused stage1: replaces QUANT1 + GEMM1+SiLU + QUANT2
# shuffle_weight preserves shape: gate_up_weight_shuffled is [E, 2*d_expert_pad, d_hidden_pad//2]
# (bytes rearranged internally but shape unchanged — see aiter/ops/shuffle.py line 24)
# The launcher reshapes to [E, 2*d_expert_pad//16, d_hidden_pad//2*16] for the kernel.
d_expert_pad_x2 = gate_up_weight_shuffled.shape[1]
d_expert_pad = d_expert_pad_x2 // 2
d_hidden_pad = gate_up_weight_shuffled.shape[2] * 2
use_fused = os.environ.get("MOE_FUSED_STAGE1", "0") == "1"
compare_mode = os.environ.get("MOE_COMPARE_STAGE1", "0") == "1"
if use_fused or compare_mode:
a2_fused, a2_scale_fused = _run_fused_stage1(
hidden_states,
gate_up_weight_shuffled,
gate_up_weight_scale_shuffled,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
state["token_num"],
state["topk"],
d_expert_pad,
d_hidden_pad,
state["block_m"],
)
if not compare_mode:
a2, a2_scale = a2_fused, a2_scale_fused
if not use_fused or compare_mode:
a1, a1_scale = _quantize_hidden_states(
hidden_states,
sorted_ids,
num_valid_ids,
state["token_num"],
state["block_m"],
)
intermediate = _get_workspace(
"intermediate",
state["token_num"],
state["topk"],
state["inter_dim"],
hidden_states.device,
hidden_states.dtype,
)
intermediate = state["metadata"].stage1(
a1,
gate_up_weight_shuffled,
down_weight_shuffled,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
intermediate,
state["topk"],
block_m=state["block_m"],
a1_scale=a1_scale,
w1_scale=_as_e8m0_scale(gate_up_weight_scale_shuffled),
sorted_weights=None,
)
a2, a2_scale = _quantize_intermediate(
intermediate,
sorted_ids,
num_valid_ids,
state["token_num"],
state["topk"],
state["block_m"],
)
if compare_mode:
import sys
# Compare a2 (fp4 values) and a2_scale between fused and CK-only
# Both are shaped [token_num, topk, ...] after _run_fused_stage1 reshaping
# CK-only a2/a2_scale come from _quantize_intermediate
a2f_flat = a2_fused.view(torch.uint8).flatten()
a2c_flat = a2.view(torch.uint8).flatten()
a2sf_flat = a2_scale_fused.view(torch.uint8).flatten()
a2sc_flat = a2_scale.view(torch.uint8).flatten()
# a2 (fp4 values) comparison
n = min(a2f_flat.shape[0], a2c_flat.shape[0])
a2_diff = (a2f_flat[:n] != a2c_flat[:n]).sum().item()
# a2_scale comparison
ns = min(a2sf_flat.shape[0], a2sc_flat.shape[0])
a2s_diff = (a2sf_flat[:ns] != a2sc_flat[:ns]).sum().item()
print(f"[COMPARE] block_m={state['block_m']} "
f"token_num={state['token_num']} topk={state['topk']} "
f"d_expert_pad={d_expert_pad} d_hidden_pad={d_hidden_pad}",
file=sys.stderr)
print(f"[COMPARE] a2 shapes: fused={list(a2_fused.shape)} ck={list(a2.shape)}",
file=sys.stderr)
print(f"[COMPARE] a2_scale shapes: fused={list(a2_scale_fused.shape)} ck={list(a2_scale.shape)}",
file=sys.stderr)
print(f"[COMPARE] a2 byte mismatches: {a2_diff}/{n} ({100*a2_diff/max(n,1):.1f}%)",
file=sys.stderr)
print(f"[COMPARE] a2_scale byte mismatches: {a2s_diff}/{ns} ({100*a2s_diff/max(ns,1):.1f}%)",
file=sys.stderr)
# Show first few mismatched positions for a2
if a2_diff > 0:
diff_mask = a2f_flat[:n] != a2c_flat[:n]
diff_pos = torch.where(diff_mask)[0][:10]
for pos in diff_pos:
p = pos.item()
print(f"[COMPARE] a2[{p}]: fused=0x{a2f_flat[p].item():02x} ck=0x{a2c_flat[p].item():02x}",
file=sys.stderr)
# Also compare the intermediate (pre-quantization) if we can
# The fused kernel quantizes inline, but let's see if the CK intermediate
# matches what fused would produce before quant
# For now just use the CK-only result for correctness
# a2, a2_scale already set from CK-only path
output = _get_workspace(
"output",
state["token_num"],
state["model_dim"],
0,
hidden_states.device,
hidden_states.dtype,
)
output.zero_()
state["metadata"].stage2(
a2,
gate_up_weight_shuffled,
down_weight_shuffled,
sorted_ids,
sorted_expert_ids,
num_valid_ids,
output,
state["topk"],
w2_scale=_as_e8m0_scale(down_weight_scale_shuffled),
a2_scale=a2_scale,
block_m=state["block_m"],
sorted_weights=sorted_weights,
)
return output
def custom_kernel(data: input_t) -> output_t:
(
hidden_states,
_, # gate_up_weight (raw, unused)
_, # down_weight (raw, unused)
_, # gate_up_weight_scale (raw, unused)
_, # down_weight_scale (raw, unused)
gate_up_weight_shuffled,
down_weight_shuffled,
gate_up_weight_scale_shuffled,
down_weight_scale_shuffled,
topk_weights,
topk_ids,
config,
) = data
hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
intermediate_pad = config["d_expert_pad"] - config["d_expert"]
gate_up_weight_shuffled.is_shuffled = True
down_weight_shuffled.is_shuffled = True
token_num = hidden_states.shape[0]
topk = topk_ids.shape[1]
num_experts = gate_up_weight_shuffled.shape[0]
_, model_dim, inter_dim = get_inter_dim(
gate_up_weight_shuffled.shape, down_weight_shuffled.shape
)
shape = _shape_key(token_num, num_experts, model_dim, inter_dim, topk)
# All shapes through fused_moe — tuned DSv3 configs exist on remote runner
# for E=257 shapes; E=33 shapes use cktile with ksplit=2
if shape in _KSPLIT_SHAPES:
os.environ["AITER_KSPLIT"] = "2"
else:
os.environ.pop("AITER_KSPLIT", None)
output = 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,
)
return output[:, : config["d_hidden"]]
scrolls · 893 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