Skip to content
KernelIndex
Search⌘K

submission 754180

Kernel-Zhang · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 3424 lines, June 9 Researcher Reciprocity License v1.0.

kmoe_00258.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-754180?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
162.2µs
#282 of 782
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4cd5fde0c066eecff9df7dc19192d3a3acce958357cd73a683b41d4fc1bcb074
license declaredunknown
license concludedunknown
authorsKernel-Zhang
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4"Stage2 boundary contract fp4 view does not match quant-sort shape."
fused-epilogueimport aiter.ops.flydsl.kernels.mfma_epilogues as _mfma_epilogues
shared-memorysmem_ptr_lines = _collect_line_numbers(

Kernel source

kmoe_00258.py3424 lines
from typing import Any, Dict, Optional, Tuple
import math
import sys

import torch
import torch.nn.functional as F

try:
    from task import input_t, output_t
except ImportError:
    input_t = Tuple[Any, ...]
    output_t = torch.Tensor

try:
    from utils import make_match_reference
except ImportError:
    make_match_reference = None


PAD_ALIGN = 256
SHUFFLE_LAYOUT = (16, 16)
RTOL = 5e-2
ATOL = 5e-2


_AITER_IMPORT_ERROR = None

try:
    import aiter
    import triton
    from aiter import ActivationType, QuantType, dtypes
    from aiter.fused_moe import fused_moe
    import aiter.ops.moe_op as _aiter_moe_op
    from aiter.ops.triton._triton_kernels.quant.fused_mxfp4_quant import (
        _fused_dynamic_mxfp4_quant_moe_sort_kernel,
    )
    from aiter.ops.shuffle import shuffle_weight
    from aiter.utility import fp4_utils
    _CK_MOE = getattr(aiter, "ck_moe", None)
except Exception as exc:
    aiter = None
    triton = None
    ActivationType = None
    QuantType = None
    dtypes = None
    fused_moe = None
    _aiter_moe_op = None
    _fused_dynamic_mxfp4_quant_moe_sort_kernel = None
    shuffle_weight = None
    fp4_utils = None
    _CK_MOE = None
    _AITER_IMPORT_ERROR = exc


_DEVICE_KIND_CACHE: Dict[Tuple[str, Optional[int]], bool] = {}
_BACKEND_CACHE: Dict[Tuple[Any, ...], Tuple[str, Optional[int]]] = {}
_CK_MOE_NT_KWARG: Optional[str] = None
_CK_MOE_NT_KWARG_READY = False
_CK_MOE_STAGE_KWARGS = None
_CK_MOE_STAGE_KWARGS_READY = False
_HIP_QUANT_PER_1X32 = None
_TORCH_QUANT_PER_1X32 = None
_LAST_HIDDEN_PREQUANT_KEY: Optional[Tuple[Any, ...]] = None
_LAST_HIDDEN_PREQUANT_VALUE: Optional[Tuple[torch.Tensor, torch.Tensor]] = None
_LAST_HIDDEN_PREQUANT_TORCH_KEY: Optional[Tuple[Any, ...]] = None
_LAST_HIDDEN_PREQUANT_TORCH_VALUE: Optional[Tuple[torch.Tensor, torch.Tensor]] = None
_LAST_FULL_EXACT_SORT_CACHE_SHAPE: Optional[Tuple[int, int, int, int, int]] = None
_KMOE_FULL_EXACT_SORT_CACHE: Dict[
    Tuple[Any, ...],
    Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],
] = {}
_KMOE_FULL_EXACT_SORT_MISS_STREAK: Dict[Tuple[int, int, int, int, int], int] = {}
_KMOE_MOEBUF_WORKSPACE_CACHE: Dict[Tuple[Any, ...], torch.Tensor] = {}
_KMOE_MOEBUF_RING_WORKSPACE_CACHE: Dict[Tuple[Any, ...], Tuple[torch.Tensor, ...]] = {}
_KMOE_MOEBUF_RING_WORKSPACE_INDEX: Dict[Tuple[Any, ...], int] = {}
_LOGGED_EVENTS = set()
_MOE_DTYPE_ALIAS_PATCHED = False
_KMOE_PENDING_FLYDSL_STAGE2_PATCH_MODE = "ptr"
_KMOE_ACTIVE_FLYDSL_STAGE2_PATCH_MODE: Optional[str] = None
_KMOE_FLYDSL_STAGE2_ORIG_SOURCE: Optional[str] = None
_KMOE_FLYDSL_STAGE2_ORIG_SOURCE_FILE: Optional[str] = None
_KMOE_FLYDSL_STAGE2_PATCHED_FN_BY_MODE: Dict[str, Any] = {}
_KMOE_FLYDSL_STAGE2_COMPILED_CACHE: Dict[Tuple[Any, ...], Any] = {}
_KMOE_FLYDSL_STAGE2_ORIG_GET_COMPILED_STAGE2 = None
_KMOE_FLYDSL_STAGE2_COMPILED_CACHE_PATCHED = False
_KMOE_STAGE2_BOUNDARY_CONTRACT_WORKSPACE_CACHE: Dict[
    Tuple[Any, ...],
    Dict[str, torch.Tensor],
] = {}


_BENCHMARK_BLOCK_M = {
    (16, 7168, 256, 257, 9): 32,
    (128, 7168, 256, 257, 9): 64,
    (512, 7168, 256, 257, 9): 64,
    (16, 7168, 512, 33, 9): 32,
    (128, 7168, 512, 33, 9): 64,
    (512, 7168, 512, 33, 9): 128,
    (512, 7168, 2048, 33, 9): 64,
}

_CK_FORCE_NON_TEMPORAL_TRUE = {
    (16, 7168, 256, 257, 9),
    (128, 7168, 256, 257, 9),
    (512, 7168, 256, 257, 9),
}

_SKIP_SPLIT_SHARED_EXPERT_SHAPES = frozenset(
    {
        (16, 7168, 512, 33, 9),
        (128, 7168, 512, 33, 9),
        (512, 7168, 512, 33, 9),
        (512, 7168, 2048, 33, 9),
    }
)

_SELECTIVE_PREQUANT_CK_SHAPES = frozenset(
    {
        (16, 7168, 256, 257, 9),
        (128, 7168, 256, 257, 9),
        (512, 7168, 256, 257, 9),
    }
)

_SELECTIVE_RAW_CK_SHAPES = frozenset(
    {
        (512, 7168, 256, 257, 9),
    }
)

_CK_FORCE_STAGE_KERNELS = {
    (512, 7168, 256, 257, 9): (
        "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
        "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
        64,
    ),
}

_FLYDSL_STAGE2_FALLBACK_KERNELS = {
    (128, 7168, 256, 257, 9): "moe_ck2stages_gemm2_64x64x128x128_1x1_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
}

_FUSED_CFG_STAGE_OVERRIDES = {
    (128, 7168, 256, 257, 9): (
        "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
        "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_atomic",
        32,
    ),
    (512, 7168, 256, 257, 9): (
        "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
        "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
        32,
    ),
}

_MANUAL_FUSED_2STAGE_SHAPES = frozenset(
    {
        (128, 7168, 256, 257, 9),
    }
)

_KMOE_FULL_EXACT_SORT_CACHE_SHAPES = frozenset(
    {
        (128, 7168, 256, 257, 9),
    }
)

_KMOE_MISS_CLONE_SORT_INPUTS_SHAPES = frozenset()

_KMOE_CACHE_HIT_NOCLONE_SORT_INPUTS_SHAPES = frozenset(
    {
        (128, 7168, 256, 257, 9),
    }
)

_KMOE_MOEBUF_RING_WORKSPACE_SLOTS = {
    (128, 7168, 256, 257, 9): 2,
}

_KMOE_FULL_EXACT_SORT_SKIP_STORE_AFTER_MISS_STREAK = {
    (128, 7168, 256, 257, 9): 1,
}

_KMOE_FULL_EXACT_SORT_MAX_CACHE_ENTRIES = {
    (128, 7168, 256, 257, 9): 1,
}

_MANUAL_FUSED_QUANT_SORT_BLOCK_SIZE_MX = {
    (128, 7168, 256, 257, 9): {
        "a1": 64,
        "a2": 64,
    },
}

_MANUAL_FUSED_STAGE2_BOUNDARY_CONTRACT_SHAPES = frozenset(
    {
        (128, 7168, 256, 257, 9),
    }
)

_KMOE_STAGE2_BOUNDARY_CONTRACT_FRESH_REASONS = frozenset(
    {
        "first_bypass",
    }
)

_STAGE2_BOUNDARY_CONTRACT_SEGMENT_ALIGN = 256


def _require_aiter() -> None:
    if _AITER_IMPORT_ERROR is not None:
        raise RuntimeError(
            "This submission expects the gpumode runtime with `aiter` installed."
        ) from _AITER_IMPORT_ERROR


def _pad_to(x: int, align: int) -> int:
    return (x + align - 1) // align * align


def _maybe_contiguous(tensor: torch.Tensor) -> torch.Tensor:
    if tensor.is_contiguous():
        return tensor
    return tensor.contiguous()


def _ensure_dtype(tensor: torch.Tensor, dtype: torch.dtype) -> torch.Tensor:
    if tensor.dtype == dtype:
        return tensor
    return tensor.to(dtype)


def _allclose(a: torch.Tensor, b: torch.Tensor) -> bool:
    if a.shape != b.shape:
        return False
    if a.dtype != b.dtype:
        b = b.to(a.dtype)
    return torch.allclose(a, b, rtol=RTOL, atol=ATOL)


def _log_once(event_key: Tuple[Any, ...], message: str) -> None:
    if event_key in _LOGGED_EVENTS:
        return
    _LOGGED_EVENTS.add(event_key)
    print(message, file=sys.stderr)


def _set_pending_flydsl_stage2_patch_mode(mode: str) -> None:
    global _KMOE_PENDING_FLYDSL_STAGE2_PATCH_MODE
    if mode not in ("ptr", "rawbuf"):
        raise ValueError(f"unsupported flydsl stage2 patch mode: {mode}")
    _KMOE_PENDING_FLYDSL_STAGE2_PATCH_MODE = mode


def _call_with_flydsl_stage2_patch_mode(stage2_callable, patch_mode: str, *args, **kwargs):
    prev_mode = _KMOE_PENDING_FLYDSL_STAGE2_PATCH_MODE
    _set_pending_flydsl_stage2_patch_mode(patch_mode)
    try:
        return stage2_callable(*args, **kwargs)
    finally:
        _set_pending_flydsl_stage2_patch_mode(prev_mode)


def _safe_tensor_version(tensor: torch.Tensor) -> Optional[int]:
    try:
        return getattr(tensor, "_version", 0)
    except RuntimeError:
        return None


def _tensor_identity_key(tensor: torch.Tensor) -> Tuple[Any, ...]:
    return (
        id(tensor),
        tensor.data_ptr(),
        tuple(tensor.shape),
        tuple(tensor.stride()),
        str(tensor.dtype),
        _safe_tensor_version(tensor),
        tensor.device.type,
        tensor.device.index,
    )


def _get_kmoe_full_exact_sort_state(
    topk_ids: torch.Tensor,
    topk_weights: torch.Tensor,
    num_experts: int,
    model_dim: int,
    moebuf_dtype: torch.dtype,
    block_m: int,
) -> Tuple[
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    Optional[torch.Tensor],
    bool,
    Optional[Tuple[Any, ...]],
]:
    cache_key = (
        _tensor_identity_key(topk_ids),
        _tensor_identity_key(topk_weights),
        num_experts,
        model_dim,
        block_m,
        str(moebuf_dtype),
    )
    cached = _KMOE_FULL_EXACT_SORT_CACHE.get(cache_key)
    if cached is not None:
        return (*cached, None, True, None)

    sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = (
        _aiter_fused_moe.moe_sorting(
            topk_ids,
            topk_weights,
            num_experts,
            model_dim,
            moebuf_dtype,
            block_m,
            None,
            None,
            0,
        )
    )
    _log_once(
        ("kmoe-full-exact-sort-cache-build", cache_key),
        "[submission] kmoe full exact sort cache build "
        f"rows={tuple(topk_ids.shape)} experts={num_experts} "
        f"block_m={block_m} device={topk_ids.device}",
    )
    return (
        sorted_ids,
        sorted_weights,
        sorted_expert_ids,
        num_valid_ids,
        moe_buf,
        False,
        cache_key,
    )


def _maybe_reset_full_exact_sort_cache(
    shape_key: Tuple[int, int, int, int, int],
) -> None:
    global _LAST_FULL_EXACT_SORT_CACHE_SHAPE

    if _LAST_FULL_EXACT_SORT_CACHE_SHAPE == shape_key:
        return
    if _KMOE_FULL_EXACT_SORT_CACHE:
        _KMOE_FULL_EXACT_SORT_CACHE.clear()
        _log_once(
            (
                "kmoe-full-exact-sort-cache-reset",
                _LAST_FULL_EXACT_SORT_CACHE_SHAPE,
                shape_key,
            ),
            "[submission] kmoe full exact sort cache reset "
            f"prev_shape={_LAST_FULL_EXACT_SORT_CACHE_SHAPE} "
            f"next_shape={shape_key}",
        )
    if _KMOE_MOEBUF_RING_WORKSPACE_CACHE:
        _KMOE_MOEBUF_RING_WORKSPACE_CACHE.clear()
        _KMOE_MOEBUF_RING_WORKSPACE_INDEX.clear()
        _log_once(
            (
                "kmoe-moebuf-ring-workspace-reset",
                _LAST_FULL_EXACT_SORT_CACHE_SHAPE,
                shape_key,
            ),
            "[submission] kmoe moebuf ring workspace reset "
            f"prev_shape={_LAST_FULL_EXACT_SORT_CACHE_SHAPE} "
            f"next_shape={shape_key}",
        )
    _LAST_FULL_EXACT_SORT_CACHE_SHAPE = shape_key


def _store_kmoe_full_exact_sort_cache(
    shape_key: Tuple[int, int, int, int, int],
    cache_key: Tuple[Any, ...],
    cached_values: Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],
) -> None:
    max_entries = _KMOE_FULL_EXACT_SORT_MAX_CACHE_ENTRIES.get(shape_key)
    if max_entries is not None and cache_key not in _KMOE_FULL_EXACT_SORT_CACHE:
        while len(_KMOE_FULL_EXACT_SORT_CACHE) >= max_entries:
            oldest_key = next(iter(_KMOE_FULL_EXACT_SORT_CACHE))
            _KMOE_FULL_EXACT_SORT_CACHE.pop(oldest_key, None)
            _log_once(
                ("kmoe-full-exact-sort-cache-evict", shape_key, max_entries),
                "[submission] kmoe full exact sort cache evict "
                f"shape={shape_key} max_entries={max_entries}",
            )
    _KMOE_FULL_EXACT_SORT_CACHE[cache_key] = cached_values


