Skip to content
KernelIndex
Search⌘K

submission 754408

Jingze · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3df005cb95fe421d50ec694cc60ab6412fb63ce6ef4c3e307494c712c5185b63
license declaredunknown
license concludedunknown
authorsJingze
imported2026-08-26

Techniques

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

fp4"""Patch or fill CK kernel names when block_m changes on FP4 shuffled path.
split-ksplitk=0,

Kernel source

submission_ck.py1451 lines
from __future__ import annotations

import functools
import os
from typing import Any, Dict, Optional, Tuple

import aiter
import torch
from aiter import ActivationType, QuantType, dtypes, logger
from aiter import get_hip_quant as get_quant
from aiter.fused_moe import (
    fp4_utils,
    fused_dynamic_mxfp4_quant_moe_sort,
)
from aiter.ops.triton.moe.moe_op_mxfp4 import fused_moe_mxfp4
from aiter.ops.triton.moe.moe_op_mxfp4_silu_fused import fused_moe_mxfp4_silu
from aiter.ops.triton.utils.moe_config_utils import get_optimal_moe_config_func
from aiter.ops.triton.utils.types import torch_to_triton_dtype

from task import input_t, output_t


TOKEN_NUM_QUANT_MOE_SORT_SWITCH = 1024
# Keep this True to enable offline fixed-shape OPUS policy below.
_USE_OPUS_MOE_SORTING = True
_USE_TENSOR_CACHE = True

# Canonical FP4 shuffled CK kernel names used by this submission.
# These names must track block_m to avoid reusing a 32-tuned pair when block_m changes.
_FP4_CK_STAGE1_BY_BLOCK = {
    32: "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
    64: "moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
    128: "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
}

_FP4_CK_STAGE2_BY_BLOCK_INTER_SMALL = {
    32: "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
    64: "moe_ck2stages_gemm2_64x64x128x128_1x1_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
    128: "moe_ck2stages_gemm2_64x128x128x128_1x1_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
}

_FP4_CK_STAGE2_BY_BLOCK_INTER_LARGE = {
    32: "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
    64: "moe_ck2stages_gemm2_256x64x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
    128: "moe_ck2stages_gemm2_256x128x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
}

# Offline precomputed shape allowlist for OPUS moe sorting.
# Only these exact shapes will enable OPUS by default.
OPUS_SORTING_ENABLED_SHAPES = {
    # benchmark: TP=8 (E=257, d_expert=256)
    "M16_E257_H7168_I256_topk9",
    "M128_E257_H7168_I256_topk9",
    "M512_E257_H7168_I256_topk9",
    # benchmark: TP=4 (E=33, d_expert=512)
    "M16_E33_H7168_I512_topk9",
    "M128_E33_H7168_I512_topk9",
    "M512_E33_H7168_I512_topk9",
    # benchmark: EP-on (E=33, d_expert=2048)
    "M512_E33_H7168_I2048_topk9",
}

_BUFFER_CACHE: Dict[Tuple[Any, ...], torch.Tensor] = {}
_STATIC_TENSOR_CACHE: Dict[Tuple[Any, ...], Any] = {}
_LAST_BUFFER_CACHE_SHAPE_KEY: Optional[Tuple[Any, ...]] = None
_DEBUG_OUTPUT_COORDS = ((0, 0), (0, 2), (0, 3), (0, 4), (0, 5))
_DEBUG_PRINTED_CUSTOM_KERNEL_CALLS = 0


def _should_print_intermediates(config: Dict[str, Any]) -> bool:
    override = config.get("print_intermediates")
    if override is not None:
        return bool(override)
    env_value = os.getenv("SUBMISSION_CK_PRINT_INTERMEDIATES")
    if env_value is not None:
        return env_value not in ("0", "false", "False", "")
    return True


def _get_print_intermediates_max_calls(config: Dict[str, Any]) -> int:
    override = config.get("print_intermediates_max_calls")
    if override is not None:
        return int(override)
    env_value = os.getenv("SUBMISSION_CK_PRINT_INTERMEDIATES_MAX_CALLS")
    if env_value is not None:
        return int(env_value)
    return 1


def _should_print_this_custom_kernel_call(config: Dict[str, Any]) -> bool:
    global _DEBUG_PRINTED_CUSTOM_KERNEL_CALLS

    if not _should_print_intermediates(config):
        return False

    max_calls = _get_print_intermediates_max_calls(config)
    if max_calls == 0:
        return False
    if max_calls > 0 and _DEBUG_PRINTED_CUSTOM_KERNEL_CALLS >= max_calls:
        return False

    _DEBUG_PRINTED_CUSTOM_KERNEL_CALLS += 1
    return True


def _debug_view_tensor_for_summary(tensor: torch.Tensor) -> Tuple[torch.Tensor, str]:
    dtype_label = str(tensor.dtype)
    if tensor.dtype == dtypes.fp4x2:
        return tensor.view(torch.uint8), f"{dtype_label}(storage_as_uint8)"
    if tensor.dtype == dtypes.fp8_e8m0:
        return tensor.view(torch.uint8), f"{dtype_label}(storage_as_uint8)"
    return tensor, dtype_label


def _tensor_debug_summary(
    tensor: Optional[torch.Tensor],
    *,
    max_items: int = 8,
    coords_2d: Tuple[Tuple[int, int], ...] = _DEBUG_OUTPUT_COORDS,
) -> str:
    if tensor is None:
        return "None"

    view_tensor, dtype = _debug_view_tensor_for_summary(tensor.detach())
    shape = tuple(view_tensor.shape)
    if view_tensor.numel() == 0:
        return f"shape={shape} dtype={dtype} empty"

    if view_tensor.ndim == 0:
        if view_tensor.is_cuda:
            view_tensor = view_tensor.cpu()
        return f"shape={shape} dtype={dtype} value={view_tensor.item()}"

    if view_tensor.is_cuda:
        view_tensor = view_tensor.cpu()

    pieces = [f"shape={shape}", f"dtype={dtype}"]

    if torch.is_floating_point(view_tensor):
        flat_float = view_tensor.float().reshape(-1)
        pieces.append(f"min={flat_float.min().item():.6f}")
        pieces.append(f"max={flat_float.max().item():.6f}")
        pieces.append(f"mean={flat_float.mean().item():.6f}")

    if view_tensor.ndim == 1:
        pieces.append(f"head={view_tensor[:max_items].tolist()}")
    elif view_tensor.ndim == 2:
        coord_values = []
        for row, col in coords_2d:
            if row < view_tensor.shape[0] and col < view_tensor.shape[1]:
                coord_values.append(f"({row},{col})={view_tensor[row, col].item()}")
        if coord_values:
            pieces.append("coords=[" + ", ".join(coord_values) + "]")
        pieces.append(f"row0_head={view_tensor[0, :max_items].tolist()}")
    else:
        flattened = view_tensor.reshape(-1, view_tensor.shape[-1])
        pieces.append(f"flat0_head={flattened[0, :max_items].tolist()}")

    return " ".join(pieces)


def _debug_print_tensor(name: str, tensor: Optional[torch.Tensor]) -> None:
    print(f"[submission_ck] {name}: {_tensor_debug_summary(tensor)}", flush=True)


def _debug_print_message(message: str) -> None:
    print(f"[submission_ck] {message}", flush=True)


def _evict_cache_on_shape_change(shape_key: Tuple[Any, ...]) -> None:
    """Keep tensor cache only for the most recent runtime shape."""
    global _LAST_BUFFER_CACHE_SHAPE_KEY
    if _LAST_BUFFER_CACHE_SHAPE_KEY != shape_key:
        _BUFFER_CACHE.clear()
        _STATIC_TENSOR_CACHE.clear()
        _LAST_BUFFER_CACHE_SHAPE_KEY = shape_key


QUANT_SORT_SWITCH_TOKENS_DEFAULT = {
    "stage1": TOKEN_NUM_QUANT_MOE_SORT_SWITCH,
    "stage2": TOKEN_NUM_QUANT_MOE_SORT_SWITCH,
}

# Offline fixed thresholds for fused quant+sort crossover.
# Keys follow _shape_key(token_num, e, model_dim, inter_dim, topk).
QUANT_SORT_SWITCH_TOKENS_BY_SHAPE: Dict[str, Dict[str, int]] = {
    # benchmark: TP=8 (E=257, d_expert=256, topk=9)
    "M16_E257_H7168_I256_topk9": {"stage1": 1536, "stage2": 1280},
    "M128_E257_H7168_I256_topk9": {"stage1": 1536, "stage2": 1280},
    "M512_E257_H7168_I256_topk9": {"stage1": 1536, "stage2": 1280},
    # benchmark: TP=4 (E=33, d_expert=512, topk=9)
    "M16_E33_H7168_I512_topk9": {"stage1": 1344, "stage2": 1088},
    "M128_E33_H7168_I512_topk9": {"stage1": 1344, "stage2": 1088},
    "M512_E33_H7168_I512_topk9": {"stage1": 1344, "stage2": 1088},
    # benchmark: EP-on (E=33, d_expert=2048, topk=9)
    "M512_E33_H7168_I2048_topk9": {"stage1": 1216, "stage2": 960},
}

# Offline fixed policy for where to apply routed expert weights.
# True: apply weights in stage1; False: apply weights in stage2.
DOWEIGHT_STAGE1_BY_SHAPE: Dict[str, bool] = {
    # benchmark: TP=8 (E=257, d_expert=256, topk=9)
    "M16_E257_H7168_I256_topk9": False,
    "M128_E257_H7168_I256_topk9": False,
    "M512_E257_H7168_I256_topk9": False,
    # benchmark: TP=4 (E=33, d_expert=512, topk=9)
    "M16_E33_H7168_I512_topk9": False,
    "M128_E33_H7168_I512_topk9": False,
    "M512_E33_H7168_I512_topk9": False,
    # benchmark: EP-on (E=33, d_expert=2048, topk=9)
    "M512_E33_H7168_I2048_topk9": False,
}


def _quant_sort_switch_tokens(
    runtime_cfg: Dict[str, Any],
    *,
    stage: str,
    token_num: int,
    topk: int,
    model_dim: int,
    inter_dim: int,
    num_experts: int,
) -> int:
    """Return token threshold for choosing fused quant+sort path."""
    if stage not in ("stage1", "stage2"):
        raise ValueError(f"invalid stage: {stage}")

    # Highest priority: explicit per-stage override from runtime config.
    stage_override_key = f"quant_sort_switch_tokens_{stage}"
    stage_override = runtime_cfg.get(stage_override_key)
    if isinstance(stage_override, int) and stage_override > 0:
        return stage_override

    # Second priority: generic override for both stages.
    generic_override = runtime_cfg.get("quant_sort_switch_tokens")
    if isinstance(generic_override, int) and generic_override > 0:
        return generic_override

    # Offline fixed lookup only: no online threshold computation.
    shape_key = _shape_key(token_num, num_experts, model_dim, inter_dim, topk)
    shape_cfg = QUANT_SORT_SWITCH_TOKENS_BY_SHAPE.get(shape_key)
    if isinstance(shape_cfg, dict):
        stage_value = shape_cfg.get(stage)
        if isinstance(stage_value, int) and stage_value > 0:
            return stage_value

    return QUANT_SORT_SWITCH_TOKENS_DEFAULT[stage]


def _should_use_fused_quant_sort(
    runtime_cfg: Dict[str, Any],
    *,
    stage: str,
    token_num: int,
    topk: int,
    model_dim: int,
    inter_dim: int,
    num_experts: int,
) -> bool:
    switch_tokens = _quant_sort_switch_tokens(
        runtime_cfg,
        stage=stage,
        token_num=token_num,
        topk=topk,
        model_dim=model_dim,
        inter_dim=inter_dim,
        num_experts=num_experts,
    )
    return token_num <= switch_tokens


def _resolve_doweight_stage1(
    runtime_cfg: Dict[str, Any],
    ck_cfg: Dict[str, Any],
    *,
    token_num: int,
    num_experts: int,
    model_dim: int,
    inter_dim: int,
    topk: int,
) -> bool:
    """Resolve whether routed weights are multiplied in stage1 or stage2.

    Priority:
    1) runtime_cfg["doweight_stage1"] bool
    2) offline shape table (default behavior)
    3) ck_cfg fallback value
    """
    runtime_override = runtime_cfg.get("doweight_stage1")
    if isinstance(runtime_override, bool):
        return runtime_override

    shape_key = _shape_key(token_num, num_experts, model_dim, inter_dim, topk)
    shape_value = DOWEIGHT_STAGE1_BY_SHAPE.get(shape_key)
    if isinstance(shape_value, bool):
        return shape_value

    return bool(ck_cfg.get("doweight_stage1", False))


def _should_use_opus_moe_sorting(
    runtime_cfg: Dict[str, Any],
    token_num: int,
    num_experts: int,
    model_dim: int,
    inter_dim: int,
    topk: int,
    topk_ids: torch.Tensor,
    topk_weights: torch.Tensor,
) -> bool:
    cfg_override = runtime_cfg.get("use_opus_moe_sorting")
    if cfg_override is not None:
        return bool(cfg_override)

    if not _USE_OPUS_MOE_SORTING:
        return False
    if not hasattr(aiter, "moe_sorting_opus_fwd"):
        return False

    # Current OPUS binding accepts int32 topk ids and fp32 topk weights.
    if topk_ids.dtype != dtypes.i32:
        return False
    if topk_weights.dtype != dtypes.fp32:
        return False

    # OPUS MP path does not support the mesh_byte_size==2 branch (topk >= 255 when tokens >= 512).
    if topk >= 255:
        return False

    shape_key = _shape_key(token_num, num_experts, model_dim, inter_dim, topk)
    cfg_shape_allowlist = runtime_cfg.get("opus_enabled_shapes")
    if isinstance(cfg_shape_allowlist, (list, tuple, set)):
        return shape_key in set(cfg_shape_allowlist)
    return shape_key in OPUS_SORTING_ENABLED_SHAPES


def _make_cache_key(
    name: str,
    device: torch.device,
    dtype: torch.dtype,
    shape: Tuple[int, ...],
) -> Tuple[Any, ...]:
    return (name, device.type, device.index, str(dtype), *shape)


def _get_cached_tensor(
    cache_key: Tuple[Any, ...],
    shape: Tuple[int, ...],
    dtype: torch.dtype,
    device: torch.device,
    *,
    enabled: bool,
) -> torch.Tensor:
    if not enabled:
        return torch.empty(shape, dtype=dtype, device=device)
    cached = _BUFFER_CACHE.get(cache_key)
    if cached is None:
        cached = torch.empty(shape, dtype=dtype, device=device)
        _BUFFER_CACHE[cache_key] = cached
    return cached


def _get_cached_value(cache_key: Tuple[Any, ...], factory):
    cached = _STATIC_TENSOR_CACHE.get(cache_key)
    if cached is None:
        cached = factory()
        _STATIC_TENSOR_CACHE[cache_key] = cached
    return cached


def _slice_expert_tensor(
    tensor: Optional[torch.Tensor],
    start_expert: int,
    expert_count: int,
) -> Optional[torch.Tensor]:
    if tensor is None:
        return None
    sliced = tensor[start_expert : start_expert + expert_count].contiguous()
    if getattr(tensor, "is_shuffled", False):
        sliced.is_shuffled = True
    return sliced


def _slice_raw_scale_by_expert(
    scale: Optional[torch.Tensor],
    start_expert: int,
    expert_count: int,
    total_experts: int,
) -> Optional[torch.Tensor]:
    if scale is None:
        return None
    if scale.shape[0] == total_experts:
        return scale[start_expert : start_expert + expert_count].contiguous()
    if scale.ndim == 2 and scale.shape[0] % total_experts == 0:
        rows_per_expert = scale.shape[0] // total_experts
        row_start = start_expert * rows_per_expert
        row_end = (start_expert + expert_count) * rows_per_expert
        return scale[row_start:row_end].contiguous()
    raise ValueError(
        f"cannot slice scale by expert: shape={tuple(scale.shape)}, total_experts={total_experts}"
    )


def _make_branch_shuffled_scale(
    raw_scale: Optional[torch.Tensor],
    start_expert: int,
    expert_count: int,
    total_experts: int,
) -> Optional[torch.Tensor]:
    branch_scale = _slice_raw_scale_by_expert(
        raw_scale,
        start_expert,
        expert_count,
        total_experts,
    )
    if branch_scale is None:
        return None
    if branch_scale.ndim > 2:
        branch_scale = branch_scale.reshape(-1, branch_scale.shape[-1]).contiguous()
    return fp4_utils.e8m0_shuffle(branch_scale)


def _get_cached_shared_branch_tensors(
    *,
    w1: torch.Tensor,
    w2: torch.Tensor,
    gate_up_weight_scale: Optional[torch.Tensor],
    down_weight_scale: Optional[torch.Tensor],
    n_routed_experts: int,
    n_shared_experts: int,
    total_experts: int,
) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor], Optional[torch.Tensor]]:
    cache_key = (
        "shared_branch_tensors",
        w1.data_ptr(),
        w2.data_ptr(),
        None if gate_up_weight_scale is None else gate_up_weight_scale.data_ptr(),
        None if down_weight_scale is None else down_weight_scale.data_ptr(),
        n_routed_experts,
        n_shared_experts,
        total_experts,
    )

    def _factory():
        shared_w1 = _slice_expert_tensor(w1, n_routed_experts, n_shared_experts)
        shared_w2 = _slice_expert_tensor(w2, n_routed_experts, n_shared_experts)
        shared_w1_scale = _slice_raw_scale_by_expert(
            gate_up_weight_scale, n_routed_experts, n_shared_experts, total_experts
        )
        shared_w2_scale = _slice_raw_scale_by_expert(
            down_weight_scale, n_routed_experts, n_shared_experts, total_experts
        )
        return shared_w1, shared_w2, shared_w1_scale, shared_w2_scale

    return _get_cached_value(cache_key, _factory)


def _get_cached_shared_sorting(
    *,
    token_num: int,
    n_shared_experts: int,
    block_m: int,
    device: torch.device,
    use_opus: bool,
    dispatch_policy: int,
) -> Tuple[
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
    torch.Tensor,
]:
    cache_key = (
        "shared_branch_sorting",
        device.type,
        device.index,
        token_num,
        n_shared_experts,
        block_m,
        use_opus,
        dispatch_policy,
    )

    def _factory():
        shared_ids = torch.arange(n_shared_experts, dtype=dtypes.i32, device=device).view(
            1, n_shared_experts
        ).expand(token_num, -1).contiguous()
        shared_weights = torch.ones(
            (token_num, n_shared_experts), dtype=dtypes.fp32, device=device
        )
        sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, _ = moe_sorting(
            shared_ids,
            shared_weights,
            n_shared_experts,
            0,
            dtypes.fp32,
            block_m,
            expert_mask=None,
            num_local_tokens=None,
            dispatch_policy=dispatch_policy,
            use_tensor_cache=False,
            use_opus=use_opus,
        )
        num_tokens_post_padded = num_valid_ids[:1].contiguous()
        return (
            shared_ids,
            shared_weights,
            sorted_ids,
            sorted_weights,
            sorted_expert_ids,
            num_valid_ids,
            num_tokens_post_padded,
        )

    return _get_cached_value(cache_key, _factory)


def _as_triton_mx_scale(scale: torch.Tensor) -> torch.Tensor:
    return scale.view(torch.uint8)


def _as_triton_mx_data(data: torch.Tensor) -> torch.Tensor:
    return data.view(torch.uint8)


def _canonicalize_triton_weight_scale(
    scale: torch.Tensor,
    weight: torch.Tensor,
) -> torch.Tensor:
    if scale.ndim == 3:
        return _as_triton_mx_scale(scale)

    if scale.ndim != 2:
        raise ValueError(
            f"unsupported shared Triton scale rank: scale_shape={tuple(scale.shape)}, weight_shape={tuple(weight.shape)}"
        )

    expert_count, out_dim, packed_k = weight.shape
    scale_k = packed_k * 2 // 32
    expected_rows = expert_count * out_dim
    if scale.shape[0] != expected_rows:
        raise ValueError(
            f"cannot reshape shared Triton scale: scale_shape={tuple(scale.shape)}, weight_shape={tuple(weight.shape)}"
        )
    if scale.shape[1] != scale_k:
        raise ValueError(
            f"unexpected shared Triton scale K dimension: scale_shape={tuple(scale.shape)}, expected_scale_k={scale_k}"
        )

    return _as_triton_mx_scale(scale.reshape(expert_count, out_dim, scale_k))


def _run_shared_triton_dense(
    *,
    hidden_states: torch.Tensor,
    shared_w1: torch.Tensor,
    shared_w2: torch.Tensor,
    shared_w1_scale: torch.Tensor,
    shared_w2_scale: torch.Tensor,
    shared_ids: torch.Tensor,
    shared_weights: torch.Tensor,
    shared_sorted_ids: torch.Tensor,
    shared_sorted_expert_ids: torch.Tensor,
    shared_num_tokens_post_padded: torch.Tensor,
    token_num: int,
    n_shared_experts: int,
    inter_dim: int,
    model_dim: int,
    use_tensor_cache: bool,
) -> torch.Tensor:
    moe_config = get_optimal_moe_config_func(hidden_states.dtype, use_mxfp4=True)(token_num)
    compute_type = torch_to_triton_dtype[hidden_states.dtype]
    a_scale = torch.ones((1,), dtype=torch.float32, device=hidden_states.device)
    b1_scale = torch.ones((n_shared_experts,), dtype=torch.float32, device=hidden_states.device)
    b2_scale = b1_scale
    shared_w1_u8 = _as_triton_mx_data(shared_w1)
    shared_w2_u8 = _as_triton_mx_data(shared_w2)
    b1_mx_scale = _canonicalize_triton_weight_scale(shared_w1_scale, shared_w1)
    b2_mx_scale = _canonicalize_triton_weight_scale(shared_w2_scale, shared_w2)

    a1, a1_scale = fp4_utils.dynamic_mxfp4_quant(hidden_states)
    a1 = _as_triton_mx_data(a1)
    a1_scale = _as_triton_mx_scale(a1_scale)

    stage1_out = _get_cached_tensor(
        _make_cache_key(
            "shared_triton_stage1_out",
            hidden_states.device,
            hidden_states.dtype,
            (token_num * n_shared_experts, inter_dim),
        ),
        (token_num * n_shared_experts, inter_dim),
        hidden_states.dtype,
        hidden_states.device,
        enabled=use_tensor_cache,
    )
    fused_moe_mxfp4_silu(
        a1,
        shared_w1_u8,
        stage1_out,
        a_scale,
        b1_scale,
        a1_scale,
        b1_mx_scale,
        shared_weights,
        shared_ids,
        shared_sorted_ids,
        shared_sorted_expert_ids,
        shared_num_tokens_post_padded,
        False,
        n_shared_experts,
        False,
        False,
        moe_config,
        compute_type,
    )

    a2, a2_scale = fp4_utils.dynamic_mxfp4_quant(stage1_out)
    a2 = _as_triton_mx_data(a2)
    a2_scale = _as_triton_mx_scale(a2_scale)

    stage2_out = _get_cached_tensor(
        _make_cache_key(
            "shared_triton_stage2_out",
            hidden_states.device,
            hidden_states.dtype,
            (token_num, n_shared_experts, model_dim),
        ),
        (token_num, n_shared_experts, model_dim),
        hidden_states.dtype,
        hidden_states.device,
        enabled=use_tensor_cache,
    )
    fused_moe_mxfp4(
        a2,
        shared_w2_u8,
        stage2_out,
        a_scale,
        b2_scale,
        a2_scale,
        b2_mx_scale,
        shared_weights,
        shared_ids,
        shared_sorted_ids,
        shared_sorted_expert_ids,
        shared_num_tokens_post_padded,
        False,
        n_shared_experts,
        False,
        False,
        moe_config,
        compute_type,
    )

    if n_shared_experts == 1:
        return stage2_out.view(token_num, model_dim)
    return stage2_out.sum(dim=1)


# benchmarks:
#   # TP=8
#   - {"dhidden": 7168, "dexpert": 256, "nroutedexperts": 256, "nexpertspertoken": 8, "nsharedexperts": 1, "bs": 16, "seed": 9371}
#   - {"dhidden": 7168, "dexpert": 256, "nroutedexperts": 256, "nexpertspertoken": 8, "nsharedexperts": 1, "bs": 128, "seed": 2291}
#   - {"dhidden": 7168, "dexpert": 256, "nroutedexperts": 256, "nexpertspertoken": 8, "nsharedexperts": 1, "bs": 512, "seed": 81934}
#   # TP=4
#   - {"dhidden": 7168, "dexpert": 512, "nroutedexperts": 32, "nexpertspertoken": 8, "nsharedexperts": 1, "bs": 16, "seed": 2291}
#   - {"dhidden": 7168, "dexpert": 512, "nroutedexperts": 32, "nexpertspertoken": 8, "nsharedexperts": 1, "bs": 128, "seed": 81934}
#   - {"dhidden": 7168, "dexpert": 512, "nroutedexperts": 32, "nexpertspertoken": 8, "nsharedexperts": 1, "bs": 512, "seed": 81934}
#   # EP on
#   - {"dhidden": 7168, "dexpert": 2048, "nroutedexperts": 32, "nexpertspertoken": 8, "nsharedexperts": 1, "bs": 512, "seed": 81934}


# Manual kernel config table (CK).
#
# Runtime override priority:
# 1) config["kernel_manual_cfg"] or config["ck_manual_cfg"]
# 2) config["kernel_profile"] or config["ck_profile"]
# 3) shape profile in CK_MANUAL_CONFIGS
# 4) CK_MANUAL_CONFIGS["default"]
#
# Supported keys:
# - block_m: int
# - kernel_name1: str
# - kernel_name2: str
# - use_non_temporal_load: bool
# - doweight_stage1: bool
# - use_shuffled: bool
CK_MANUAL_CONFIGS: Dict[str, Dict[str, Any]] = {
    "default": {
        "block_m": 32,
        "use_non_temporal_load": False,
        "doweight_stage1": False,
        "use_shuffled": True,
    },
    # TP=8, dexpert=256, E=257, topk=9
    "M16_E257_H7168_I256_topk9": {
        "block_m": 32,
    },
    "M128_E257_H7168_I256_topk9": {
        "block_m": 32,
    },
    "M512_E257_H7168_I256_topk9": {
        "block_m": 32,
    },
    # TP=4, dexpert=512, E=33, topk=9
    # No exact tuned rows for E=33 in current CSVs; keep CK heuristic dispatch.
    "M16_E33_H7168_I512_topk9": {
        "block_m": 32,
    },
    "M128_E33_H7168_I512_topk9": {
        "block_m": 32,
    },
    "M512_E33_H7168_I512_topk9": {
        "block_m": 128,
    },
    # EP-on, dexpert=2048, E=33, topk=9
    # No exact row found in current CSV; keep CK heuristic fallback.
    "M512_E33_H7168_I2048_topk9": {
        "block_m": 128,
    },
}


def _moe_sorting_impl(
    topk_ids,
    topk_weights,
    num_experts,
    model_dim,
    moebuf_dtype,
    block_size,
    expert_mask,
    num_local_tokens,
    dispatch_policy,
    use_opus,
    use_tensor_cache,
):
    device = topk_ids.device
    M, topk = topk_ids.shape
    max_num_tokens_padded = topk_ids.numel() + num_experts * block_size - topk

    max_num_m_blocks = (max_num_tokens_padded + block_size - 1) // block_size
    sorted_ids = _get_cached_tensor(
        _make_cache_key(
            "moe_sorting_sorted_ids",
            device,
            dtypes.i32,
            (max_num_tokens_padded,),
        ),
        (max_num_tokens_padded,),
        dtypes.i32,
        device,
        enabled=use_tensor_cache,
    )
    sorted_weights = _get_cached_tensor(
        _make_cache_key(
            "moe_sorting_sorted_weights",
            device,
            dtypes.fp32,
            (max_num_tokens_padded,),
        ),
        (max_num_tokens_padded,),
        dtypes.fp32,
        device,
        enabled=use_tensor_cache,
    )
    sorted_expert_ids = _get_cached_tensor(
        _make_cache_key(
            "moe_sorting_sorted_expert_ids",
            device,
            dtypes.i32,
            (max_num_m_blocks,),
        ),
        (max_num_m_blocks,),
        dtypes.i32,
        device,
        enabled=use_tensor_cache,
    )
    num_valid_ids = _get_cached_tensor(
        _make_cache_key("moe_sorting_num_valid_ids", device, dtypes.i32, (2,)),
        (2,),
        dtypes.i32,
        device,
        enabled=use_tensor_cache,
    )
    moe_buf = _get_cached_tensor(
        _make_cache_key(
            "moe_final_out",
            device,
            moebuf_dtype,
            (M, model_dim),
        ),
        (M, model_dim),
        moebuf_dtype,
        device,
        enabled=use_tensor_cache,
    )

    fwd_fn = aiter.moe_sorting_opus_fwd if use_opus else aiter.moe_sorting_fwd
    fwd_fn(
        topk_ids,
        topk_weights,
        sorted_ids,
        sorted_weights,
        sorted_expert_ids,
        num_valid_ids,
        moe_buf,
        num_experts,
        block_size,
        expert_mask,
        num_local_tokens,
        dispatch_policy,
    )
    return sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf


def moe_sorting(
    topk_ids,
    topk_weights,
    num_experts,
    model_dim,
    moebuf_dtype,
    block_size,
    expert_mask=None,
    num_local_tokens=None,
    dispatch_policy=0,
    use_tensor_cache=True,
    use_opus: Optional[bool] = None,
):
    return _moe_sorting_impl(
        topk_ids,
        topk_weights,
        num_experts,
        model_dim,
        moebuf_dtype,
        block_size,
        expert_mask,
        num_local_tokens,
        dispatch_policy,
        use_opus=(_USE_OPUS_MOE_SORTING if use_opus is None else bool(use_opus)),
        use_tensor_cache=use_tensor_cache,
    )


# Lru cache will using hash to create key, which makes error when w1,w2 shape is symint.
# We can use torch.compile(dynamic=False) to avoid
@functools.lru_cache(maxsize=2048)
def get_inter_dim(w1_shape, w2_shape):
    E, _, model_dim = w1_shape
    E, model_dim, inter_dim = w2_shape

    int4_war = model_dim // w1_shape[-1]
    inter_dim *= int4_war
    return E, model_dim, inter_dim


def _shape_key(token_num: int, e: int, model_dim: int, inter_dim: int, topk: int) -> str:
    return f"M{token_num}_E{e}_H{model_dim}_I{inter_dim}_topk{topk}"


def _resolve_ck_cfg(
    runtime_cfg: Dict[str, Any],
    token_num: int,
    e: int,
    model_dim: int,
    inter_dim: int,
    topk: int,
) -> Dict[str, Any]:
    cfg = CK_MANUAL_CONFIGS["default"].copy()

    profile = runtime_cfg.get("kernel_profile", runtime_cfg.get("ck_profile"))
    if isinstance(profile, str) and profile in CK_MANUAL_CONFIGS:
        cfg.update(CK_MANUAL_CONFIGS[profile])

    shape_profile = _shape_key(token_num, e, model_dim, inter_dim, topk)
    if shape_profile in CK_MANUAL_CONFIGS:
        cfg.update(CK_MANUAL_CONFIGS[shape_profile])

    manual_cfg = runtime_cfg.get("kernel_manual_cfg", runtime_cfg.get("ck_manual_cfg"))
    if isinstance(manual_cfg, dict):
        cfg.update(manual_cfg)
    if cfg["block_m"] <= 0:
        raise ValueError(f"invalid block_m: {cfg['block_m']}")
    return cfg


def _ensure_shuffled_flag(tensor: torch.Tensor) -> None:
    # CK dispatch relies on this runtime flag for preshuffle path selection.
    if not getattr(tensor, "is_shuffled", False):
        tensor.is_shuffled = True


def _select_fp4_kernel_names_for_block_m(
    block_m: int,
    inter_dim: int,
    kernel_name1: str,
    kernel_name2: str,
    *,
    use_shuffled: bool,
    w1_dtype: torch.dtype,
    w2_dtype: torch.dtype,
    allow_stage1_name_patch: bool,
) -> Tuple[str, str]:
    """Patch or fill CK kernel names when block_m changes on FP4 shuffled path.

    We only touch names for the preshuffled FP4 path because this submission's
    manual table currently hardcodes a 32-tuned pair for multiple shapes.
    """
    if not use_shuffled:
        return kernel_name1, kernel_name2
    if w1_dtype != dtypes.fp4x2 or w2_dtype != dtypes.fp4x2:
        return kernel_name1, kernel_name2
    if block_m not in _FP4_CK_STAGE1_BY_BLOCK:
        return kernel_name1, kernel_name2

    expected_stage1 = _FP4_CK_STAGE1_BY_BLOCK[block_m]
    stage2_table = (
        _FP4_CK_STAGE2_BY_BLOCK_INTER_SMALL
        if inter_dim <= 256
        else _FP4_CK_STAGE2_BY_BLOCK_INTER_LARGE
    )
    expected_stage2 = stage2_table[block_m]

    legacy_stage1_32 = _FP4_CK_STAGE1_BY_BLOCK[32]
    legacy_stage2_32 = _FP4_CK_STAGE2_BY_BLOCK_INTER_SMALL[32]

    # Fill empty stage2 name from the per-block defaults.
    if not kernel_name2:
        kernel_name2 = expected_stage2

    # If user only changes block_m and leaves the old 32 stage2 name, rewrite safely.
    if block_m != 32:
        if kernel_name2 == legacy_stage2_32:
            kernel_name2 = expected_stage2

    # Stage1 FP4 named kernels do not match blockscale/per_1x128 dispatch.
    if allow_stage1_name_patch:
        if not kernel_name1:
            kernel_name1 = expected_stage1
        if block_m != 32 and kernel_name1 == legacy_stage1_32:
            kernel_name1 = expected_stage1

    return kernel_name1, kernel_name2


def _run_stage1(
    a1: torch.Tensor,
    w1: torch.Tensor,
    w2: torch.Tensor,
    sorted_ids: torch.Tensor,
    sorted_expert_ids: torch.Tensor,
    num_valid_ids: torch.Tensor,
    out: Optional[torch.Tensor],
    topk: int,
    kernel_name: str,
    w1_scale: Optional[torch.Tensor],
    a1_scale: Optional[torch.Tensor],
    block_m: int,
    sorted_weights: Optional[torch.Tensor],
    use_non_temporal_load: bool,
    dst_type: Optional[torch.dtype],
):
    assert out is not None
    aiter.ck_moe_stage1_fwd(
        a1,
        w1,
        w2,
        sorted_ids,
        sorted_expert_ids,
        num_valid_ids,
        out,
        topk,
        kernelName=kernel_name,
        w1_scale=w1_scale,
        a1_scale=a1_scale,
        block_m=block_m,
        sorted_weights=sorted_weights,
        quant_type=QuantType.per_1x32,
        activation=ActivationType.Silu,
        splitk=0,
        use_non_temporal_load=use_non_temporal_load,
        dst_type=dst_type,
    )
    return out


def _run_stage2(
    a2: torch.Tensor,
    w1: torch.Tensor,
    w2: torch.Tensor,
    sorted_ids: torch.Tensor,
    sorted_expert_ids: torch.Tensor,
    num_valid_ids: torch.Tensor,
    out: torch.Tensor,
    topk: int,
    kernel_name: str,
    w2_scale: Optional[torch.Tensor],
    a2_scale: Optional[torch.Tensor],
    block_m: int,
    sorted_weights: Optional[torch.Tensor],
    use_non_temporal_load: bool,
) -> torch.Tensor:
    aiter.ck_moe_stage2_fwd(
        a2,
        w1,
        w2,
        sorted_ids,
        sorted_expert_ids,
        num_valid_ids,
        out,
        topk,
        kernelName=kernel_name,
        w2_scale=w2_scale,
        a2_scale=a2_scale,
        block_m=block_m,
        sorted_weights=sorted_weights,
        quant_type=QuantType.per_1x32,
        activation=ActivationType.Silu,
        use_non_temporal_load=use_non_temporal_load,
    )
    return out


def custom_kernel(data: input_t) -> output_t:
    (
        hidden_states,
        gate_up_weight,
        down_weight,
        gate_up_weight_scale,
        down_weight_scale,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        gate_up_weight_scale_shuffled,
        down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        config,
    ) = data

    token_num, topk = topk_ids.shape
    e, model_dim, inter_dim = get_inter_dim(
        gate_up_weight_shuffled.shape, down_weight_shuffled.shape
    )
    ck_cfg = _resolve_ck_cfg(config, token_num, e, model_dim, inter_dim, topk)

    use_shuffled = ck_cfg.get("use_shuffled", config.get("use_shuffled", True))
    if use_shuffled:
        w1 = gate_up_weight_shuffled
        w2 = down_weight_shuffled
        w1_scale = gate_up_weight_scale_shuffled
        w2_scale = down_weight_scale_shuffled
        _ensure_shuffled_flag(w1)
        _ensure_shuffled_flag(w2)
    else:
        w1 = gate_up_weight
        w2 = down_weight
        w1_scale = gate_up_weight_scale
        w2_scale = down_weight_scale

    e, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)

    stage1_quant_type = QuantType.per_1x32
    quant_func_stage1 = get_quant(stage1_quant_type)
    quant_func_stage2 = get_quant(QuantType.per_1x32)

    # Keep cache entries only for one active shape.
    _evict_cache_on_shape_change(
        (
            hidden_states.device,
            hidden_states.dtype,
            token_num,
            topk,
            e,
            model_dim,
            inter_dim,
        )
    )

    block_m = ck_cfg["block_m"]
    use_tensor_cache = config.get("use_tensor_cache", _USE_TENSOR_CACHE)
    kernel_name1 = ck_cfg.get("kernel_name1", "")
    kernel_name2 = ck_cfg.get("kernel_name2", "")
    kernel_name1, kernel_name2 = _select_fp4_kernel_names_for_block_m(
        block_m,
        inter_dim,
        kernel_name1,
        kernel_name2,
        use_shuffled=use_shuffled,
        w1_dtype=w1.dtype,
        w2_dtype=w2.dtype,
        allow_stage1_name_patch=True,
    )
    use_non_temporal_load = ck_cfg.get("use_non_temporal_load", False)
    doweight_stage1 = _resolve_doweight_stage1(
        config,
        ck_cfg,
        token_num=token_num,
        num_experts=e,
        model_dim=model_dim,
        inter_dim=inter_dim,
        topk=topk,
    )
    moe_sorting_dispatch_policy = config.get("moe_sorting_dispatch_policy", 0)
    debug_print_intermediates = _should_print_this_custom_kernel_call(config)

    stage1_dst_type: Optional[torch.dtype] = None

    w1_scale_fp8 = w1_scale.view(dtypes.fp8_e8m0) if w1.dtype == dtypes.fp4x2 else w1_scale
    w2_scale_fp8 = w2_scale.view(dtypes.fp8_e8m0) if w2.dtype == dtypes.fp4x2 else w2_scale

    def _run_one_group(
        ids_group: torch.Tensor,
        weights_group: torch.Tensor,
        group_topk: int,
        *,
        group_num_experts: int,
        group_w1: torch.Tensor,
        group_w2: torch.Tensor,
        group_w1_scale_fp8: Optional[torch.Tensor],
        group_w2_scale_fp8: Optional[torch.Tensor],
        group_kernel_name1: str,
        group_name: str,
        weights_are_unit: bool = False,
        precomputed_sorting: Optional[
            Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]
        ] = None,
        output_buffer: Optional[torch.Tensor] = None,
        zero_output_buffer: bool = True,
    ) -> torch.Tensor:
        group_kernel_name2 = "" if weights_are_unit else kernel_name2
        if debug_print_intermediates:
            _debug_print_message(
                f"group={group_name} topk={group_topk} num_experts={group_num_experts} weights_are_unit={weights_are_unit}"
            )
            _debug_print_tensor(f"{group_name}.ids_group", ids_group)
            _debug_print_tensor(f"{group_name}.weights_group", weights_group)
        if precomputed_sorting is None:
            use_opus_sorting = _should_use_opus_moe_sorting(
                config,
                token_num,
                group_num_experts,
                model_dim,
                inter_dim,
                group_topk,
                ids_group,
                weights_group,
            )
            sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_out = moe_sorting(
                ids_group,
                weights_group,
                group_num_experts,
                model_dim,
                hidden_states.dtype,
                block_m,
                expert_mask=None,
                num_local_tokens=None,
                dispatch_policy=moe_sorting_dispatch_policy,
                use_tensor_cache=use_tensor_cache,
                use_opus=use_opus_sorting,
            )
            if output_buffer is not None:
                moe_out = output_buffer
                if zero_output_buffer:
                    moe_out.zero_()
        else:
            sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids = precomputed_sorting
            if output_buffer is None:
                moe_out = _get_cached_tensor(
                    _make_cache_key(
                        f"moe_final_out_{group_name}",
                        hidden_states.device,
                        hidden_states.dtype,
                        (token_num, model_dim),
                    ),
                    (token_num, model_dim),
                    hidden_states.dtype,
                    hidden_states.device,
                    enabled=use_tensor_cache,
                )
            else:
                moe_out = output_buffer
            if zero_output_buffer:
                moe_out.zero_()

        if debug_print_intermediates:
            _debug_print_tensor(f"{group_name}.sorted_ids", sorted_ids)
            _debug_print_tensor(f"{group_name}.sorted_weights", sorted_weights)
            _debug_print_tensor(f"{group_name}.sorted_expert_ids", sorted_expert_ids)
            _debug_print_tensor(f"{group_name}.num_valid_ids", num_valid_ids)

        use_fused_stage1 = _should_use_fused_quant_sort(
            config,
            stage="stage1",
            token_num=token_num,
            topk=group_topk,
            model_dim=model_dim,
            inter_dim=inter_dim,
            num_experts=group_num_experts,
        )
        if use_fused_stage1:
            a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
                hidden_states,
                sorted_ids=sorted_ids,
                num_valid_ids=num_valid_ids,
                token_num=token_num,
                topk=1,
                block_size=block_m,
            )
        else:
            a1, a1_scale = quant_func_stage1(
                hidden_states,
                scale=None,
                quant_dtype=dtypes.fp4x2,
                num_rows=None,
            )
            a1_scale = fp4_utils.moe_mxfp4_sort(
                a1_scale,
                sorted_ids=sorted_ids,
                num_valid_ids=num_valid_ids,
                token_num=token_num,
                block_size=block_m,
            )

        if debug_print_intermediates:
            _debug_print_message(f"group={group_name} use_fused_stage1={use_fused_stage1}")
            _debug_print_tensor(f"{group_name}.a1", a1)
            _debug_print_tensor(f"{group_name}.a1_scale", a1_scale)

        stage1_out_dtype = hidden_states.dtype
        stage1_out_buf = _get_cached_tensor(
            _make_cache_key(
                f"ck_stage1_out_{group_name}",
                hidden_states.device,
                stage1_out_dtype,
                (token_num, group_topk, inter_dim),
            ),
            (token_num, group_topk, inter_dim),
            stage1_out_dtype,
            hidden_states.device,
            enabled=use_tensor_cache,
        )

        stage1_result = _run_stage1(
            a1,
            group_w1,
            group_w2,
            sorted_ids,
            sorted_expert_ids,
            num_valid_ids,
            stage1_out_buf,
            group_topk,
            group_kernel_name1,
            w1_scale=group_w1_scale_fp8,
            a1_scale=a1_scale,
            block_m=block_m,
            sorted_weights=(sorted_weights if doweight_stage1 and not weights_are_unit else None),
            use_non_temporal_load=use_non_temporal_load,
            dst_type=stage1_dst_type,
        )

        if debug_print_intermediates:
            _debug_print_tensor(f"{group_name}.stage1_result", stage1_result)

        a2_scale: Optional[torch.Tensor] = None
        a2 = stage1_result.view(-1, inter_dim)
        use_fused_stage2 = _should_use_fused_quant_sort(
            config,
            stage="stage2",
            token_num=token_num,
            topk=group_topk,
            model_dim=model_dim,
            inter_dim=inter_dim,
            num_experts=group_num_experts,
        )
        if use_fused_stage2:
            a2, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
                a2,
                sorted_ids=sorted_ids,
                num_valid_ids=num_valid_ids,
                token_num=token_num,
                topk=group_topk,
                block_size=block_m,
            )
        else:
            a2, a2_scale = quant_func_stage2(
                a2,
                scale=a2_scale,
                quant_dtype=dtypes.fp4x2,
                num_rows=None,
                num_rows_factor=group_topk,
            )
            a2_scale = fp4_utils.moe_mxfp4_sort(
                a2_scale[: token_num * group_topk, :].view(token_num, group_topk, -1),
                sorted_ids=sorted_ids,
                num_valid_ids=num_valid_ids,
                token_num=token_num,
                block_size=block_m,
            )
        a2 = a2.view(token_num, group_topk, -1)

        if debug_print_intermediates:
            _debug_print_message(f"group={group_name} use_fused_stage2={use_fused_stage2}")
            _debug_print_tensor(f"{group_name}.a2", a2)
            _debug_print_tensor(f"{group_name}.a2_scale", a2_scale)

        _run_stage2(
            a2,
            group_w1,
            group_w2,
            sorted_ids,
            sorted_expert_ids,
            num_valid_ids,
            moe_out,
            group_topk,
            group_kernel_name2,
            w2_scale=group_w2_scale_fp8,
            a2_scale=a2_scale,
            block_m=block_m,
            sorted_weights=(None if doweight_stage1 or weights_are_unit else sorted_weights),
            use_non_temporal_load=use_non_temporal_load,
        )
        if debug_print_intermediates:
            _debug_print_tensor(f"{group_name}.moe_out", moe_out)
        return moe_out

    n_shared_experts = int(config.get("n_shared_experts", 0))
    if n_shared_experts < 0 or n_shared_experts > topk:
        raise ValueError(f"invalid n_shared_experts: {n_shared_experts}")

    routed_topk = topk - n_shared_experts
    n_routed_experts = int(config.get("n_routed_experts", e - n_shared_experts))
    if n_routed_experts < 0 or n_routed_experts + n_shared_experts > e:
        raise ValueError(
            f"invalid expert partition: n_routed_experts={n_routed_experts}, n_shared_experts={n_shared_experts}, total_experts={e}"
        )

    split_shared_experts = bool(config.get("split_shared_experts", False))

    if debug_print_intermediates:
        _debug_print_message(
            "custom_kernel "
            f"token_num={token_num} topk={topk} num_experts={e} model_dim={model_dim} inter_dim={inter_dim} "
            f"block_m={block_m} use_shuffled={use_shuffled} doweight_stage1={doweight_stage1} "
            f"split_shared_experts={split_shared_experts} use_tensor_cache={use_tensor_cache}"
        )
        _debug_print_tensor("hidden_states", hidden_states)
        _debug_print_tensor("topk_ids", topk_ids)
        _debug_print_tensor("topk_weights", topk_weights)
        _debug_print_message(
            f"kernel_name1={kernel_name1 or '<empty>'} kernel_name2={kernel_name2 or '<empty>'}"
        )

    if n_shared_experts == 0 or not split_shared_experts:
        moe_out = _run_one_group(
            topk_ids,
            topk_weights,
            topk,
            group_num_experts=e,
            group_w1=w1,
            group_w2=w2,
            group_w1_scale_fp8=w1_scale_fp8,
            group_w2_scale_fp8=w2_scale_fp8,
            group_kernel_name1=kernel_name1,
            group_name="all",
        )
    else:
        moe_out: Optional[torch.Tensor] = None
        if routed_topk > 0:
            routed_out = _run_one_group(
                topk_ids[:, :routed_topk].contiguous(),
                topk_weights[:, :routed_topk].contiguous(),
                routed_topk,
                group_num_experts=e,
                group_w1=w1,
                group_w2=w2,
                group_w1_scale_fp8=w1_scale_fp8,
                group_w2_scale_fp8=w2_scale_fp8,
                group_kernel_name1=kernel_name1,
                group_name="routed",
            )
            moe_out = routed_out

        shared_w1, shared_w2, shared_w1_scale, shared_w2_scale = _get_cached_shared_branch_tensors(
            w1=gate_up_weight,
            w2=down_weight,
            gate_up_weight_scale=gate_up_weight_scale,
            down_weight_scale=down_weight_scale,
            n_routed_experts=n_routed_experts,
            n_shared_experts=n_shared_experts,
            total_experts=e,
        )
        shared_ids, shared_weights, shared_sorted_ids, shared_sorted_weights, shared_sorted_expert_ids, shared_num_valid_ids, shared_num_tokens_post_padded = _get_cached_shared_sorting(
            token_num=token_num,
            n_shared_experts=n_shared_experts,
            block_m=get_optimal_moe_config_func(hidden_states.dtype, use_mxfp4=True)(token_num)["BLOCK_SIZE_M"],
            device=hidden_states.device,
            use_opus=False,
            dispatch_policy=moe_sorting_dispatch_policy,
        )
        use_shared_triton = bool(config.get("use_shared_triton", False))
        if use_shared_triton:
            shared_out = _run_shared_triton_dense(
                hidden_states=hidden_states,
                shared_w1=shared_w1,
                shared_w2=shared_w2,
                shared_w1_scale=shared_w1_scale,
                shared_w2_scale=shared_w2_scale,
                shared_ids=shared_ids,
                shared_weights=shared_weights,
                shared_sorted_ids=shared_sorted_ids,
                shared_sorted_expert_ids=shared_sorted_expert_ids,
                shared_num_tokens_post_padded=shared_num_tokens_post_padded,
                token_num=token_num,
                n_shared_experts=n_shared_experts,
                inter_dim=inter_dim,
                model_dim=model_dim,
                use_tensor_cache=use_tensor_cache,
            )
        else:
            shared_out = _run_one_group(
                shared_ids,
                shared_weights,
                n_shared_experts,
                group_num_experts=n_shared_experts,
                group_w1=shared_w1,
                group_w2=shared_w2,
                group_w1_scale_fp8=shared_w1_scale,
                group_w2_scale_fp8=shared_w2_scale,
                group_kernel_name1=kernel_name1,
                group_name="shared",
                weights_are_unit=True,
                precomputed_sorting=(
                    shared_sorted_ids,
                    shared_sorted_weights,
                    shared_sorted_expert_ids,
                    shared_num_valid_ids,
                ),
            )
        if moe_out is None:
            moe_out = shared_out
        else:
            moe_out.add_(shared_out)

    d_hidden = config.get("d_hidden", hidden_states.shape[1])
    final_out = moe_out[:, :d_hidden]
    if debug_print_intermediates:
        _debug_print_tensor("final_out", final_out)
    return final_out
scrolls · 1451 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