Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 MoEsuite of 7 cases
AMD Instinct MI355X
169.4µs
#306 of 782
2026-03-08

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.

fp4Reinterpret byte-packed FP4 / E8M0 tensors as raw uint8."""
num-warps = 4num_warps=4,
stages = 2num_stages=2,
tile-k = 256BLOCK_K = 256 # bf16 elements per K iteration

Kernel 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