Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 MoEsuite of 7 cases
AMD Instinct MI355X
168.3µs
#298 of 782
2026-03-15

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-epiloguelds_out_alias_anchor = """ # Alias the same underlying LDS bytes as f16/bf16 for epilogue shuffle.
split-ksplitk=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