def _get_kmoe_moebuf_workspace(
    token_num: int,
    model_dim: int,
    dtype: torch.dtype,
    device: torch.device,
    ring_depth: int = 1,
) -> torch.Tensor:
    workspace_key = (
        token_num,
        model_dim,
        str(dtype),
        device.type,
        device.index,
    )
    if ring_depth > 1:
        ring_workspace_key = (*workspace_key, ring_depth)
        cached_ring = _KMOE_MOEBUF_RING_WORKSPACE_CACHE.get(ring_workspace_key)
        if cached_ring is None:
            cached_ring = tuple(
                torch.empty((token_num, model_dim), dtype=dtype, device=device)
                for _ in range(ring_depth)
            )
            _KMOE_MOEBUF_RING_WORKSPACE_CACHE[ring_workspace_key] = cached_ring
            _KMOE_MOEBUF_RING_WORKSPACE_INDEX[ring_workspace_key] = 0
            _log_once(
                ("kmoe-moebuf-ring-workspace-build", ring_workspace_key),
                "[submission] kmoe moebuf ring workspace build "
                f"shape={(token_num, model_dim)} slots={ring_depth} device={device}",
            )
        slot = _KMOE_MOEBUF_RING_WORKSPACE_INDEX.get(ring_workspace_key, 0)
        _KMOE_MOEBUF_RING_WORKSPACE_INDEX[ring_workspace_key] = (
            slot + 1
        ) % ring_depth
        return cached_ring[slot]
    cached = _KMOE_MOEBUF_WORKSPACE_CACHE.get(workspace_key)
    if cached is None:
        cached = torch.empty((token_num, model_dim), dtype=dtype, device=device)
        _KMOE_MOEBUF_WORKSPACE_CACHE[workspace_key] = cached
        _log_once(
            ("kmoe-moebuf-workspace-build", workspace_key),
            "[submission] kmoe moebuf workspace build "
            f"shape={(token_num, model_dim)} device={device}",
        )
    return cached


def _patch_moe_dtype_aliases() -> None:
    global _MOE_DTYPE_ALIAS_PATCHED

    if _MOE_DTYPE_ALIAS_PATCHED or _aiter_moe_op is None:
        return

    alias_names = {"fp4x2", "float4_e2m1fn_x2", "torch.float4_e2m1fn_x2"}
    fp4_dtype = dtypes.fp4x2
    alias_names.add(str(fp4_dtype))
    alias_names.add(str(fp4_dtype).replace("torch.", ""))

    if hasattr(torch, "float4_e2m1fn_x2"):
        _aiter_moe_op.dtype2str_dict.setdefault(torch.float4_e2m1fn_x2, "fp4x2")

    _aiter_moe_op.dtype2str_dict.setdefault(fp4_dtype, "fp4x2")
    _aiter_moe_op.str2dtype_dict.setdefault("fp4x2", fp4_dtype)

    for alias in alias_names:
        if alias:
            _aiter_moe_op.dtype2str_dict.setdefault(alias, "fp4x2")
            _aiter_moe_op.str2dtype_dict.setdefault(alias, fp4_dtype)

    _MOE_DTYPE_ALIAS_PATCHED = True


def _manual_fused_quant_sort_block_size_mx(
    shape_key: Tuple[int, int, int, int, int],
    stage: str,
) -> Optional[int]:
    stage_cfg = _MANUAL_FUSED_QUANT_SORT_BLOCK_SIZE_MX.get(shape_key)
    if stage_cfg is None:
        return None
    return stage_cfg.get(stage)


def _ceil_div(x: int, y: int) -> int:
    return (x + y - 1) // y


def _quant_sort_blockscale_shape(
    sorted_rows: int,
    n: int,
) -> Tuple[int, int, int, int, int]:
    scale_n = _ceil_div(n, 32)
    return (
        _ceil_div(sorted_rows, 32),
        _ceil_div(scale_n, 8),
        4,
        16,
        4,
    )


def _get_stage2_boundary_contract(
    shape_key: Tuple[int, int, int, int, int],
    token_num: int,
    topk: int,
    inter_dim: int,
    sorted_rows: int,
    device: torch.device,
    workspace_reason: str = "default",
) -> Optional[Dict[str, torch.Tensor]]:
    if shape_key not in _MANUAL_FUSED_STAGE2_BOUNDARY_CONTRACT_SHAPES:
        return None
    if inter_dim % 2 != 0:
        return None

    a2_fp4_u8_shape = (token_num * topk, inter_dim // 2)
    a2_scale_u8_shape = _quant_sort_blockscale_shape(sorted_rows, inter_dim)

    total_bytes = 0

    def reserve(size_bytes: int) -> Tuple[int, int]:
        nonlocal total_bytes
        total_bytes = _pad_to(total_bytes, _STAGE2_BOUNDARY_CONTRACT_SEGMENT_ALIGN)
        offset = total_bytes
        total_bytes += size_bytes
        return offset, size_bytes

    layout = {
        "a2_fp4_u8": reserve(math.prod(a2_fp4_u8_shape)),
        "a2_scale_u8": reserve(math.prod(a2_scale_u8_shape)),
    }
    cache_workspace = (
        workspace_reason not in _KMOE_STAGE2_BOUNDARY_CONTRACT_FRESH_REASONS
    )
    contract = None
    workspace_key = None
    if cache_workspace:
        workspace_key = (
            shape_key,
            token_num,
            topk,
            inter_dim,
            sorted_rows,
            device.type,
            device.index,
            workspace_reason,
        )
        contract = _KMOE_STAGE2_BOUNDARY_CONTRACT_WORKSPACE_CACHE.get(workspace_key)
    if contract is None:
        storage = torch.empty((total_bytes,), dtype=torch.uint8, device=device)

        def make_u8_view(name: str, shape: Tuple[int, ...]) -> torch.Tensor:
            offset, size_bytes = layout[name]
            return storage.narrow(0, offset, size_bytes).view(shape)

        contract = {
            "storage": storage,
            "a2_fp4_u8": make_u8_view("a2_fp4_u8", a2_fp4_u8_shape),
            "a2_scale_u8": make_u8_view("a2_scale_u8", a2_scale_u8_shape),
        }
        if cache_workspace:
            _KMOE_STAGE2_BOUNDARY_CONTRACT_WORKSPACE_CACHE[workspace_key] = contract
            event_key = ("kmoe-stage2-boundary-contract-workspace-build", workspace_key)
            workspace_mode = "cached"
        else:
            event_key = (
                "kmoe-stage2-boundary-contract-workspace-build-fresh",
                shape_key,
                token_num,
                topk,
                inter_dim,
                sorted_rows,
                device.type,
                device.index,
                workspace_reason,
            )
            workspace_mode = "fresh"
        _log_once(
            event_key,
            "[submission] kmoe stage2-boundary contract workspace build "
            f"shape={shape_key} sorted_rows={sorted_rows} bytes={total_bytes} "
            f"reason={workspace_reason} mode={workspace_mode} device={device}",
        )

    _log_once(
        ("manual2stage-stage2-boundary-contract", shape_key),
        "[submission] manual fused 2stage stage2-boundary contract active "
        f"shape={shape_key} sorted_rows={sorted_rows} bytes={total_bytes} "
        "pack=a2_fp4+a2_scale",
    )
    return contract


def _fused_dynamic_mxfp4_quant_moe_sort_tuned(
    x: torch.Tensor,
    sorted_ids: torch.Tensor,
    num_valid_ids: torch.Tensor,
    token_num: int,
    topk: int,
    block_size: int = 32,
    block_size_mx: Optional[int] = None,
    shape_key: Optional[Tuple[int, int, int, int, int]] = None,
    stage: Optional[str] = None,
    x_fp4_u8: Optional[torch.Tensor] = None,
    blockscale_e8m0_sorted_u8: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
    if (
        block_size_mx is None
        or triton is None
        or _fused_dynamic_mxfp4_quant_moe_sort_kernel is None
    ):
        return _aiter_fused_moe.fused_dynamic_mxfp4_quant_moe_sort(
            x,
            sorted_ids=sorted_ids,
            num_valid_ids=num_valid_ids,
            token_num=token_num,
            topk=topk,
            block_size=block_size,
        )

    M, N = x.shape
    assert (N // 2) % 2 == 0

    mxfp4_quant_block_size = 32
    x_fp4_shape = (M, N // 2)
    if x_fp4_u8 is None:
        x_fp4 = torch.empty(x_fp4_shape, dtype=torch.uint8, device=x.device)
    else:
        if (
            x_fp4_u8.dtype != torch.uint8
            or x_fp4_u8.device != x.device
            or tuple(x_fp4_u8.shape) != x_fp4_shape
        ):
            raise RuntimeError(
                "Stage2 boundary contract fp4 view does not match quant-sort shape."
            )
        x_fp4 = x_fp4_u8
    scale_n_valid = triton.cdiv(N, mxfp4_quant_block_size)
    scale_n = scale_n_valid

    block_size_m = 32
    block_size_n = 8
    block_size_m_u32 = 16
    block_size_n_u32 = 4

    n_i = scale_n
    m_o, n_o = sorted_ids.shape[0], n_i
    assert (n_i // 2) % 2 == 0
    assert block_size % block_size_m == 0

    blockscale_shape = (
        triton.cdiv(m_o, block_size_m),
        triton.cdiv(n_o, block_size_n),
        block_size_n_u32,
        block_size_m_u32,
        4,
    )
    if blockscale_e8m0_sorted_u8 is None:
        blockscale_e8m0_sorted = torch.empty(
            blockscale_shape,
            dtype=torch.uint8,
            device=x.device,
        )
    else:
        if (
            blockscale_e8m0_sorted_u8.dtype != torch.uint8
            or blockscale_e8m0_sorted_u8.device != x.device
            or tuple(blockscale_e8m0_sorted_u8.shape) != blockscale_shape
        ):
            raise RuntimeError(
                "Stage2 boundary contract scale view does not match quant-sort shape."
            )
        blockscale_e8m0_sorted = blockscale_e8m0_sorted_u8

    num_pid = triton.cdiv(M, block_size_mx) * scale_n + triton.cdiv(
        m_o, block_size_m
    ) * triton.cdiv(n_i, block_size_n)
    launch_kwargs = {
        "token_num": token_num,
        "N_i": n_i,
        "MXFP4_QUANT_BLOCK_SIZE": mxfp4_quant_block_size,
        "BLOCK_SIZE_Mx": block_size_mx,
        "BLOCK_SIZE_M": block_size_m // 2,
        "BLOCK_SIZE_N": block_size_n // 2,
        "TOPK": topk,
    }
    kernel_arg_names = tuple(
        getattr(_fused_dynamic_mxfp4_quant_moe_sort_kernel, "arg_names", ())
    )
    if "M_i" in kernel_arg_names:
        launch_kwargs["M_i"] = M
        _log_once(
            (
                "manual2stage-quant-sort-launcher-argnames-has-mi",
                shape_key,
                stage,
                block_size_mx,
            ),
            "[submission] quant-sort launcher compat add M_i "
            f"shape={shape_key} stage={stage} block_size_mx={block_size_mx}",
        )

    launcher = _fused_dynamic_mxfp4_quant_moe_sort_kernel[(num_pid,)]
    base_args = (
        x,
        x_fp4,
        sorted_ids,
        num_valid_ids,
        blockscale_e8m0_sorted,
        M,
        N,
        scale_n,
        *x.stride(),
        *x_fp4.stride(),
        *blockscale_e8m0_sorted.stride(),
    )
    try:
        launcher(*base_args, **launch_kwargs)
    except TypeError as exc:
        if "M_i" not in str(exc) or "M_i" in launch_kwargs:
            raise
        launch_kwargs["M_i"] = M
        _log_once(
            (
                "manual2stage-quant-sort-launcher-fallback-with-mi",
                shape_key,
                stage,
                block_size_mx,
            ),
            "[submission] quant-sort launcher fallback-with-M_i "
            f"shape={shape_key} stage={stage} block_size_mx={block_size_mx}",
        )
        launcher(*base_args, **launch_kwargs)

    return (
        x_fp4.view(dtypes.fp4x2),
        blockscale_e8m0_sorted.view(dtypes.fp8_e8m0).view(-1, n_o),
    )


def _run_manual_fused_quant_sort_with_override(
    shape_key: Tuple[int, int, int, int, int],
    stage: str,
    x: torch.Tensor,
    sorted_ids: torch.Tensor,
    num_valid_ids: torch.Tensor,
    token_num: int,
    topk: int,
    block_size: int,
    x_fp4_u8: Optional[torch.Tensor] = None,
    blockscale_e8m0_sorted_u8: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
    block_size_mx = _manual_fused_quant_sort_block_size_mx(shape_key, stage)
    if block_size_mx is not None:
        _log_once(
            ("manual2stage-quant-sort-tuned", shape_key, stage, block_size_mx),
            "[submission] manual fused quant-sort tuned "
            f"shape={shape_key} stage={stage} block_size_mx={block_size_mx}",
        )
    if x_fp4_u8 is not None or blockscale_e8m0_sorted_u8 is not None:
        _log_once(
            ("manual2stage-quant-sort-stage2-boundary-contract", shape_key, stage),
            "[submission] manual fused quant-sort stage2-boundary contract "
            f"shape={shape_key} stage={stage}",
        )
    return _fused_dynamic_mxfp4_quant_moe_sort_tuned(
        x,
        sorted_ids=sorted_ids,
        num_valid_ids=num_valid_ids,
        token_num=token_num,
        topk=topk,
        block_size=block_size,
        block_size_mx=block_size_mx,
        shape_key=shape_key,
        stage=stage,
        x_fp4_u8=x_fp4_u8,
        blockscale_e8m0_sorted_u8=blockscale_e8m0_sorted_u8,
    )


def _device_is_mi355x(device: torch.device) -> bool:
    if device.type != "cuda":
        return False
    key = (device.type, device.index)
    cached = _DEVICE_KIND_CACHE.get(key)
    if cached is not None:
        return cached
    try:
        name = torch.cuda.get_device_name(device).upper()
    except Exception:
        name = ""
    is_mi355x = "MI355" in name or "GFX950" in name
    _DEVICE_KIND_CACHE[key] = is_mi355x
    return is_mi355x


def _shape_key(
    hidden_states: torch.Tensor,
    topk_ids: torch.Tensor,
    config: Dict[str, int],
) -> Tuple[int, int, int, int, int]:
    return (
        hidden_states.shape[0],
        config["d_hidden"],
        config["d_expert"],
        config.get("n_routed_experts", 0) + config.get("n_shared_experts", 0),
        topk_ids.shape[1],
    )


def _backend_key(
    hidden_states: torch.Tensor,
    topk_ids: torch.Tensor,
    config: Dict[str, int],
) -> Tuple[Any, ...]:
    return (
        str(hidden_states.device),
        _shape_key(hidden_states, topk_ids, config),
    )


def _candidate_block_m(
    hidden_states: torch.Tensor,
    topk_ids: torch.Tensor,
    config: Dict[str, int],
) -> Optional[int]:
    exact = _BENCHMARK_BLOCK_M.get(_shape_key(hidden_states, topk_ids, config))
    if exact is not None:
        return exact

    token_count = hidden_states.shape[0]
    d_expert = config["d_expert"]
    if token_count <= 32:
        return 32
    if token_count <= 128:
        return 64
    if d_expert >= 2048:
        return 64
    return 128


def _forced_ck_non_temporal_load(
    hidden_states: torch.Tensor,
    topk_ids: torch.Tensor,
    config: Dict[str, int],
) -> Optional[bool]:
    if _shape_key(hidden_states, topk_ids, config) in _CK_FORCE_NON_TEMPORAL_TRUE:
        return True
    return None


def _get_hip_quant_per_1x32():
    global _HIP_QUANT_PER_1X32
    if _HIP_QUANT_PER_1X32 is None:
        _HIP_QUANT_PER_1X32 = aiter.get_hip_quant(QuantType.per_1x32)
    return _HIP_QUANT_PER_1X32


def _get_torch_quant_per_1x32():
    global _TORCH_QUANT_PER_1X32
    if _TORCH_QUANT_PER_1X32 is None:
        _TORCH_QUANT_PER_1X32 = aiter.get_torch_quant(QuantType.per_1x32)
    return _TORCH_QUANT_PER_1X32


def _get_hidden_prequant(
    hidden_states: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
    global _LAST_HIDDEN_PREQUANT_KEY, _LAST_HIDDEN_PREQUANT_VALUE

    key = (
        id(hidden_states),
        hidden_states.data_ptr(),
        hidden_states.numel(),
        _safe_tensor_version(hidden_states),
    )
    if _LAST_HIDDEN_PREQUANT_KEY == key and _LAST_HIDDEN_PREQUANT_VALUE is not None:
        return _LAST_HIDDEN_PREQUANT_VALUE

    quant_func = _get_hip_quant_per_1x32()
    hidden_q, hidden_scale = quant_func(hidden_states, quant_dtype=dtypes.fp4x2)
    hidden_q = _maybe_contiguous(hidden_q)
    hidden_scale = _maybe_contiguous(hidden_scale)
    _LAST_HIDDEN_PREQUANT_KEY = key
    _LAST_HIDDEN_PREQUANT_VALUE = (hidden_q, hidden_scale)
    return hidden_q, hidden_scale


def _get_hidden_prequant_torch(
    hidden_states: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
    global _LAST_HIDDEN_PREQUANT_TORCH_KEY, _LAST_HIDDEN_PREQUANT_TORCH_VALUE

    key = (
        id(hidden_states),
        hidden_states.data_ptr(),
        hidden_states.numel(),
        _safe_tensor_version(hidden_states),
    )
    if (
        _LAST_HIDDEN_PREQUANT_TORCH_KEY == key
        and _LAST_HIDDEN_PREQUANT_TORCH_VALUE is not None
    ):
        return _LAST_HIDDEN_PREQUANT_TORCH_VALUE

    quant_func = _get_torch_quant_per_1x32()
    hidden_q, hidden_scale = quant_func(hidden_states, quant_dtype=dtypes.fp4x2)
    hidden_q = _maybe_contiguous(hidden_q)
    hidden_scale = _maybe_contiguous(hidden_scale)
    _LAST_HIDDEN_PREQUANT_TORCH_KEY = key
    _LAST_HIDDEN_PREQUANT_TORCH_VALUE = (hidden_q, hidden_scale)
    return hidden_q, hidden_scale


def _should_prequant_ck(
    hidden_states: torch.Tensor,
    topk_ids: torch.Tensor,
    config: Dict[str, int],
) -> bool:
    shape = _shape_key(hidden_states, topk_ids, config)
    return shape in _SELECTIVE_PREQUANT_CK_SHAPES and shape not in _SELECTIVE_RAW_CK_SHAPES


def _forced_ck_stage_kernels(
    hidden_states: torch.Tensor,
    topk_ids: torch.Tensor,
    config: Dict[str, int],
) -> Optional[Tuple[str, str, int]]:
    return _CK_FORCE_STAGE_KERNELS.get(_shape_key(hidden_states, topk_ids, config))


def _should_use_manual_fused_2stage(
    hidden_states: torch.Tensor,
    topk_ids: torch.Tensor,
    config: Dict[str, int],
) -> bool:
    return _shape_key(hidden_states, topk_ids, config) in _MANUAL_FUSED_2STAGE_SHAPES


def _prepare_common(
    data: input_t,
) -> Tuple[
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    Dict[str, int],
    int,
    int,
]:
    _require_aiter()

    (
        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

    d_hidden = config["d_hidden"]
    d_expert = config["d_expert"]
    d_hidden_pad = config.get("d_hidden_pad", _pad_to(d_hidden, PAD_ALIGN))
    d_expert_pad = config.get("d_expert_pad", _pad_to(d_expert, PAD_ALIGN))

    if gate_up_weight_shuffled is None:
        gate_up_weight_shuffled = shuffle_weight(gate_up_weight, layout=SHUFFLE_LAYOUT)
    if down_weight_shuffled is None:
        down_weight_shuffled = shuffle_weight(down_weight, layout=SHUFFLE_LAYOUT)
    if gate_up_weight_scale_shuffled is None:
        gate_up_weight_scale_shuffled = fp4_utils.e8m0_shuffle(gate_up_weight_scale)
    if down_weight_scale_shuffled is None:
        down_weight_scale_shuffled = fp4_utils.e8m0_shuffle(down_weight_scale)

    hidden_states = _maybe_contiguous(_ensure_dtype(hidden_states, torch.bfloat16))
    topk_weights = _maybe_contiguous(_ensure_dtype(topk_weights, torch.float32))
    topk_ids = _maybe_contiguous(_ensure_dtype(topk_ids, torch.int32))

    gate_up_weight_shuffled = _maybe_contiguous(gate_up_weight_shuffled)
    down_weight_shuffled = _maybe_contiguous(down_weight_shuffled)
    gate_up_weight_scale_shuffled = _maybe_contiguous(gate_up_weight_scale_shuffled)
    down_weight_scale_shuffled = _maybe_contiguous(down_weight_scale_shuffled)

    return (
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        gate_up_weight_scale,
        down_weight_scale,
        gate_up_weight_scale_shuffled,
        down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        config,
        d_hidden_pad - d_hidden,
        d_expert_pad - d_expert,
    )


def _run_fused(
    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:
    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,
    )


def _run_ck_with_a1_scale(
    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,
    block_m: int,
    a1_scale: Optional[torch.Tensor],
    non_temporal_load: Optional[bool] = None,
    extra_kwargs: Optional[Dict[str, Any]] = None,
) -> torch.Tensor:
    base_kwargs = {
        "w1_scale": gate_up_weight_scale_shuffled,
        "w2_scale": down_weight_scale_shuffled,
        "a1_scale": a1_scale,
        "a2_scale": None,
        "block_m": block_m,
        "expert_mask": None,
    }

    def call(more_kwargs: Optional[Dict[str, Any]] = None) -> torch.Tensor:
        merged_kwargs = {}
        if extra_kwargs is not None:
            merged_kwargs.update(extra_kwargs)
        if more_kwargs is not None:
            merged_kwargs.update(more_kwargs)
        return _CK_MOE(
            hidden_states,
            gate_up_weight_shuffled,
            down_weight_shuffled,
            topk_weights,
            topk_ids,
            **base_kwargs,
            **merged_kwargs,
        )

    if non_temporal_load is None:
        return call()

    global _CK_MOE_NT_KWARG, _CK_MOE_NT_KWARG_READY
    if _CK_MOE_NT_KWARG_READY:
        if _CK_MOE_NT_KWARG is None:
            return call()
        try:
            return call({_CK_MOE_NT_KWARG: non_temporal_load})
        except Exception:
            _CK_MOE_NT_KWARG = None
            return call()

    for kwarg_name in ("non_temporal_load", "use_non_temporal_load"):
        try:
            out = call({kwarg_name: non_temporal_load})
            _CK_MOE_NT_KWARG = kwarg_name
            _CK_MOE_NT_KWARG_READY = True
            return out
        except TypeError:
            continue
        except Exception:
            _CK_MOE_NT_KWARG = None
            _CK_MOE_NT_KWARG_READY = True
            return call()

    _CK_MOE_NT_KWARG = None
    _CK_MOE_NT_KWARG_READY = True
    return call()


def _run_ck(
    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,
    block_m: int,
    non_temporal_load: Optional[bool] = None,
) -> torch.Tensor:
    return _run_ck_with_a1_scale(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        gate_up_weight_scale_shuffled,
        down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        block_m,
        None,
        non_temporal_load,
    )


def _run_ck_prequant(
    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,
    block_m: int,
    non_temporal_load: Optional[bool] = None,
) -> torch.Tensor:
    hidden_q, hidden_scale = _get_hidden_prequant(hidden_states)
    return _run_ck_with_a1_scale(
        hidden_q,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        gate_up_weight_scale_shuffled,
        down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        block_m,
        hidden_scale,
        non_temporal_load,
    )


def _run_ck_prequant_with_stage_kernels(
    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,
    block_m: int,
    non_temporal_load: Optional[bool],
    stage1_kernel: str,
    stage2_kernel: str,
) -> torch.Tensor:
    hidden_q, hidden_scale = _get_hidden_prequant(hidden_states)

    global _CK_MOE_STAGE_KWARGS, _CK_MOE_STAGE_KWARGS_READY

    def call(stage_kwargs: Dict[str, str]) -> torch.Tensor:
        return _run_ck_with_a1_scale(
            hidden_q,
            gate_up_weight_shuffled,
            down_weight_shuffled,
            gate_up_weight_scale_shuffled,
            down_weight_scale_shuffled,
            topk_weights,
            topk_ids,
            block_m,
            hidden_scale,
            non_temporal_load,
            stage_kwargs,
        )

    if _CK_MOE_STAGE_KWARGS_READY:
        if _CK_MOE_STAGE_KWARGS is None:
            raise RuntimeError("CK MoE stage-kernel kwargs are unavailable.")
        return call(
            {
                _CK_MOE_STAGE_KWARGS[0]: stage1_kernel,
                _CK_MOE_STAGE_KWARGS[1]: stage2_kernel,
            }
        )

    for kw1, kw2 in (
        ("kernelName1", "kernelName2"),
        ("kernel_name1", "kernel_name2"),
        ("stage1_kernelName", "stage2_kernelName"),
        ("stage1_kernel_name", "stage2_kernel_name"),
        ("gemm1_kernelName", "gemm2_kernelName"),
        ("gemm1_kernel_name", "gemm2_kernel_name"),
    ):
        try:
            out = call({kw1: stage1_kernel, kw2: stage2_kernel})
            _CK_MOE_STAGE_KWARGS = (kw1, kw2)
            _CK_MOE_STAGE_KWARGS_READY = True
            return out
        except TypeError:
            continue
        except Exception:
            _CK_MOE_STAGE_KWARGS = None
            _CK_MOE_STAGE_KWARGS_READY = True
            raise

    _CK_MOE_STAGE_KWARGS = None
    _CK_MOE_STAGE_KWARGS_READY = True
    raise RuntimeError("CK MoE stage-kernel kwargs are unavailable.")


def _run_fused_prequant(
    hidden_q: torch.Tensor,
    hidden_scale: 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:
    return fused_moe(
        hidden_q,
        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=hidden_scale,
        a2_scale=None,
        dtype=torch.bfloat16,
        hidden_pad=hidden_pad,
        intermediate_pad=intermediate_pad,
    )


def _run_manual_fused_2stage_prequant(
    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[str, int],
    hidden_pad: int,
    intermediate_pad: int,
    stage1_kernel: str,
    stage2_kernel: str,
    block_m: int,
) -> torch.Tensor:
    shape_key = _shape_key(hidden_states, topk_ids, config)
    num_experts = config.get("n_routed_experts", 0) + config.get("n_shared_experts", 0)
    token_num = hidden_states.shape[0]
    topk = topk_ids.shape[1]
    is_shuffled = getattr(gate_up_weight_shuffled, "is_shuffled", False)
    _, model_dim, inter_dim = _aiter_fused_moe.get_inter_dim(
        gate_up_weight_shuffled.shape,
        down_weight_shuffled.shape,
    )

    _log_once(
        ("manual2stage-enter", shape_key),
        "[submission] manual fused 2stage enter "
        f"shape={shape_key} block_m={block_m} prequant=fused_quant_sort "
        f"kernel1='{stage1_kernel}' kernel2='{stage2_kernel}'",
    )
    moebuf_ring_slots = _KMOE_MOEBUF_RING_WORKSPACE_SLOTS.get(shape_key, 1)
    if moebuf_ring_slots > 1:
        _log_once(
            ("manual2stage-moebuf-ring-workspace", shape_key, moebuf_ring_slots),
            "[submission] manual fused 2stage moebuf ring workspace active "
            f"shape={shape_key} slots={moebuf_ring_slots}",
        )
    _patch_moe_dtype_aliases()
    use_exact_sort_cache = shape_key in _KMOE_FULL_EXACT_SORT_CACHE_SHAPES
    cache_hit = False
    hit_clone_sort_inputs = False
    miss_clone_sort_inputs = False
    miss_streak = 0
    skip_cache_store = False
    bypass_exact_sort_cache = False
    first_bypass_miss_streak: Optional[int] = None
    if use_exact_sort_cache:
        skip_store_after = _KMOE_FULL_EXACT_SORT_SKIP_STORE_AFTER_MISS_STREAK.get(
            shape_key
        )
        if skip_store_after is not None:
            first_bypass_miss_streak = skip_store_after + 2
        current_miss_streak = _KMOE_FULL_EXACT_SORT_MISS_STREAK.get(shape_key, 0)
        bypass_exact_sort_cache = (
            skip_store_after is not None and current_miss_streak > skip_store_after
        )
        if bypass_exact_sort_cache:
            miss_streak = current_miss_streak + 1
            _KMOE_FULL_EXACT_SORT_MISS_STREAK[shape_key] = miss_streak
            skip_cache_store = True
            cache_hit = False
            if _KMOE_FULL_EXACT_SORT_CACHE:
                _KMOE_FULL_EXACT_SORT_CACHE.clear()
            _log_once(
                ("kmoe-full-exact-sort-cache-bypass", shape_key),
                "[submission] kmoe full exact sort cache bypass "
                f"shape={shape_key} miss_streak={miss_streak}",
            )
            (
                cached_sorted_ids,
                cached_sorted_weights,
                cached_sorted_expert_ids,
                cached_num_valid_ids,
                moe_buf,
            ) = _aiter_fused_moe.moe_sorting(
                topk_ids,
                topk_weights,
                num_experts,
                model_dim,
                hidden_states.dtype,
                block_m,
                None,
                None,
                0,
            )
            if shape_key in _KMOE_MISS_CLONE_SORT_INPUTS_SHAPES:
                sorted_ids = cached_sorted_ids.clone()
                sorted_weights = cached_sorted_weights.clone()
                sorted_expert_ids = cached_sorted_expert_ids.clone()
                num_valid_ids = cached_num_valid_ids.clone()
                miss_clone_sort_inputs = True
            else:
                sorted_ids = cached_sorted_ids
                sorted_weights = cached_sorted_weights
                sorted_expert_ids = cached_sorted_expert_ids
                num_valid_ids = cached_num_valid_ids
                if shape_key in _KMOE_CACHE_HIT_NOCLONE_SORT_INPUTS_SHAPES:
                    _log_once(
                        ("kmoe-miss-sort-inputs-noclone-bypass", shape_key),
                        "[submission] kmoe miss sort inputs noclone bypass "
                        f"shape={shape_key}",
                    )
        else:
            (
                cached_sorted_ids,
                cached_sorted_weights,
                cached_sorted_expert_ids,
                cached_num_valid_ids,
                cached_moe_buf,
                cache_hit,
                cache_key_to_store,
            ) = _get_kmoe_full_exact_sort_state(
                topk_ids,
                topk_weights,
                num_experts,
                model_dim,
                hidden_states.dtype,
                block_m,
            )
            if cache_hit:
                _KMOE_FULL_EXACT_SORT_MISS_STREAK[shape_key] = 0
                if shape_key in _KMOE_CACHE_HIT_NOCLONE_SORT_INPUTS_SHAPES:
                    sorted_ids = cached_sorted_ids
                    sorted_weights = cached_sorted_weights
                    sorted_expert_ids = cached_sorted_expert_ids
                    num_valid_ids = cached_num_valid_ids
                    _log_once(
                        ("kmoe-cache-hit-sort-inputs-noclone", shape_key),
                        "[submission] kmoe cache-hit sort inputs noclone "
                        f"shape={shape_key}",
                    )
                else:
                    sorted_ids = cached_sorted_ids.clone()
                    sorted_weights = cached_sorted_weights.clone()
                    sorted_expert_ids = cached_sorted_expert_ids.clone()
                    num_valid_ids = cached_num_valid_ids.clone()
                    hit_clone_sort_inputs = True
                moe_buf = _get_kmoe_moebuf_workspace(
                    token_num,
                    model_dim,
                    hidden_states.dtype,
                    hidden_states.device,
                    moebuf_ring_slots,
                )
            else:
                miss_streak = current_miss_streak + 1
                _KMOE_FULL_EXACT_SORT_MISS_STREAK[shape_key] = miss_streak
                if shape_key in _KMOE_MISS_CLONE_SORT_INPUTS_SHAPES:
                    sorted_ids = cached_sorted_ids.clone()
                    sorted_weights = cached_sorted_weights.clone()
                    sorted_expert_ids = cached_sorted_expert_ids.clone()
                    num_valid_ids = cached_num_valid_ids.clone()
                    miss_clone_sort_inputs = True
                else:
                    sorted_ids = cached_sorted_ids
                    sorted_weights = cached_sorted_weights
                    sorted_expert_ids = cached_sorted_expert_ids
                    num_valid_ids = cached_num_valid_ids
                    if shape_key in _KMOE_CACHE_HIT_NOCLONE_SORT_INPUTS_SHAPES:
                        _log_once(
                            ("kmoe-miss-sort-inputs-noclone", shape_key),
                            "[submission] kmoe miss sort inputs noclone "
                            f"shape={shape_key}",
                        )
                if cached_moe_buf is None:
                    raise RuntimeError("Exact-sort cache miss did not return moe_buf.")
                moe_buf = cached_moe_buf
                if cache_key_to_store is None:
                    raise RuntimeError("Exact-sort cache miss did not return cache key.")
                if skip_store_after is not None and miss_streak > skip_store_after:
                    skip_cache_store = True
                    if _KMOE_FULL_EXACT_SORT_CACHE:
                        _KMOE_FULL_EXACT_SORT_CACHE.clear()
                        _log_once(
                            (
                                "kmoe-full-exact-sort-cache-clear-skip-store",
                                shape_key,
                                miss_streak,
                            ),
                            "[submission] kmoe full exact sort cache clear "
                            f"shape={shape_key} miss_streak={miss_streak}",
                        )
                else:
                    _store_kmoe_full_exact_sort_cache(
                        shape_key,
                        cache_key_to_store,
                        (
                            cached_sorted_ids,
                            cached_sorted_weights,
                            cached_sorted_expert_ids,
                            cached_num_valid_ids,
                        ),
                    )
                    _log_once(
                        ("kmoe-full-exact-sort-cache-store-immediate", cache_key_to_store),
                        "[submission] kmoe full exact sort cache immediate store "
                        f"shape={shape_key}",
                    )
        _log_once(
            ("manual2stage-cached-full-exact-sort", shape_key, cache_hit, hit_clone_sort_inputs),
            "[submission] manual fused 2stage cached full exact sort "
            f"shape={shape_key} clone_inputs={int(hit_clone_sort_inputs)} "
            f"cache_hit={int(cache_hit)} use_sort_moebuf={int(not cache_hit)} "
            f"miss_clone_sort_inputs={int(miss_clone_sort_inputs)} "
            f"miss_streak={miss_streak} skip_cache_store={int(skip_cache_store)} "
            f"bypass_cache={int(bypass_exact_sort_cache)}",
        )
    else:
        sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = (
            _aiter_fused_moe.moe_sorting(
                topk_ids,
                topk_weights,
                num_experts,
                model_dim,
                hidden_states.dtype,
                block_m,
                None,
                None,
                0,
            )
        )
    _log_once(
        ("manual2stage-sorted", shape_key),
        f"[submission] manual fused 2stage sorted shape={shape_key}",
    )
    stage2_boundary_contract_reason: Optional[str] = None
    enable_stage2_boundary_contract = False
    if shape_key in _MANUAL_FUSED_STAGE2_BOUNDARY_CONTRACT_SHAPES and not cache_hit:
        if not bypass_exact_sort_cache and miss_streak == 1:
            enable_stage2_boundary_contract = True
            stage2_boundary_contract_reason = "first_miss"
        elif (
            bypass_exact_sort_cache
            and first_bypass_miss_streak is not None
            and miss_streak == first_bypass_miss_streak
        ):
            enable_stage2_boundary_contract = True
            stage2_boundary_contract_reason = "first_bypass"
    if enable_stage2_boundary_contract:
        stage2_boundary_workspace_mode = (
            "fresh"
            if stage2_boundary_contract_reason
            in _KMOE_STAGE2_BOUNDARY_CONTRACT_FRESH_REASONS
            else "cached"
        )
        _log_once(
            (
                "manual2stage-stage2-boundary-contract-miss-side-enabled",
                shape_key,
                stage2_boundary_contract_reason,
            ),
            "[submission] manual fused 2stage stage2-boundary contract miss-side enabled "
            f"shape={shape_key} reason={stage2_boundary_contract_reason} "
            f"workspace_mode={stage2_boundary_workspace_mode} "
            f"miss_streak={miss_streak} bypass_cache={int(bypass_exact_sort_cache)}",
        )
        stage2_boundary_contract = _get_stage2_boundary_contract(
            shape_key,
            token_num,
            topk,
            inter_dim,
            sorted_ids.shape[0],
            hidden_states.device,
            stage2_boundary_contract_reason or "default",
        )
    else:
        if shape_key in _MANUAL_FUSED_STAGE2_BOUNDARY_CONTRACT_SHAPES:
            _log_once(
                ("manual2stage-stage2-boundary-contract-miss-side-disabled", shape_key),
                "[submission] manual fused 2stage stage2-boundary contract miss-side disabled "
                f"shape={shape_key} cache_hit={int(cache_hit)} miss_streak={miss_streak} "
                f"bypass_cache={int(bypass_exact_sort_cache)}",
            )
        stage2_boundary_contract = None
    flydsl_stage2_patch_mode = "rawbuf" if stage2_boundary_contract is not None else "ptr"
    if stage2_kernel.startswith("flydsl_"):
        _log_once(
            (
                "manual2stage-stage2-patch-mode",
                shape_key,
                flydsl_stage2_patch_mode,
                bool(stage2_boundary_contract),
            ),
            "[submission] manual fused 2stage stage2 patch mode "
            f"shape={shape_key} mode={flydsl_stage2_patch_mode} "
            f"boundary_contract={int(stage2_boundary_contract is not None)} "
            f"cache_hit={int(cache_hit)} miss_streak={miss_streak} "
            f"bypass_cache={int(bypass_exact_sort_cache)}",
        )
    a1, a1_scale = _run_manual_fused_quant_sort_with_override(
        shape_key,
        "a1",
        hidden_states,
        sorted_ids,
        num_valid_ids,
        token_num,
        1,
        block_m,
    )

    metadata = _aiter_fused_moe.get_2stage_cfgs(
        _aiter_fused_moe.get_padded_M(token_num),
        model_dim,
        inter_dim,
        num_experts,
        topk,
        moe_buf.dtype,
        dtypes.fp4x2,
        dtypes.fp4x2,
        QuantType.per_1x32,
        True,
        ActivationType.Silu,
        False,
        hidden_pad,
        intermediate_pad,
        is_shuffled,
    )
    _log_once(
        ("manual2stage-stage1-enter", shape_key),
        f"[submission] manual fused 2stage stage1 enter shape={shape_key} ksplit={metadata.ksplit}",
    )
    a2 = torch.empty(
        (token_num, topk, inter_dim),
        dtype=moe_buf.dtype,
        device=hidden_states.device,
    )
    try:
        a2 = metadata.stage1(
            a1,
            gate_up_weight_shuffled,
            down_weight_shuffled,
            sorted_ids,
            sorted_expert_ids,
            num_valid_ids,
            a2,
            topk,
            block_m=block_m,
            a1_scale=a1_scale,
            w1_scale=(
                gate_up_weight_scale_shuffled.view(dtypes.fp8_e8m0)
                if gate_up_weight_shuffled.dtype == dtypes.fp4x2
                else gate_up_weight_scale_shuffled
            ),
            sorted_weights=None,
        )
    except Exception as exc:
        _log_once(
            ("manual2stage-stage1-failure", shape_key),
            "[submission] manual fused 2stage stage1 failure "
            f"shape={shape_key} exc={type(exc).__name__}: {exc}",
        )
        raise

    _log_once(
        ("manual2stage-stage1-success", shape_key),
        f"[submission] manual fused 2stage stage1 success shape={shape_key}",
    )

    if metadata.ksplit > 1 and is_shuffled:
        a2_scale = None
    else:
        _log_once(
            ("manual2stage-quant2-enter", shape_key),
            f"[submission] manual fused 2stage quant2 enter shape={shape_key}",
        )
        a2, a2_scale = _run_manual_fused_quant_sort_with_override(
            shape_key,
            "a2",
            a2.view(-1, inter_dim),
            sorted_ids,
            num_valid_ids,
            token_num,
            topk,
            block_m,
            (
                None
                if stage2_boundary_contract is None
                else stage2_boundary_contract["a2_fp4_u8"]
            ),
            (
                None
                if stage2_boundary_contract is None
                else stage2_boundary_contract["a2_scale_u8"]
            ),
        )
        a2 = a2.view(token_num, topk, -1)

    _log_once(
        ("manual2stage-stage2-enter", shape_key),
        f"[submission] manual fused 2stage stage2 enter shape={shape_key}",
    )
    try:
        _call_with_flydsl_stage2_patch_mode(
            metadata.stage2,
            flydsl_stage2_patch_mode,
            a2,
            gate_up_weight_shuffled,
            down_weight_shuffled,
            sorted_ids,
            sorted_expert_ids,
            num_valid_ids,
            moe_buf,
            topk,
            w2_scale=(
                down_weight_scale_shuffled.view(dtypes.fp8_e8m0)
                if down_weight_shuffled.dtype == dtypes.fp4x2
                else down_weight_scale_shuffled
            ),
            a2_scale=a2_scale,
            block_m=block_m,
            sorted_weights=sorted_weights,
        )
    except Exception as exc:
        _log_once(
            ("manual2stage-stage2-failure", shape_key),
            "[submission] manual fused 2stage stage2 failure "
            f"shape={shape_key} exc={type(exc).__name__}: {exc}",
        )
        raise

    out = moe_buf
    _log_once(
        ("manual2stage-stage2-success", shape_key),
        f"[submission] manual fused 2stage stage2 success shape={shape_key}",
    )
    _log_once(
        ("manual2stage-success", shape_key),
        f"[submission] manual fused 2stage success shape={shape_key}",
    )
    return out


def _can_split_shared_expert(
    topk_ids: torch.Tensor,
    config: Dict[str, int],
) -> bool:
    return (
        config.get("n_routed_experts", 0) == 32
        and config.get("n_shared_experts", 0) == 1
        and config.get("n_experts_per_token", 0) == 8
        and topk_ids.shape[1] == config.get("total_top_k", topk_ids.shape[1]) == 9
    )


def _run_split_shared_expert(
    hidden_states: torch.Tensor,
    gate_up_weight_shuffled: torch.Tensor,
    down_weight_shuffled: torch.Tensor,
    gate_up_weight_scale: torch.Tensor,
    down_weight_scale: 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[str, int],
    hidden_pad: int,
    intermediate_pad: int,
) -> torch.Tensor:
    routed_experts = config["n_routed_experts"]
    shared_experts = config["n_shared_experts"]
    routed_topk = config["n_experts_per_token"]

    hidden_q, hidden_scale = _get_hidden_prequant(hidden_states)

    routed_gate_scale_shuffled = _maybe_contiguous(
        fp4_utils.e8m0_shuffle(_maybe_contiguous(gate_up_weight_scale[:routed_experts]))
    )
    routed_down_scale_shuffled = _maybe_contiguous(
        fp4_utils.e8m0_shuffle(_maybe_contiguous(down_weight_scale[:routed_experts]))
    )
    shared_gate_scale_shuffled = _maybe_contiguous(
        fp4_utils.e8m0_shuffle(
            _maybe_contiguous(
                gate_up_weight_scale[routed_experts : routed_experts + shared_experts]
            )
        )
    )
    shared_down_scale_shuffled = _maybe_contiguous(
        fp4_utils.e8m0_shuffle(
            _maybe_contiguous(
                down_weight_scale[routed_experts : routed_experts + shared_experts]
            )
        )
    )

    routed_out = _run_fused_prequant(
        hidden_q,
        hidden_scale,
        _maybe_contiguous(gate_up_weight_shuffled[:routed_experts]),
        _maybe_contiguous(down_weight_shuffled[:routed_experts]),
        routed_gate_scale_shuffled,
        routed_down_scale_shuffled,
        _maybe_contiguous(topk_weights[:, :routed_topk]),
        _maybe_contiguous(topk_ids[:, :routed_topk]),
        hidden_pad,
        intermediate_pad,
    )

    shared_out = _run_fused_prequant(
        hidden_q,
        hidden_scale,
        _maybe_contiguous(
            gate_up_weight_shuffled[routed_experts : routed_experts + shared_experts]
        ),
        _maybe_contiguous(
            down_weight_shuffled[routed_experts : routed_experts + shared_experts]
        ),
        shared_gate_scale_shuffled,
        shared_down_scale_shuffled,
        _maybe_contiguous(topk_weights[:, routed_topk : routed_topk + shared_experts]),
        _maybe_contiguous(
            topk_ids[:, routed_topk : routed_topk + shared_experts] - routed_experts
        ),
        hidden_pad,
        intermediate_pad,
    )

    return (routed_out.float() + shared_out.float()).to(hidden_states.dtype)


@torch.inference_mode()
def generate_input(
    dhidden: int,
    dexpert: int,
    nroutedexperts: int,
    nexpertspertoken: int,
    nsharedexperts: int,
    bs: int,
    seed: int,
) -> input_t:
    _require_aiter()

    d_hidden = dhidden
    d_expert = dexpert
    n_routed_experts = nroutedexperts
    n_shared_experts = nsharedexperts
    routed_top_k = nexpertspertoken
    total_top_k = routed_top_k + n_shared_experts
    e_total = n_routed_experts + n_shared_experts
    m = bs

    d_hidden_pad = _pad_to(d_hidden, PAD_ALIGN)
    d_expert_pad = _pad_to(d_expert, PAD_ALIGN)

    config = {
        "d_hidden": d_hidden,
        "d_expert": d_expert,
        "d_hidden_pad": d_hidden_pad,
        "d_expert_pad": d_expert_pad,
        "n_routed_experts": n_routed_experts,
        "n_shared_experts": n_shared_experts,
        "n_experts_per_token": routed_top_k,
        "total_top_k": total_top_k,
        "bs": m,
    }

    gen = torch.Generator(device="cuda")
    gen.manual_seed(seed)

    hidden_states = torch.randn(
        (m, d_hidden),
        device="cuda",
        dtype=torch.bfloat16,
        generator=gen,
    )

    router_weight = torch.randn(
        (n_routed_experts, d_hidden),
        device="cuda",
        dtype=torch.bfloat16,
        generator=gen,
    ) / math.sqrt(d_hidden)
    router_logits = F.linear(hidden_states, router_weight)
    scores = router_logits.softmax(dim=-1)
    routed_weights, routed_ids = torch.topk(scores, k=routed_top_k, dim=-1, sorted=False)
    routed_weights = routed_weights.to(torch.float32)
    routed_ids = routed_ids.to(torch.int32)

    shared_ids = torch.arange(
        n_routed_experts,
        e_total,
        device="cuda",
        dtype=torch.int32,
    ).unsqueeze(0).expand(m, -1)
    shared_weights = torch.ones(
        (m, n_shared_experts),
        device="cuda",
        dtype=torch.float32,
    )

    topk_ids = torch.cat([routed_ids, shared_ids], dim=-1)
    topk_weights = torch.cat([routed_weights, shared_weights], dim=-1)

    gate_up_bf16 = torch.randn(
        (e_total, 2 * d_expert_pad, d_hidden_pad),
        device="cuda",
        dtype=torch.bfloat16,
        generator=gen,
    ) / math.sqrt(d_hidden)
    down_bf16 = torch.randn(
        (e_total, d_hidden_pad, d_expert_pad),
        device="cuda",
        dtype=torch.bfloat16,
        generator=gen,
    ) / math.sqrt(d_expert)

    torch_quant = aiter.get_torch_quant(QuantType.per_1x32)
    gate_up_weight, gate_up_weight_scale = torch_quant(
        gate_up_bf16,
        quant_dtype=dtypes.fp4x2,
    )
    down_weight, down_weight_scale = torch_quant(
        down_bf16,
        quant_dtype=dtypes.fp4x2,
    )

    gate_up_weight = gate_up_weight.view(e_total, 2 * d_expert_pad, d_hidden_pad // 2)
    down_weight = down_weight.view(e_total, d_hidden_pad, d_expert_pad // 2)

    gate_up_weight_shuffled = shuffle_weight(gate_up_weight, layout=SHUFFLE_LAYOUT)
    down_weight_shuffled = shuffle_weight(down_weight, layout=SHUFFLE_LAYOUT)
    gate_up_weight_scale_shuffled = fp4_utils.e8m0_shuffle(gate_up_weight_scale)
    down_weight_scale_shuffled = fp4_utils.e8m0_shuffle(down_weight_scale)

    return (
        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,
    )


@torch.inference_mode()
def ref_kernel(data: input_t) -> output_t:
    (
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        _gate_up_weight_scale,
        _down_weight_scale,
        gate_up_weight_scale_shuffled,
        down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        _config,
        hidden_pad,
        intermediate_pad,
    ) = _prepare_common(data)

    return _run_fused(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        gate_up_weight_scale_shuffled,
        down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        hidden_pad,
        intermediate_pad,
    )


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    (
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        gate_up_weight_scale,
        down_weight_scale,
        gate_up_weight_scale_shuffled,
        down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        config,
        hidden_pad,
        intermediate_pad,
    ) = _prepare_common(data)

    shape_key = _shape_key(hidden_states, topk_ids, config)
    _maybe_reset_full_exact_sort_cache(shape_key)
    cache_key = _backend_key(hidden_states, topk_ids, config)
    cached_backend = _BACKEND_CACHE.get(cache_key)
    forced_non_temporal_load = _forced_ck_non_temporal_load(hidden_states, topk_ids, config)
    forced_stage_kernels = _forced_ck_stage_kernels(hidden_states, topk_ids, config)
    fused_stage_override = _FUSED_CFG_STAGE_OVERRIDES.get(shape_key)
    use_manual_fused_2stage = _should_use_manual_fused_2stage(
        hidden_states,
        topk_ids,
        config,
    )
    use_ck_prequant = _should_prequant_ck(hidden_states, topk_ids, config)
    on_mi355x = _device_is_mi355x(hidden_states.device)
    full_topk = topk_ids.shape[1] == config.get("total_top_k", topk_ids.shape[1])

    if cached_backend is not None:
        if cached_backend[0] == "split_fused":
            try:
                return _run_split_shared_expert(
                    hidden_states,
                    gate_up_weight_shuffled,
                    down_weight_shuffled,
                    gate_up_weight_scale,
                    down_weight_scale,
                    gate_up_weight_scale_shuffled,
                    down_weight_scale_shuffled,
                    topk_weights,
                    topk_ids,
                    config,
                    hidden_pad,
                    intermediate_pad,
                )
            except Exception:
                _BACKEND_CACHE[cache_key] = ("fused", None)
        elif cached_backend[0] == "manual2stage":
            try:
                if fused_stage_override is None:
                    raise RuntimeError("Missing fused stage override for manual route.")
                return _run_manual_fused_2stage_prequant(
                    hidden_states,
                    gate_up_weight_shuffled,
                    down_weight_shuffled,
                    gate_up_weight_scale_shuffled,
                    down_weight_scale_shuffled,
                    topk_weights,
                    topk_ids,
                    config,
                    hidden_pad,
                    intermediate_pad,
                    fused_stage_override[0],
                    fused_stage_override[1],
                    fused_stage_override[2],
                )
            except Exception as exc:
                _log_once(
                    ("manual2stage-fallback", shape_key),
                    "[submission] manual fused 2stage fallback "
                    f"shape={shape_key} exc={type(exc).__name__}: {exc}",
                )
                _BACKEND_CACHE[cache_key] = ("fused", None)
        elif cached_backend[0] == "ck_stage":
            try:
                if forced_stage_kernels is None:
                    raise RuntimeError("Missing CK stage-kernel override.")
                return _run_ck_prequant_with_stage_kernels(
                    hidden_states,
                    gate_up_weight_shuffled,
                    down_weight_shuffled,
                    gate_up_weight_scale_shuffled,
                    down_weight_scale_shuffled,
                    topk_weights,
                    topk_ids,
                    forced_stage_kernels[2],
                    forced_non_temporal_load,
                    forced_stage_kernels[0],
                    forced_stage_kernels[1],
                )
            except Exception:
                _BACKEND_CACHE[cache_key] = ("ck", cached_backend[1])
                return _run_ck_prequant(
                    hidden_states,
                    gate_up_weight_shuffled,
                    down_weight_shuffled,
                    gate_up_weight_scale_shuffled,
                    down_weight_scale_shuffled,
                    topk_weights,
                    topk_ids,
                    cached_backend[1],
                    forced_non_temporal_load,
                )
        elif cached_backend[0] == "ck":
            try:
                if use_ck_prequant:
                    return _run_ck_prequant(
                        hidden_states,
                        gate_up_weight_shuffled,
                        down_weight_shuffled,
                        gate_up_weight_scale_shuffled,
                        down_weight_scale_shuffled,
                        topk_weights,
                        topk_ids,
                        cached_backend[1],
                        forced_non_temporal_load,
                    )
                return _run_ck(
                    hidden_states,
                    gate_up_weight_shuffled,
                    down_weight_shuffled,
                    gate_up_weight_scale_shuffled,
                    down_weight_scale_shuffled,
                    topk_weights,
                    topk_ids,
                    cached_backend[1],
                    forced_non_temporal_load,
                )
            except Exception:
                _BACKEND_CACHE[cache_key] = ("fused", None)
        return _run_fused(
            hidden_states,
            gate_up_weight_shuffled,
            down_weight_shuffled,
            gate_up_weight_scale_shuffled,
            down_weight_scale_shuffled,
            topk_weights,
            topk_ids,
            hidden_pad,
            intermediate_pad,
        )

    if _can_split_shared_expert(topk_ids, config) and shape_key not in _SKIP_SPLIT_SHARED_EXPERT_SHAPES:
        try:
            _BACKEND_CACHE[cache_key] = ("split_fused", None)
            return _run_split_shared_expert(
                hidden_states,
                gate_up_weight_shuffled,
                down_weight_shuffled,
                gate_up_weight_scale,
                down_weight_scale,
                gate_up_weight_scale_shuffled,
                down_weight_scale_shuffled,
                topk_weights,
                topk_ids,
                config,
                hidden_pad,
                intermediate_pad,
            )
        except Exception:
            _BACKEND_CACHE[cache_key] = ("fused", None)

    if use_manual_fused_2stage and fused_stage_override is not None and on_mi355x and full_topk:
        try:
            _BACKEND_CACHE[cache_key] = ("manual2stage", None)
            return _run_manual_fused_2stage_prequant(
                hidden_states,
                gate_up_weight_shuffled,
                down_weight_shuffled,
                gate_up_weight_scale_shuffled,
                down_weight_scale_shuffled,
                topk_weights,
                topk_ids,
                config,
                hidden_pad,
                intermediate_pad,
                fused_stage_override[0],
                fused_stage_override[1],
                fused_stage_override[2],
            )
        except Exception as exc:
            _log_once(
                ("manual2stage-fallback", shape_key),
                "[submission] manual fused 2stage fallback "
                f"shape={shape_key} exc={type(exc).__name__}: {exc}",
            )
            _BACKEND_CACHE[cache_key] = ("fused", None)

    if fused_stage_override is not None and on_mi355x and full_topk:
        try:
            _BACKEND_CACHE[cache_key] = ("fused", None)
            return _run_fused(
                hidden_states,
                gate_up_weight_shuffled,
                down_weight_shuffled,
                gate_up_weight_scale_shuffled,
                down_weight_scale_shuffled,
                topk_weights,
                topk_ids,
                hidden_pad,
                intermediate_pad,
            )
        except Exception:
            _BACKEND_CACHE.pop(cache_key, None)

    benchmark_block_m = _BENCHMARK_BLOCK_M.get(shape_key)
    if (
        _CK_MOE is not None
        and benchmark_block_m is not None
        and on_mi355x
        and full_topk
    ):
        if forced_stage_kernels is not None and use_ck_prequant:
            try:
                _BACKEND_CACHE[cache_key] = ("ck_stage", forced_stage_kernels[2])
                return _run_ck_prequant_with_stage_kernels(
                    hidden_states,
                    gate_up_weight_shuffled,
                    down_weight_shuffled,
                    gate_up_weight_scale_shuffled,
                    down_weight_scale_shuffled,
                    topk_weights,
                    topk_ids,
                    forced_stage_kernels[2],
                    forced_non_temporal_load,
                    forced_stage_kernels[0],
                    forced_stage_kernels[1],
                )
            except Exception:
                _BACKEND_CACHE[cache_key] = ("ck", benchmark_block_m)
        try:
            _BACKEND_CACHE[cache_key] = ("ck", benchmark_block_m)
            if use_ck_prequant:
                return _run_ck_prequant(
                    hidden_states,
                    gate_up_weight_shuffled,
                    down_weight_shuffled,
                    gate_up_weight_scale_shuffled,
                    down_weight_scale_shuffled,
                    topk_weights,
                    topk_ids,
                    benchmark_block_m,
                    forced_non_temporal_load,
                )
            return _run_ck(
                hidden_states,
                gate_up_weight_shuffled,
                down_weight_shuffled,
                gate_up_weight_scale_shuffled,
                down_weight_scale_shuffled,
                topk_weights,
                topk_ids,
                benchmark_block_m,
                    forced_non_temporal_load,
                )
        except Exception:
            _BACKEND_CACHE[cache_key] = ("fused", None)

    _BACKEND_CACHE[cache_key] = ("fused", None)
    return _run_fused(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        gate_up_weight_scale_shuffled,
        down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        hidden_pad,
        intermediate_pad,
    )


kernel = custom_kernel
custom_generate_input = generate_input
create_input = generate_input
reference_kernel = ref_kernel

if make_match_reference is not None:
    check_implementation = make_match_reference(custom_kernel, rtol=RTOL, atol=ATOL)
else:
    check_implementation = None


import functools
import aiter.fused_moe as _aiter_fused_moe


_FLYDSL_STAGE2_SOURCE_PROBED = False


def _probe_flydsl_stage2_builder_source() -> None:
    global _FLYDSL_STAGE2_SOURCE_PROBED
    if _FLYDSL_STAGE2_SOURCE_PROBED or aiter is None:
        return

    try:
        import ast as _ast
        import hashlib as _hashlib
        import inspect as _inspect
        import aiter.ops.flydsl.kernels.mixed_moe_gemm_2stage as _mixed_stage2
        import aiter.ops.flydsl.kernels.mfma_epilogues as _mfma_epilogues
        import aiter.ops.flydsl.moe_kernels as _flydsl_moe_kernels
    except Exception as exc:
        _log_once(
            ("flydsl-stage2-source-probe-import-failure", type(exc).__name__),
            "[submission] kmoe probe flydsl stage2 skipped "
            f"reason=import_failure exc={type(exc).__name__}",
        )
        return

    def _find_matches(lines, predicate):
        return [idx for idx, line in enumerate(lines) if predicate(line)]

    def _collect_line_numbers(lines, predicate):
        return [idx + 1 for idx in _find_matches(lines, predicate)]

    def _emit_contexts(
        digest: str,
        lines,
        label: str,
        predicate,
        radius: int = 6,
        max_matches: int = 1,
    ) -> None:
        matches = _find_matches(lines, predicate)
        if not matches:
            _log_once(
                ("flydsl-stage2-source-probe-context", digest, label, "missing"),
                "[submission] kmoe probe flydsl stage2 "
                f"context={label} status=missing",
            )
            return
        for match_idx, center in enumerate(matches[:max_matches], 1):
            start = max(0, center - radius)
            end = min(len(lines), center + radius + 1)
            snippet = "\n".join(
                f"{idx + 1}:{lines[idx]}" for idx in range(start, end)
            )
            _log_once(
                ("flydsl-stage2-source-probe-context", digest, label, center),
                "[submission] kmoe probe flydsl stage2 "
                f"context={label} idx={match_idx} line={center + 1}\n{snippet}",
            )

    def _emit_line_contexts(
        digest: str,
        lines,
        label: str,
        line_numbers,
        radius: int = 6,
        max_matches: int = 1,
    ) -> None:
        if not line_numbers:
            _log_once(
                ("flydsl-stage2-source-probe-line-context", digest, label, "missing"),
                "[submission] kmoe probe flydsl stage2 "
                f"context={label} status=missing",
            )
            return
        for match_idx, line_no in enumerate(line_numbers[:max_matches], 1):
            center = int(line_no) - 1
            start = max(0, center - radius)
            end = min(len(lines), center + radius + 1)
            snippet = "\n".join(
                f"{idx + 1}:{lines[idx]}" for idx in range(start, end)
            )
            _log_once(
                ("flydsl-stage2-source-probe-line-context", digest, label, center),
                "[submission] kmoe probe flydsl stage2 "
                f"context={label} idx={match_idx} line={line_no}\n{snippet}",
            )

    def _ast_name(node) -> str:
        if node is None:
            return "None"
        if isinstance(node, _ast.Name):
            return node.id
        if isinstance(node, _ast.Attribute):
            base = _ast_name(node.value)
            return f"{base}.{node.attr}" if base else node.attr
        try:
            return _ast.unparse(node)
        except Exception:
            return type(node).__name__

    def _probe_ast(label: str, digest: str, source_text: str):
        try:
            module_ast = _ast.parse(source_text)
        except SyntaxError as exc:
            _log_once(
                (
                    "flydsl-stage2-source-probe-ast-failure",
                    digest,
                    label,
                    type(exc).__name__,
                ),
                "[submission] kmoe probe flydsl stage2 ast "
                f"target={label} reason=parse_failure exc={type(exc).__name__}",
            )
            return {}, {}

        parents = {}
        for parent in _ast.walk(module_ast):
            for child in _ast.iter_child_nodes(parent):
                parents[id(child)] = parent

        def _scope(node) -> str:
            parts = []
            cur = parents.get(id(node))
            while cur is not None:
                if isinstance(
                    cur, (_ast.FunctionDef, _ast.AsyncFunctionDef, _ast.ClassDef)
                ):
                    parts.append(cur.name)
                cur = parents.get(id(cur))
            parts.reverse()
            return "/".join(parts) if parts else "<module>"

        interesting_defs = {
            "compile_mixed_moe_gemm2": {
                "compile_mixed_moe_gemm2",
                "moe_gemm2",
                "write_row_to_lds",
                "precompute_row",
                "store_pair",
            },
            "c_shuffle_epilog": {
                "c_shuffle_epilog",
                "_write_row",
                "_do_store_row",
            },
            "default_epilog": {
                "default_epilog",
            },
        }.get(label, set())
        interesting_calls = {
            "c_shuffle_epilog",
            "default_epilog",
            "precompute_row",
            "write_row_to_lds",
            "store_pair",
        }

        def_lines = {}
        call_infos = []
        for node in _ast.walk(module_ast):
            if isinstance(node, (_ast.FunctionDef, _ast.AsyncFunctionDef, _ast.ClassDef)):
                if node.name in interesting_defs:
                    def_lines.setdefault(node.name, []).append(node.lineno)
                continue
            if not isinstance(node, _ast.Call):
                continue
            func_name = _ast_name(node.func)
            short_name = func_name.rsplit(".", 1)[-1]
            if short_name not in interesting_calls:
                continue
            kwargs = []
            for kw in node.keywords:
                if kw.arg is None:
                    continue
                kwargs.append(f"{kw.arg}={_ast_name(kw.value)}")
            call_infos.append(
                {
                    "name": short_name,
                    "func": func_name,
                    "line": node.lineno,
                    "scope": _scope(node),
                    "kwargs": kwargs,
                }
            )

        call_lines = {}
        for info in call_infos:
            call_lines.setdefault(info["name"], []).append(info["line"])

        def_summary = ",".join(
            f"{name}:{def_lines[name]}" for name in sorted(def_lines)
        ) or "-"
        call_summary = ",".join(
            f"{name}:{call_lines[name]}" for name in sorted(call_lines)
        ) or "-"
        _log_once(
            ("flydsl-stage2-source-probe-ast-summary", digest, label),
            "[submission] kmoe probe flydsl stage2 ast "
            f"target={label} defs={def_summary} calls={call_summary}",
        )
        for info in call_infos:
            kw_text = "|".join(info["kwargs"]) if info["kwargs"] else "-"
            _log_once(
                (
                    "flydsl-stage2-source-probe-ast-call",
                    digest,
                    label,
                    info["name"],
                    info["line"],
                ),
                "[submission] kmoe probe flydsl stage2 ast "
                f"target={label} call={info['name']} func={info['func']} "
                f"line={info['line']} scope={info['scope']} kwargs={kw_text}",
            )
        return def_lines, call_lines

    targets = [
        ("compile_mixed_moe_gemm2", _mixed_stage2.compile_mixed_moe_gemm2),
    ]

    for label, fn in targets:
        target_fn = _inspect.unwrap(fn)
        try:
            source = _inspect.getsource(target_fn)
            source_file = _inspect.getsourcefile(target_fn) or "unknown"
            firstlineno = getattr(
                getattr(target_fn, "__code__", None), "co_firstlineno", -1
            )
        except (OSError, TypeError) as exc:
            _log_once(
                ("flydsl-stage2-source-probe-source-failure", label, type(exc).__name__),
                "[submission] kmoe probe flydsl stage2 "
                f"target={label} reason=source_failure exc={type(exc).__name__}",
            )
            continue

        lines = source.splitlines()
        normalized = "\n".join(line.rstrip() for line in lines).strip()
        digest = _hashlib.sha1(normalized.encode("utf-8")).hexdigest()[:16]
        ast_def_lines = {}
        ast_call_lines = {}
        if label in ("compile_mixed_moe_gemm2", "c_shuffle_epilog", "default_epilog"):
            ast_def_lines, ast_call_lines = _probe_ast(label, digest, source)

        if label == "compile_mixed_moe_gemm2":
            module_name_lines = _collect_line_numbers(
                lines, lambda line: "module_name" in line
            )
            lds_out_bytes_lines = _collect_line_numbers(
                lines, lambda line: "lds_out_bytes" in line
            )
            lds_total_bytes_lines = _collect_line_numbers(
                lines, lambda line: "lds_total_bytes" in line
            )
            lds_alloc_bytes_lines = _collect_line_numbers(
                lines, lambda line: "lds_alloc_bytes" in line
            )
            lds_x_decl_lines = _collect_line_numbers(
                lines, lambda line: '_state["lds_x_decl"]' in line
            )
            allocate_array_lines = _collect_line_numbers(
                lines, lambda line: "allocator.allocate_array(" in line
            )
            base_ptr_lines = _collect_line_numbers(
                lines, lambda line: "base_ptr" in line
            )
            lds_x_ptr_lines = _collect_line_numbers(
                lines, lambda line: "lds_x_ptr" in line
            )
            lds_out_lines = _collect_line_numbers(
                lines, lambda line: "lds_out" in line
            )
            smem_ptr_lines = _collect_line_numbers(
                lines, lambda line: "SmemPtr(" in line
            )
            write_row_lines = _collect_line_numbers(
                lines, lambda line: "def write_row_to_lds" in line
            )
            write_row_call_lines = _collect_line_numbers(
                lines,
                lambda line: "write_row_to_lds(" in line
                and "def write_row_to_lds" not in line,
            )
            precompute_row_lines = _collect_line_numbers(
                lines, lambda line: "def precompute_row" in line
            )
            precompute_row_call_lines = _collect_line_numbers(
                lines,
                lambda line: "precompute_row(" in line
                and "def precompute_row" not in line,
            )
            fused2_assign_lines = _collect_line_numbers(
                lines, lambda line: "fused2 = buffer_ops.buffer_load(" in line
            )
            row_valid_lines = _collect_line_numbers(
                lines, lambda line: "row_valid" in line
            )
            vector_store_lds_out_lines = _collect_line_numbers(
                lines, lambda line: "vector.store(v1, lds_out" in line
            )
            sorted_rsrc_load_lines = _collect_line_numbers(
                lines,
                lambda line: "sorted_rsrc" in line
                and "buffer_ops.buffer_load" in line,
            )
            sorted_w_rsrc_lines = _collect_line_numbers(
                lines, lambda line: "sorted_w_rsrc" in line
            )
            sorted_w_rsrc_load_lines = _collect_line_numbers(
                lines,
                lambda line: "sorted_w_rsrc" in line
                and "buffer_ops.buffer_load" in line,
            )
            doweight_stage2_lines = _collect_line_numbers(
                lines, lambda line: "doweight_stage2" in line
            )
            tw_assign_lines = _collect_line_numbers(
                lines, lambda line: "tw =" in line
            )
            tw_pf_use_lines = _collect_line_numbers(
                lines,
                lambda line: "tw_pf[" in line
                or "tw = tw_pf" in line
                or "tw_pf is not None" in line,
            )
            lds_tid_lines = _collect_line_numbers(
                lines, lambda line: "lds_tid" in line
            )
            memref_load_lds_tid_lines = _collect_line_numbers(
                lines, lambda line: "memref.load(lds_tid" in line
            )
            tw_pf_lines = _collect_line_numbers(lines, lambda line: "tw_pf" in line)
            use_cshuffle_epilog_lines = _collect_line_numbers(
                lines, lambda line: "_use_cshuffle_epilog" in line
            )
            ast_cshuffle_call_lines = ast_call_lines.get("c_shuffle_epilog", [])
            ast_store_pair_def_lines = ast_def_lines.get("store_pair", [])
            summary = (
                "[submission] kmoe probe flydsl stage2 "
                f"target={label} file={source_file} firstlineno={firstlineno} "
                f"sha1={digest} lines={len(lines)} "
                f"wrapper_type={type(fn).__name__} unwrapped_type={type(target_fn).__name__} "
                f"used_unwrap={int(target_fn is not fn)} "
                f"module_name={module_name_lines} "
                f"lds_out_bytes={lds_out_bytes_lines} "
                f"lds_total_bytes={lds_total_bytes_lines} "
                f"lds_alloc_bytes={lds_alloc_bytes_lines} "
                f"lds_x_decl={lds_x_decl_lines} "
                f"allocate_array={allocate_array_lines} "
                f"base_ptr={base_ptr_lines} "
                f"lds_x_ptr={lds_x_ptr_lines} "
                f"lds_out={lds_out_lines} "
                f"SmemPtr={smem_ptr_lines} "
                f"write_row={write_row_lines} "
                f"write_row_call={write_row_call_lines} "
                f"precompute_row={precompute_row_lines} "
                f"precompute_row_call={precompute_row_call_lines} "
                f"fused2_assign={fused2_assign_lines} "
                f"row_valid={row_valid_lines} "
                f"vector_store_lds_out={vector_store_lds_out_lines} "
                f"sorted_rsrc_load={sorted_rsrc_load_lines} "
                f"sorted_w_rsrc={sorted_w_rsrc_lines} "
                f"sorted_w_rsrc_load={sorted_w_rsrc_load_lines} "
                f"doweight_stage2={doweight_stage2_lines} "
                f"tw_assign={tw_assign_lines} "
                f"tw_pf_use={tw_pf_use_lines} "
                f"lds_tid={lds_tid_lines} "
                f"memref_load_lds_tid={memref_load_lds_tid_lines} "
                f"tw_pf={tw_pf_lines} "
                f"use_cshuffle_epilog={use_cshuffle_epilog_lines} "
                f"ast_c_shuffle_epilog_call={ast_cshuffle_call_lines} "
                f"ast_store_pair_def={ast_store_pair_def_lines}"
            )
            _log_once(("flydsl-stage2-source-probe-summary", digest), summary)
            for context_label, predicate, radius, max_matches in (
                ("module_name", lambda line: "module_name" in line, 6, 1),
                ("lds_out_bytes", lambda line: "lds_out_bytes" in line, 6, 1),
                ("lds_total_bytes", lambda line: "lds_total_bytes" in line, 6, 1),
                ("lds_alloc_bytes", lambda line: "lds_alloc_bytes" in line, 6, 1),
                ("lds_x_decl", lambda line: '_state["lds_x_decl"]' in line, 6, 1),
                ("allocate_array", lambda line: "allocator.allocate_array(" in line, 6, 1),
                ("base_ptr", lambda line: "base_ptr" in line, 6, 1),
                ("lds_x_ptr", lambda line: "lds_x_ptr" in line, 6, 1),
                ("lds_out", lambda line: "lds_out" in line, 6, 2),
                ("SmemPtr", lambda line: "SmemPtr(" in line, 6, 1),
                ("write_row_to_lds", lambda line: "def write_row_to_lds" in line, 8, 1),
                (
                    "write_row_call",
                    lambda line: "write_row_to_lds(" in line
                    and "def write_row_to_lds" not in line,
                    6,
                    2,
                ),
                ("precompute_row", lambda line: "def precompute_row" in line, 8, 1),
                (
                    "precompute_row_call",
                    lambda line: "precompute_row(" in line
                    and "def precompute_row" not in line,
                    6,
                    2,
                ),
                (
                    "fused2_assign",
                    lambda line: "fused2 = buffer_ops.buffer_load(" in line,
                    6,
                    1,
                ),
                ("row_valid", lambda line: "row_valid" in line, 4, 2),
                (
                    "vector_store_lds_out",
                    lambda line: "vector.store(v1, lds_out" in line,
                    6,
                    1,
                ),
                (
                    "sorted_rsrc_load",
                    lambda line: "sorted_rsrc" in line
                    and "buffer_ops.buffer_load" in line,
                    6,
                    1,
                ),
                ("sorted_w_rsrc", lambda line: "sorted_w_rsrc" in line, 6, 3),
                (
                    "sorted_w_rsrc_load",
                    lambda line: "sorted_w_rsrc" in line
                    and "buffer_ops.buffer_load" in line,
                    6,
                    3,
                ),
                ("doweight_stage2", lambda line: "doweight_stage2" in line, 6, 3),
                ("tw_assign", lambda line: "tw =" in line, 6, 3),
                (
                    "tw_pf_use",
                    lambda line: "tw_pf[" in line
                    or "tw = tw_pf" in line
                    or "tw_pf is not None" in line,
                    6,
                    3,
                ),
                ("memref_load_lds_tid", lambda line: "memref.load(lds_tid" in line, 6, 1),
                ("tw_pf", lambda line: "tw_pf" in line, 4, 2),
                ("use_cshuffle_epilog", lambda line: "_use_cshuffle_epilog" in line, 4, 2),
                ("c_shuffle_epilog", lambda line: "c_shuffle_epilog(" in line, 6, 1),
            ):
                _emit_contexts(
                    digest,
                    lines,
                    context_label,
                    predicate,
                    radius=radius,
                    max_matches=max_matches,
                )
            _emit_line_contexts(
                digest,
                lines,
                "ast_c_shuffle_epilog_call",
                ast_cshuffle_call_lines,
                radius=8,
                max_matches=1,
            )
            _emit_line_contexts(
                digest,
                lines,
                "ast_store_pair_def",
                ast_store_pair_def_lines,
                radius=8,
                max_matches=1,
            )
        elif label == "c_shuffle_epilog":
            precomputed_rows_lines = _collect_line_numbers(
                lines, lambda line: "_precomputed_rows = []" in line
            )
            row_ctx_raw_lines = _collect_line_numbers(
                lines, lambda line: "row_ctx_raw =" in line
            )
            summary = (
                "[submission] kmoe probe flydsl stage2 "
                f"target={label} file={source_file} firstlineno={firstlineno} "
                f"sha1={digest} lines={len(lines)} "
                f"wrapper_type={type(fn).__name__} unwrapped_type={type(target_fn).__name__} "
                f"used_unwrap={int(target_fn is not fn)} "
                f"precomputed_rows={precomputed_rows_lines} "
                f"row_ctx_raw={row_ctx_raw_lines} "
                f"ast__write_row={ast_def_lines.get('_write_row', [])} "
                f"ast__do_store_row={ast_def_lines.get('_do_store_row', [])} "
                f"ast_default_epilog_call={ast_call_lines.get('default_epilog', [])} "
                f"ast_precompute_row_call={ast_call_lines.get('precompute_row', [])} "
                f"ast_write_row_to_lds_call={ast_call_lines.get('write_row_to_lds', [])} "
                f"ast_store_pair_call={ast_call_lines.get('store_pair', [])}"
            )
            _log_once(
                ("flydsl-stage2-source-probe-target", label, digest),
                summary,
            )
            for context_label, predicate, radius, max_matches in (
                ("_write_row", lambda line: "def _write_row" in line, 6, 1),
                (
                    "default_epilog_call",
                    lambda line: "default_epilog(" in line,
                    6,
                    1,
                ),
                (
                    "precomputed_rows",
                    lambda line: "_precomputed_rows = []" in line,
                    6,
                    1,
                ),
                ("row_ctx_raw", lambda line: "row_ctx_raw =" in line, 6, 1),
                ("_do_store_row", lambda line: "def _do_store_row" in line, 6, 1),
                (
                    "store_pair_call",
                    lambda line: "store_pair(" in line
                    and "def store_pair" not in line,
                    6,
                    1,
                ),
            ):
                _emit_contexts(
                    digest,
                    lines,
                    context_label,
                    predicate,
                    radius=radius,
                    max_matches=max_matches,
                )
        elif label == "default_epilog":
            body_row_lines = _collect_line_numbers(
                lines, lambda line: "body_row(" in line
            )
            summary = (
                "[submission] kmoe probe flydsl stage2 "
                f"target={label} file={source_file} firstlineno={firstlineno} "
                f"sha1={digest} lines={len(lines)} "
                f"wrapper_type={type(fn).__name__} unwrapped_type={type(target_fn).__name__} "
                f"used_unwrap={int(target_fn is not fn)} "
                f"body_row={body_row_lines}"
            )
            _log_once(
                ("flydsl-stage2-source-probe-target", label, digest),
                summary,
            )
            _emit_contexts(
                digest,
                lines,
                "body_row_call",
                lambda line: "body_row(" in line,
                radius=6,
                max_matches=1,
            )
        else:
            _log_once(
                ("flydsl-stage2-source-probe-target", label, digest),
                "[submission] kmoe probe flydsl stage2 "
                f"target={label} file={source_file} firstlineno={firstlineno} "
                f"sha1={digest} lines={len(lines)} "
                f"wrapper_type={type(fn).__name__} unwrapped_type={type(target_fn).__name__} "
                f"used_unwrap={int(target_fn is not fn)}",
            )

    _FLYDSL_STAGE2_SOURCE_PROBED = True


if False:
    _probe_flydsl_stage2_builder_source()


_FLYDSL_STAGE2_SORTEDIDX_LDS_PATCHED = False


def _install_flydsl_stage2_multi_mode_compiled_cache(_flydsl_moe_kernels) -> None:
    global _KMOE_FLYDSL_STAGE2_ORIG_GET_COMPILED_STAGE2
    global _KMOE_FLYDSL_STAGE2_COMPILED_CACHE_PATCHED
    if _KMOE_FLYDSL_STAGE2_COMPILED_CACHE_PATCHED:
        return

    target_fn = getattr(_flydsl_moe_kernels, "_get_compiled_stage2", None)
    if target_fn is None:
        _log_once(
            ("flydsl-stage2-compiled-cache-install-missing-target",),
            "[submission] kmoe flydsl stage2 compiled cache install skipped "
            "reason=missing_get_compiled_stage2",
        )
        return

    _KMOE_FLYDSL_STAGE2_ORIG_GET_COMPILED_STAGE2 = target_fn

    def _patched_get_compiled_stage2(
        model_dim: int,
        inter_dim: int,
        experts: int,
        topk: int,
        tile_m: int,
        tile_n: int,
        tile_k: int,
        doweight: bool,
        a_dtype: str,
        b_dtype: str,
        out_dtype: str,
        accumulate: bool = True,
    ):
        patch_mode = (
            _KMOE_ACTIVE_FLYDSL_STAGE2_PATCH_MODE
            or _KMOE_PENDING_FLYDSL_STAGE2_PATCH_MODE
        )
        cache_key = (
            patch_mode,
            model_dim,
            inter_dim,
            experts,
            topk,
            tile_m,
            tile_n,
            tile_k,
            doweight,
            a_dtype,
            b_dtype,
            out_dtype,
            accumulate,
        )
        cached = _KMOE_FLYDSL_STAGE2_COMPILED_CACHE.get(cache_key)
        if cached is not None:
            _log_once(
                ("flydsl-stage2-compiled-cache-hit", cache_key),
                "[submission] kmoe flydsl stage2 compiled cache hit "
                f"patch_mode={patch_mode} tile={tile_m}x{tile_n}x{tile_k} "
                f"accumulate={int(accumulate)}",
            )
            return cached

        try:
            target_fn.cache_clear()
        except Exception:
            pass
        compiled = target_fn(
            model_dim=model_dim,
            inter_dim=inter_dim,
            experts=experts,
            topk=topk,
            tile_m=tile_m,
            tile_n=tile_n,
            tile_k=tile_k,
            doweight=doweight,
            a_dtype=a_dtype,
            b_dtype=b_dtype,
            out_dtype=out_dtype,
            accumulate=accumulate,
        )
        _KMOE_FLYDSL_STAGE2_COMPILED_CACHE[cache_key] = compiled
        _log_once(
            ("flydsl-stage2-compiled-cache-build", cache_key),
            "[submission] kmoe flydsl stage2 compiled cache build "
            f"patch_mode={patch_mode} tile={tile_m}x{tile_n}x{tile_k} "
            f"accumulate={int(accumulate)}",
        )
        return compiled

    _flydsl_moe_kernels._get_compiled_stage2 = _patched_get_compiled_stage2
    _KMOE_FLYDSL_STAGE2_COMPILED_CACHE_PATCHED = True
    _log_once(
        ("flydsl-stage2-compiled-cache-install",),
        "[submission] kmoe flydsl stage2 compiled cache multi-mode install",
    )


def _patch_flydsl_stage2_sortedidx_lds_prefetch(patch_mode: str = "ptr") -> None:
    global _FLYDSL_STAGE2_SORTEDIDX_LDS_PATCHED, _KMOE_ACTIVE_FLYDSL_STAGE2_PATCH_MODE
    global _KMOE_FLYDSL_STAGE2_ORIG_SOURCE, _KMOE_FLYDSL_STAGE2_ORIG_SOURCE_FILE
    if aiter is None:
        return
    if patch_mode not in ("ptr", "rawbuf"):
        raise ValueError(f"unsupported flydsl stage2 patch mode: {patch_mode}")
    if _KMOE_ACTIVE_FLYDSL_STAGE2_PATCH_MODE == patch_mode:
        return

    try:
        import ast as _ast
        import hashlib as _hashlib
        import inspect as _inspect
        import aiter.ops.flydsl.kernels.mixed_moe_gemm_2stage as _mixed_stage2
        import aiter.ops.flydsl.moe_kernels as _flydsl_moe_kernels
    except Exception as exc:
        _log_once(
            ("flydsl-stage2-sortedidx-lds-patch-import-failure", type(exc).__name__),
            "[submission] kmoe patch flydsl stage2 sortedidx lds skipped "
            f"reason=import_failure exc={type(exc).__name__}",
        )
        return

    _install_flydsl_stage2_multi_mode_compiled_cache(_flydsl_moe_kernels)

    cached_patched_fn = _KMOE_FLYDSL_STAGE2_PATCHED_FN_BY_MODE.get(patch_mode)
    if cached_patched_fn is not None:
        _mixed_stage2.compile_mixed_moe_gemm2 = cached_patched_fn
        _KMOE_ACTIVE_FLYDSL_STAGE2_PATCH_MODE = patch_mode
        _log_once(
            ("flydsl-stage2-sortedidx-lds-patch-reuse", patch_mode),
            "[submission] kmoe patch flydsl stage2 sortedidx lds reuse "
            f"patch_mode={patch_mode}",
        )
        return

    target_fn = getattr(_mixed_stage2, "compile_mixed_moe_gemm2", None)
    if target_fn is None:
        _log_once(
            ("flydsl-stage2-sortedidx-lds-patch-missing-target",),
            "[submission] kmoe patch flydsl stage2 sortedidx lds skipped "
            "reason=missing_compile_mixed_moe_gemm2",
        )
        return

    if patch_mode == "rawbuf":
        _mixed_stage2.compile_mixed_moe_gemm2 = target_fn
        _KMOE_FLYDSL_STAGE2_PATCHED_FN_BY_MODE[patch_mode] = target_fn
        _KMOE_ACTIVE_FLYDSL_STAGE2_PATCH_MODE = patch_mode
        _log_once(
            ("flydsl-stage2-sortedidx-lds-patch-rawbuf-orig",),
            "[submission] kmoe patch flydsl stage2 sortedidx lds bypass "
            "mode=orig_builder_rawbuf patch_mode=rawbuf",
        )
        return

    if _KMOE_FLYDSL_STAGE2_ORIG_SOURCE is None:
        target_unwrapped = _inspect.unwrap(target_fn)
        try:
            source = _inspect.getsource(target_unwrapped)
            source_file = _inspect.getsourcefile(target_unwrapped) or "unknown"
        except (OSError, TypeError) as exc:
            _log_once(
                ("flydsl-stage2-sortedidx-lds-patch-source-failure", type(exc).__name__),
                "[submission] kmoe patch flydsl stage2 sortedidx lds skipped "
                f"reason=source_failure exc={type(exc).__name__}",
            )
            return
        _KMOE_FLYDSL_STAGE2_ORIG_SOURCE = source
        _KMOE_FLYDSL_STAGE2_ORIG_SOURCE_FILE = source_file
    else:
        source = _KMOE_FLYDSL_STAGE2_ORIG_SOURCE
        source_file = _KMOE_FLYDSL_STAGE2_ORIG_SOURCE_FILE or "unknown"

    patch_tag = {
        "ptr": "_vscale_fix3_sortedidxlds18ptr",
        "rawbuf": "_vscale_fix4_sortedidxlds18rawbufhotloop",
    }[patch_mode]
    mode_label = {
        "ptr": "old_builder_sortedidx_lds_v18_ptr",
        "rawbuf": "old_builder_sortedidx_lds_v18_rawbuf_hotloop",
    }[patch_mode]
    rowctx_label = {
        "ptr": "idx0_i32_acc_ptr_runtime",
        "rawbuf": "idx0_i32_acc_rawbuf_miss_only",
    }[patch_mode]
    if patch_mode == "rawbuf":
        store_pair_body = (
            "frag_v = frag._value if hasattr(frag, \"_value\") else frag\n"
            "if bool(accumulate):\n"
            "    idx0 = row_ctx\n"
            "    col_i32 = arith.index_cast(i32, col_g0)\n"
            "    idx_elem = idx0 + col_i32\n"
            "    idx_elem_even = idx_elem & arith.constant(0xFFFFFFFE, type=i32)\n"
            "    byte_off = idx_elem_even * arith.constant(2, type=i32)\n"
            "    atomic_add_f16x2(frag, byte_off)\n"
            "    return\n"
            "row_byte_base = row_ctx\n"
            "col_idx = col_g0\n"
            "byte_off_col = col_idx * arith.constant(out_elem_bytes, index=True)\n"
            "ptr_addr_idx = row_byte_base + byte_off_col\n"
            "out_ptr = buffer_ops.create_llvm_ptr(ptr_addr_idx, address_space=1)\n"
            "out_ptr_v = out_ptr._value if hasattr(out_ptr, \"_value\") else out_ptr\n"
            "llvm.StoreOp(frag_v, out_ptr_v, alignment=4)\n"
        )
    else:
        store_pair_body = (
            "frag_v = frag._value if hasattr(frag, \"_value\") else frag\n"
            "if bool(accumulate):\n"
            "    idx0 = row_ctx\n"
            "    col_i32 = arith.index_cast(i32, col_g0)\n"
            "    idx_elem = idx0 + col_i32\n"
            "    idx_elem_even = idx_elem & arith.constant(0xFFFFFFFE, type=i32)\n"
            "    byte_off = idx_elem_even * arith.constant(2, type=i32)\n"
            "    byte_off_idx = arith.index_cast(ir.IndexType.get(), byte_off)\n"
            "    ptr_addr_idx = out_base_idx + byte_off_idx\n"
            "    out_ptr = buffer_ops.create_llvm_ptr(ptr_addr_idx, address_space=1)\n"
            "    out_ptr_v = out_ptr._value if hasattr(out_ptr, \"_value\") else out_ptr\n"
            "    llvm.AtomicRMWOp(\n"
            "        llvm.AtomicBinOp.fadd,\n"
            "        out_ptr_v,\n"
            "        frag_v,\n"
            "        llvm.AtomicOrdering.monotonic,\n"
            "        syncscope=\"agent\",\n"
            "        alignment=4,\n"
            "    )\n"
            "    return\n"
            "row_byte_base = row_ctx\n"
            "col_idx = col_g0\n"
            "byte_off_col = col_idx * arith.constant(out_elem_bytes, index=True)\n"
            "ptr_addr_idx = row_byte_base + byte_off_col\n"
            "out_ptr = buffer_ops.create_llvm_ptr(ptr_addr_idx, address_space=1)\n"
            "out_ptr_v = out_ptr._value if hasattr(out_ptr, \"_value\") else out_ptr\n"
            "llvm.StoreOp(frag_v, out_ptr_v, alignment=4)\n"
        )

    normalized = "\n".join(line.rstrip() for line in source.splitlines()).strip()
    digest = _hashlib.sha1(normalized.encode("utf-8")).hexdigest()[:16]

    def _replace_once(text: str, label: str, old: str, new: str) -> str:
        if old not in text:
            raise RuntimeError(label)
        return text.replace(old, new, 1)

    def _rewrite_rowctx_contract(text: str) -> tuple[str, bool, bool]:
        module_ast = _ast.parse(text)

        class _RowCtxFixer(_ast.NodeTransformer):
            def __init__(self):
                self.precompute_fixed = False
                self.store_pair_fixed = False

            def visit_FunctionDef(self, node):
                node = self.generic_visit(node)
                if not isinstance(node, _ast.FunctionDef):
                    return node
                if node.name == "precompute_row" and not self.precompute_fixed:
                    node.body = _ast.parse(
                        "fused2 = memref.load(lds_tid, [row_local])\n"
                        "row_i32 = arith.index_cast(i32, row)\n"
                        "row_valid0 = arith.cmpu(row_i32, num_valid_i32, \"ult\")\n"
                        "t = fused2 & mask24_i32\n"
                        "s = fused2 >> 24\n"
                        "t_ok = arith.cmpu(t, tokens_i32, \"ult\")\n"
                        "s_ok = arith.cmpu(s, topk_i32_v, \"ult\")\n"
                        "row_valid = arith.andi(row_valid0, arith.andi(t_ok, s_ok))\n"
                        "if bool(accumulate):\n"
                        "    t_safe = arith.select(row_valid, t, arith.constant(0, type=i32))\n"
                        "    idx0 = t_safe * arith.constant(model_dim, type=i32)\n"
                        "    return (idx0, row_valid)\n"
                        "t_safe = arith.select(row_valid, t, arith.constant(0, type=i32))\n"
                        "s_safe = arith.select(row_valid, s, arith.constant(0, type=i32))\n"
                        "t_idx = arith.index_cast(ir.IndexType.get(), t_safe)\n"
                        "s_idx = arith.index_cast(ir.IndexType.get(), s_safe)\n"
                        "n_byte_stride = arith.constant(model_dim * out_elem_bytes, index=True)\n"
                        "row_byte_base = out_base_idx + (t_idx * arith.constant(topk, index=True) + s_idx) * n_byte_stride\n"
                        "return (row_byte_base, row_valid)\n"
                    ).body
                    self.precompute_fixed = True
                    return node
                if node.name == "store_pair" and not self.store_pair_fixed:
                    node.body = _ast.parse(store_pair_body).body
                    self.store_pair_fixed = True
                    return node
                return node

        fixer = _RowCtxFixer()
        module_ast = fixer.visit(module_ast)
        _ast.fix_missing_locations(module_ast)
        return _ast.unparse(module_ast), fixer.precompute_fixed, fixer.store_pair_fixed

    def _rewrite_k_main2_loop(text: str) -> tuple[str, bool]:
        module_ast = _ast.parse(text)

        class _LoopFixer(_ast.NodeTransformer):
            def __init__(self):
                self.fixed = False

            def visit_For(self, node):
                node = self.generic_visit(node)
                if self.fixed:
                    return node
                if not isinstance(node, _ast.For):
                    return node
                if not isinstance(node.target, _ast.Name) or node.target.id != "k_iv":
                    return node
                loop_call = node.iter
                if not isinstance(loop_call, _ast.Call):
                    return node
                if not isinstance(loop_call.func, _ast.Name) or loop_call.func.id != "range":
                    return node
                if len(loop_call.args) != 3:
                    return node
                stop_arg = loop_call.args[1]
                if not isinstance(stop_arg, _ast.Name) or stop_arg.id != "c_k_main2":
                    return node

                fixed_if = _ast.parse(
                    "if k_main2_py > 0:\n"
                    "    for k_iv_py in range_constexpr(0, k_main2_py, tile_k * 2):\n"
                    "        k_iv = k_iv_py\n"
                    "        pass\n"
                ).body[0]
                fixed_for = fixed_if.body[0]
                fixed_for.body = (
                    [_ast.parse("k_iv = k_iv_py").body[0]] + list(node.body)
                )
                fixed_for.orelse = list(node.orelse)
                self.fixed = True
                return _ast.copy_location(fixed_if, node)

        fixer = _LoopFixer()
        module_ast = fixer.visit(module_ast)
        _ast.fix_missing_locations(module_ast)
        return _ast.unparse(module_ast), fixer.fixed

    try:
        patched_source = source
        patched_source = _replace_once(
            patched_source,
            "module_name_tag",
            '        f"_vscale_fix3"',
            f'        f"{patch_tag}"',
        )
        patched_source = _replace_once(
            patched_source,
            "lds_tid_bytes",
            "            lds_total_bytes = max(lds_x_bytes, lds_out_bytes)\n",
            "            lds_tid_bytes = int(tile_m) * 4\n"
            "            lds_total_bytes = max(lds_x_bytes, lds_out_bytes) + lds_tid_bytes\n",
        )
        patched_source = _replace_once(
            patched_source,
            "lds_tid_alias",
            "            # Buffer resources.\n",
            "            # lds_tid: alias LDS after max(x, out) for sorted_idx preload\n"
            "            _lds_x_b = 2 * int(tile_m) * int(lds_stride) * int(a_elem_bytes)\n"
            "            _lds_out_b = 2 * int(tile_m) * int(tile_n) if _use_cshuffle_epilog else 0\n"
            "            _lds_tid_off = max(_lds_x_b, _lds_out_b)\n"
            "            lds_tid = SmemPtr(\n"
            "                base_ptr, lds_x_ptr.byte_offset + _lds_tid_off, i32, shape=(tile_m,)\n"
            "            ).get()\n\n"
            "            # Buffer resources.\n",
        )
        patched_source = _replace_once(
            patched_source,
            "sortedidx_prologue",
            "                store_x_tile_to_lds(x_regs0, lds_base_cur)\n"
            "                gpu.barrier()\n",
            "                store_x_tile_to_lds(x_regs0, lds_base_cur)\n"
            "                # Preload sorted_idx into lds_tid for epilogue precompute_row.\n"
            "                if int(tile_m) % 4 == 0:\n"
            "                    _c_tile_m4_i32 = arith.constant(tile_m // 4, type=i32)\n"
            "                    _tid4_tx_i32 = arith.index_cast(i32, tx)\n"
            "                    _tid4_in_range = arith.cmpu(_tid4_tx_i32, _c_tile_m4_i32, \"ult\")\n"
            "                    _if_tid4 = scf.IfOp(_tid4_in_range)\n"
            "                    with ir.InsertionPoint(_if_tid4.then_block):\n"
            "                        _tid4_mul_i32 = _tid4_tx_i32 * arith.constant(4, type=i32)\n"
            "                        _tid4_col = arith.index_cast(ir.IndexType.get(), _tid4_mul_i32)\n"
            "                        _tid4_row = bx_m + _tid4_col\n"
            "                        _tid4_val = buffer_ops.buffer_load(\n"
            "                            sorted_rsrc, _tid4_row, vec_width=4, dtype=i32\n"
            "                        )\n"
            "                        vector.store(_tid4_val, lds_tid, [_tid4_col])\n"
            "                        scf.YieldOp([])\n"
            "                else:\n"
            "                    _c_tile_m_i32 = arith.constant(tile_m, type=i32)\n"
            "                    _tid_tx_i32 = arith.index_cast(i32, tx)\n"
            "                    _tid_in_range = arith.cmpu(_tid_tx_i32, _c_tile_m_i32, \"ult\")\n"
            "                    _if_tid = scf.IfOp(_tid_in_range)\n"
            "                    with ir.InsertionPoint(_if_tid.then_block):\n"
            "                        _tid_row = bx_m + tx\n"
            "                        _tid_val = buffer_ops.buffer_load(\n"
            "                            sorted_rsrc, _tid_row, vec_width=1, dtype=i32\n"
            "                        )\n"
            "                        memref.store(_tid_val, lds_tid, [tx])\n"
            "                        scf.YieldOp([])\n"
            "                gpu.barrier()\n",
        )
        patched_source, rowctx_precompute_fixed, rowctx_store_pair_fixed = (
            _rewrite_rowctx_contract(patched_source)
        )
        if not rowctx_precompute_fixed:
            raise RuntimeError("rowctx_ast_precompute")
        if not rowctx_store_pair_fixed:
            raise RuntimeError("rowctx_ast_store_pair")
        patched_source, k_loop_fixed = _rewrite_k_main2_loop(patched_source)
        if not k_loop_fixed:
            raise RuntimeError("k_main2_loop")
    except RuntimeError as exc:
        _log_once(
            ("flydsl-stage2-sortedidx-lds-patch-missing-pattern", digest, str(exc)),
            "[submission] kmoe patch flydsl stage2 sortedidx lds skipped "
            f"reason=missing_pattern sha1={digest} label={exc}",
        )
        return

    if patched_source == source:
        _log_once(
            ("flydsl-stage2-sortedidx-lds-patch-noop", digest),
            "[submission] kmoe patch flydsl stage2 sortedidx lds skipped "
            f"reason=noop sha1={digest}",
        )
        return

    new_digest = _hashlib.sha1(
        "\n".join(line.rstrip() for line in patched_source.splitlines()).strip().encode(
            "utf-8"
        )
    ).hexdigest()[:16]

    module_globals = vars(_mixed_stage2)
    try:
        exec(compile(patched_source, source_file, "exec"), module_globals, module_globals)
    except Exception as exc:
        _log_once(
            ("flydsl-stage2-sortedidx-lds-patch-exec-failure", digest, type(exc).__name__),
            "[submission] kmoe patch flydsl stage2 sortedidx lds skipped "
            f"reason=exec_failure sha1={digest} exc={type(exc).__name__}: {exc}",
        )
        return

    patched_fn = module_globals.get("compile_mixed_moe_gemm2")
    if patched_fn is None:
        _log_once(
            ("flydsl-stage2-sortedidx-lds-patch-missing-rebound", digest),
            "[submission] kmoe patch flydsl stage2 sortedidx lds skipped "
            f"reason=missing_rebound sha1={digest}",
        )
        return

    _mixed_stage2.compile_mixed_moe_gemm2 = patched_fn
    _KMOE_FLYDSL_STAGE2_PATCHED_FN_BY_MODE[patch_mode] = patched_fn
    try:
        _flydsl_moe_kernels._get_compiled_stage2.cache_clear()
    except Exception:
        pass

    _FLYDSL_STAGE2_SORTEDIDX_LDS_PATCHED = True
    _KMOE_ACTIVE_FLYDSL_STAGE2_PATCH_MODE = patch_mode
    _log_once(
        ("flydsl-stage2-sortedidx-lds-patch-active", digest, new_digest, patch_mode),
        "[submission] kmoe patch flydsl stage2 sortedidx lds active "
        f"old_sha1={digest} new_sha1={new_digest} mode={mode_label} "
        f"k_loop_fix={int(k_loop_fixed)} preload=dwordx4 rowctx={rowctx_label} "
        f"patch_mode={patch_mode}",
    )


_ORIG_GET_BLOCK_SIZE_M = _aiter_fused_moe.get_block_size_M
_ORIG_GET_2STAGE_CFGS = _aiter_fused_moe.get_2stage_cfgs
_FUSED_MOE_GLOBALS = fused_moe.__globals__
_ORIG_FLYDSL_STAGE2_WRAPPER = _aiter_fused_moe._flydsl_stage2_wrapper


def _patched_flydsl_stage2_wrapper(
    inter_states,
    w1,
    w2,
    sorted_token_ids,
    sorted_expert_ids,
    num_valid_ids,
    out,
    topk,
    kernelName="",
    w2_scale=None,
    a2_scale=None,
    sorted_weights=None,
    **kwargs,
):
    patch_mode = _KMOE_PENDING_FLYDSL_STAGE2_PATCH_MODE
    _patch_flydsl_stage2_sortedidx_lds_prefetch(patch_mode)
    _log_once(
        ("kmoe-flydsl-stage2-wrapper-patch-mode", kernelName, patch_mode),
        "[submission] kmoe flydsl stage2 wrapper patch mode "
        f"kernel='{kernelName}' patch_mode={patch_mode}",
    )
    return _ORIG_FLYDSL_STAGE2_WRAPPER(
        inter_states,
        w1,
        w2,
        sorted_token_ids,
        sorted_expert_ids,
        num_valid_ids,
        out,
        topk,
        kernelName=kernelName,
        w2_scale=w2_scale,
        a2_scale=a2_scale,
        sorted_weights=sorted_weights,
        **kwargs,
    )


_aiter_fused_moe._flydsl_stage2_wrapper = _patched_flydsl_stage2_wrapper
_FUSED_MOE_GLOBALS["_flydsl_stage2_wrapper"] = _patched_flydsl_stage2_wrapper
_BLOCK_M_OVERRIDES = {
    (512, 9, 33, 512): 128,
    (512, 9, 33, 2048): 64,
}


@functools.lru_cache(maxsize=2048)
def _patched_get_block_size_M(token, topk, expert, inter_dim):
    return _BLOCK_M_OVERRIDES.get(
        (token, topk, expert, inter_dim),
        _ORIG_GET_BLOCK_SIZE_M(token, topk, expert, inter_dim),
    )


_aiter_fused_moe.get_block_size_M = _patched_get_block_size_M
_FUSED_MOE_GLOBALS["get_block_size_M"] = _patched_get_block_size_M


@functools.lru_cache(maxsize=2048)
def _patched_get_2stage_cfgs(
    token,
    model_dim,
    inter_dim,
    expert,
    topk,
    dtype,
    q_dtype_a,
    q_dtype_w,
    q_type,
    use_g1u1,
    activation,
    doweight_stage1,
    hidden_pad,
    intermediate_pad,
    is_shuffled=True,
):
    metadata = _ORIG_GET_2STAGE_CFGS(
        token,
        model_dim,
        inter_dim,
        expert,
        topk,
        dtype,
        q_dtype_a,
        q_dtype_w,
        q_type,
        use_g1u1,
        activation,
        doweight_stage1,
        hidden_pad,
        intermediate_pad,
        is_shuffled,
    )
    override = _FUSED_CFG_STAGE_OVERRIDES.get((token, model_dim, inter_dim, expert, topk))
    if override is None:
        return metadata
    if (
        activation != ActivationType.Silu
        or dtype != dtypes.bf16
        or q_dtype_a != dtypes.fp4x2
        or q_dtype_w != dtypes.fp4x2
        or q_type != QuantType.per_1x32
        or not use_g1u1
        or doweight_stage1
        or not is_shuffled
        or metadata.run_1stage
    ):
        return metadata

    shape_key = (token, model_dim, inter_dim, expert, topk)
    stage1_kernel, stage2_kernel, block_m = override
    flydsl_available = bool(
        hasattr(_aiter_fused_moe, "is_flydsl_available")
        and _aiter_fused_moe.is_flydsl_available()
    )
    resolved_stage2_kernel = stage2_kernel
    if stage2_kernel.startswith("flydsl_") and not flydsl_available:
        resolved_stage2_kernel = _FLYDSL_STAGE2_FALLBACK_KERNELS.get(
            shape_key,
            stage2_kernel,
        )
        print(
            "[submission] fused cfg flydsl stage2 unavailable, "
            f"fallback to CK for shape={shape_key} kernel2='{resolved_stage2_kernel}'",
            file=sys.stderr,
        )
    print(
        f"[submission] fused cfg override active token={token} model_dim={model_dim} "
        f"inter_dim={inter_dim} expert={expert} topk={topk} block_m={block_m} "
        f"kernel1='{stage1_kernel}' kernel2='{resolved_stage2_kernel}'",
        file=sys.stderr,
    )
    stage1 = metadata.stage1
    if isinstance(stage1, functools.partial):
        stage1_kwargs = dict(stage1.keywords or {})
        stage1_kwargs["kernelName"] = stage1_kernel
        stage1 = functools.partial(stage1.func, *stage1.args, **stage1_kwargs)

    stage2 = metadata.stage2
    if resolved_stage2_kernel.startswith("flydsl_") and flydsl_available:
        stage2 = functools.partial(
            _aiter_fused_moe._flydsl_stage2_wrapper,
            kernelName=resolved_stage2_kernel,
        )
        print(
            "[submission] fused cfg flydsl stage2 active "
            f"shape={shape_key} kernel2='{resolved_stage2_kernel}'",
            file=sys.stderr,
        )
    elif isinstance(stage2, functools.partial):
        stage2_kwargs = dict(stage2.keywords or {})
        stage2_kwargs["kernelName"] = resolved_stage2_kernel
        stage2 = functools.partial(stage2.func, *stage2.args, **stage2_kwargs)

    return _aiter_fused_moe.MOEMetadata(
        stage1=stage1,
        stage2=stage2,
        block_m=block_m,
        ksplit=metadata.ksplit,
        run_1stage=metadata.run_1stage,
        has_bias=metadata.has_bias,
        use_non_temporal_load=metadata.use_non_temporal_load,
    )


_aiter_fused_moe.get_2stage_cfgs = _patched_get_2stage_cfgs
_FUSED_MOE_GLOBALS["get_2stage_cfgs"] = _patched_get_2stage_cfgs
scrolls · 3424 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