# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Kernel tests for ROCm AITER fused MoE.

This file owns the ROCm-specific fused-MoE custom-op path:
- custom-op registration and fake-tensor support
- enablement gating for ``VLLM_ROCM_USE_AITER_MOE`` and shared experts
- BF16 accuracy for the fused MoE kernel on representative shapes
- gfx950-only AITER MXFP4 W4A16 MoE support, accuracy, and determinism
- MoE-facing FP8 group-quant activation quality
- deterministic routing and representative gfx942 / gfx950 coverage

Generic fused-MoE backend selection and non-ROCm kernel coverage live in the
generic MoE test files under ``tests/kernels/moe``.

Raw ROCm AITER helper-op coverage such as ``group_fp8_quant`` and
``per_tensor_quant`` lives in ``tests/kernels/core/test_rocm_aiter_ops.py``.
This file only keeps the MoE-shaped integration angle for those helpers.
"""

import importlib
import math
import warnings
from typing import Any, NamedTuple
from unittest import mock

import pytest
import torch
import torch.nn.functional as F

from tests.kernels.utils import _assert_deterministic
from vllm.platforms import current_platform
from vllm.platforms.rocm import on_cdna, on_gfx942, on_gfx950
from vllm.utils.torch_utils import set_random_seed

pytestmark = pytest.mark.skipif(
    not current_platform.is_rocm(), reason="ROCm-specific tests"
)


# Helpers -----------------------------------------------------------------


def _reload_envs():
    import vllm.envs as envs

    return importlib.reload(envs)


@pytest.fixture(autouse=True)
def _restore_rocm_env_state():
    """Restore global env + AITER flag state after every test."""
    yield
    _reload_envs()
    from vllm._aiter_ops import rocm_aiter_ops

    rocm_aiter_ops.refresh_env_variables()


def _assert_aiter_supported() -> None:
    from vllm._aiter_ops import is_aiter_found_and_supported

    assert is_aiter_found_and_supported(), (
        "aiter is required on supported ROCm hardware for this test"
    )


def _format_observed_rate(count: int, total: int) -> str:
    return f"{count / total:.4%} ({count}/{total})"


def _format_allowed_rate(rate: float, total: int) -> str:
    allowed_count = int(rate * total)
    return f"{rate:.4%} (<= {allowed_count}/{total})"


def _quantile(values: torch.Tensor, q: float) -> float:
    if values.numel() == 0:
        return 0.0
    return torch.quantile(values, q).item()


def _assert_close_budget(
    actual: torch.Tensor,
    expected: torch.Tensor,
    *,
    label: str,
    atol: float,
    rtol: float = 0.0,
    pass_rate: float = 0.99999,
    max_violation_factor: float = 3.0,
) -> None:
    actual_f = actual.detach().float().flatten()
    expected_f = expected.detach().float().flatten()
    abs_diff = (actual_f - expected_f).abs()
    allowed = atol + rtol * expected_f.abs()

    total = abs_diff.numel()
    within = abs_diff <= allowed
    passed = int(within.sum().item())
    failed = total - passed
    allowed_fail_rate = 1.0 - pass_rate

    max_abs = abs_diff.max().item()
    mean_abs = abs_diff.mean().item()
    p99_abs = _quantile(abs_diff, 0.99)
    p999_abs = _quantile(abs_diff, 0.999)
    worst_ratio = (abs_diff / allowed.clamp_min(1e-12)).max().item()
    max_atol = max_violation_factor * atol
    above_max_count = int((abs_diff > max_atol).sum().item())

    msg = (
        "[rocm_aiter_moe] "
        f"{label}: "
        f"pass={passed / total:.4%} ({passed}/{total}) "
        f"fail={_format_observed_rate(failed, total)} "
        f"allowed_fail={_format_allowed_rate(allowed_fail_rate, total)} "
        f"atol={atol:g} "
        f"rtol={rtol:g} "
        f"abs>{max_atol:g}={_format_observed_rate(above_max_count, total)} "
        f"allowed_above_max={_format_allowed_rate(0.0, total)} "
        f"max_abs={max_abs:.6g} "
        f"mean_abs={mean_abs:.6g} "
        f"p99_abs={p99_abs:.6g} "
        f"p999_abs={p999_abs:.6g} "
        f"worst_ratio={worst_ratio:.6g}"
    )
    print(msg)
    if failed > 0:
        warnings.warn(msg, stacklevel=2)

    assert passed / total >= pass_rate, msg
    assert max_abs <= max_atol, msg
    assert mean_abs <= atol * 0.25, msg


def _assert_group_quant_quality(
    actual: torch.Tensor,
    expected: torch.Tensor,
    *,
    label: str,
    preferred_rel: float = 0.06,
    pass_rate: float = 0.9999,
    max_rel: float = 0.55,
    max_fail_rate: float = 0.000005,
    mean_limit: float = 0.03,
) -> None:
    rel = (
        (actual.float() - expected.float()).abs()
        / expected.float().abs().clamp_min(1e-5)
    ).flatten()
    total = rel.numel()
    within_preferred_count = int((rel <= preferred_rel).sum().item())
    fail_count = total - within_preferred_count
    allowed_fail_rate = 1.0 - pass_rate
    above_max_count = int((rel > max_rel).sum().item())
    mean_rel = rel.mean().item()
    max_rel_err = rel.max().item()
    p99 = _quantile(rel, 0.99)
    p999 = _quantile(rel, 0.999)

    msg = (
        "[rocm_aiter_moe] "
        f"{label}: "
        f"rel<={preferred_rel:g} pass={within_preferred_count / total:.4%} "
        f"({within_preferred_count}/{total}) "
        f"fail={_format_observed_rate(fail_count, total)} "
        f"allowed_fail={_format_allowed_rate(allowed_fail_rate, total)} "
        f"rel>{max_rel:g}={_format_observed_rate(above_max_count, total)} "
        f"allowed_above_max={_format_allowed_rate(max_fail_rate, total)} "
        f"mean_rel={mean_rel:.6g} "
        f"max_rel={max_rel_err:.6g} "
        f"p99={p99:.6g} "
        f"p999={p999:.6g}"
    )
    print(msg)
    if fail_count > 0:
        warnings.warn(msg, stacklevel=2)

    assert within_preferred_count / total >= pass_rate, msg
    assert above_max_count / total <= max_fail_rate, msg
    assert mean_rel < mean_limit, msg


def ref_moe_forward(
    hidden_states: torch.Tensor,
    w1: torch.Tensor,
    w2: torch.Tensor,
    topk_weights: torch.Tensor,
    topk_ids: torch.Tensor,
    activation: str = "silu",
) -> torch.Tensor:
    """Float32 mask-based MoE reference for the ROCm fused-MoE kernel."""
    num_tokens, hidden_dim = hidden_states.shape
    num_experts = w1.shape[0]
    intermediate_dim = w1.shape[1] // 2
    topk = topk_ids.shape[1]
    device = hidden_states.device

    hidden_states_f = hidden_states.float()
    w1_f = w1.float()
    w2_f = w2.float()

    expanded = hidden_states_f.view(num_tokens, 1, hidden_dim).expand(
        num_tokens, topk, hidden_dim
    )
    expanded = expanded.reshape(num_tokens * topk, hidden_dim)
    output = torch.zeros(
        num_tokens * topk, hidden_dim, dtype=torch.float32, device=device
    )
    flat_topk_ids = topk_ids.view(-1).long()

    for expert_idx in range(num_experts):
        expert_mask = flat_topk_ids == expert_idx
        if expert_mask.sum() == 0:
            continue
        gate_up = expanded[expert_mask] @ w1_f[expert_idx].T
        gate = gate_up[:, :intermediate_dim]
        up = gate_up[:, intermediate_dim:]
        if activation == "silu":
            act = F.silu(gate) * up
        elif activation == "gelu":
            act = F.gelu(gate) * up
        else:
            raise ValueError(f"Unknown activation: {activation}")
        output[expert_mask] = act @ w2_f[expert_idx].T

    output = output.view(num_tokens, topk, hidden_dim)
    weights = topk_weights.float().view(num_tokens, topk, 1)
    return (output * weights).sum(dim=1)


class AiterMxfp4MoeCase(NamedTuple):
    hidden_states: torch.Tensor
    w1_kernel: torch.Tensor
    w2_kernel: torch.Tensor
    w1_ref: torch.Tensor
    w2_ref: torch.Tensor
    topk_weights: torch.Tensor
    topk_ids: torch.Tensor
    moe_config: Any
    quant_config: Any


def _make_topk_ids(
    num_tokens: int,
    num_experts: int,
    topk: int,
    *,
    device: str = "cuda",
) -> torch.Tensor:
    """Generate distinct expert IDs per token with the same top-k shape as
    production routing."""
    router_logits = torch.randn(num_tokens, num_experts, device=device)
    _, topk_ids = torch.topk(torch.softmax(router_logits, dim=-1), k=topk, dim=-1)
    return topk_ids.to(torch.int32)


def _shuffle_moe_weights(
    w1: torch.Tensor,
    w2: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Shuffle MoE weights into the AITER CK layout used in production."""
    from vllm._aiter_ops import rocm_aiter_ops

    w1_shuffled, w2_shuffled = rocm_aiter_ops.shuffle_weights(w1, w2)
    w1_shuffled.is_shuffled = True
    w2_shuffled.is_shuffled = True
    return w1_shuffled, w2_shuffled


def _make_moe_case(
    *,
    num_tokens: int,
    hidden_dim: int,
    intermediate_dim: int,
    num_experts: int,
    topk: int,
    seed: int,
) -> dict[str, torch.Tensor]:
    torch.set_default_device("cuda")
    set_random_seed(seed)

    hidden_states = torch.randn(num_tokens, hidden_dim, dtype=torch.bfloat16)
    w1 = torch.randn(
        num_experts,
        intermediate_dim * 2,
        hidden_dim,
        dtype=torch.bfloat16,
    ) / math.sqrt(hidden_dim)
    w2 = torch.randn(
        num_experts,
        hidden_dim,
        intermediate_dim,
        dtype=torch.bfloat16,
    ) / math.sqrt(intermediate_dim)
    topk_weights = torch.rand(num_tokens, topk, dtype=torch.float32)
    topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
    topk_ids = _make_topk_ids(num_tokens, num_experts, topk)
    return {
        "hidden_states": hidden_states,
        "w1": w1,
        "w2": w2,
        "topk_weights": topk_weights,
        "topk_ids": topk_ids,
    }


def _interleave_gate_up_rows(t: torch.Tensor) -> torch.Tensor:
    """Reorder contiguous ``[gate; up]`` rows into gpt-oss interleaved order."""
    e, two_i = t.shape[0], t.shape[1]
    i, rest = two_i // 2, t.shape[2:]
    perm = (0, 2, 1, *range(3, 3 + len(rest)))
    return t.view(e, 2, i, *rest).permute(*perm).contiguous().view(e, two_i, *rest)


def _make_aiter_mxfp4_moe_case(
    *,
    num_tokens: int,
    hidden_dim: int,
    intermediate_dim: int,
    num_experts: int,
    topk: int,
    seed: int,
) -> AiterMxfp4MoeCase:
    from triton_kernels.numerics_details.mxfp import (
        downcast_to_mxfp,
        upcast_from_mxfp,
    )

    from tests.kernels.moe.utils import make_dummy_moe_config
    from vllm.model_executor.layers.fused_moe.oracle.mxfp4 import (
        Mxfp4MoeBackend,
        convert_gpt_oss_weight_to_mxfp4_moe_kernel_format,
        make_mxfp4_moe_quant_config,
    )

    torch.set_default_device("cuda")
    set_random_seed(seed)

    hidden_states = torch.randn(num_tokens, hidden_dim, dtype=torch.bfloat16)
    w1 = torch.randn(
        num_experts,
        intermediate_dim * 2,
        hidden_dim,
        dtype=torch.bfloat16,
    ) / math.sqrt(hidden_dim)
    w2 = torch.randn(
        num_experts,
        hidden_dim,
        intermediate_dim,
        dtype=torch.bfloat16,
    ) / math.sqrt(intermediate_dim)
    topk_weights = torch.rand(num_tokens, topk, dtype=torch.float32)
    topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
    topk_ids = _make_topk_ids(num_tokens, num_experts, topk)

    w1_q, w1_scale = downcast_to_mxfp(w1, torch.uint8, axis=-1)
    w2_q, w2_scale = downcast_to_mxfp(w2, torch.uint8, axis=-1)

    w1_ref = upcast_from_mxfp(w1_q, w1_scale, torch.bfloat16, axis=-1)
    w2_ref = upcast_from_mxfp(w2_q, w2_scale, torch.bfloat16, axis=-1)

    (
        w1_kernel,
        w2_kernel,
        w1_scale_kernel,
        w2_scale_kernel,
        _,
        _,
    ) = convert_gpt_oss_weight_to_mxfp4_moe_kernel_format(
        mxfp4_backend=Mxfp4MoeBackend.AITER,
        layer=torch.nn.Module(),
        # w13 rows must be gpt-oss-interleaved; the converter de-interleaves them
        # back to the contiguous gate/up blocks that ``w1_ref`` / ``ref_moe_forward``
        # use. w2 has no gate/up split, so it stays as-is.
        w13_weight=_interleave_gate_up_rows(w1_q),
        w2_weight=w2_q.clone(),
        w13_weight_scale=_interleave_gate_up_rows(w1_scale),
        w2_weight_scale=w2_scale.clone(),
    )

    moe_config = make_dummy_moe_config(
        num_experts=num_experts,
        experts_per_token=topk,
        hidden_dim=hidden_dim,
        intermediate_size=intermediate_dim,
        in_dtype=torch.bfloat16,
    )
    quant_config = make_mxfp4_moe_quant_config(
        Mxfp4MoeBackend.AITER,
        w1_scale=w1_scale_kernel,
        w2_scale=w2_scale_kernel,
    )

    return AiterMxfp4MoeCase(
        hidden_states=hidden_states,
        w1_kernel=w1_kernel,
        w2_kernel=w2_kernel,
        w1_ref=w1_ref,
        w2_ref=w2_ref,
        topk_weights=topk_weights,
        topk_ids=topk_ids,
        moe_config=moe_config,
        quant_config=quant_config,
    )


def _run_fused_moe(
    hidden_states: torch.Tensor,
    w1: torch.Tensor,
    w2: torch.Tensor,
    topk_weights: torch.Tensor,
    topk_ids: torch.Tensor,
    *,
    activation_method: int,
    quant_method: int,
) -> torch.Tensor:
    w1_shuffled, w2_shuffled = _shuffle_moe_weights(w1, w2)
    return torch.ops.vllm.rocm_aiter_fused_moe(
        hidden_states,
        w1_shuffled,
        w2_shuffled,
        topk_weights,
        topk_ids,
        expert_mask=None,
        activation_method=activation_method,
        quant_method=quant_method,
        doweight_stage1=False,
    )


# Custom op tests ---------------------------------------------------------


def test_aiter_fused_moe_custom_op_registered():
    """The main fused-MoE custom op should stay registered for runtime use."""
    _assert_aiter_supported()
    import vllm._aiter_ops as aiter_ops  # noqa: F401

    assert hasattr(torch.ops.vllm, "rocm_aiter_fused_moe")
    assert callable(torch.ops.vllm.rocm_aiter_fused_moe)


def test_aiter_asm_moe_tkw1_custom_op_registered():
    """The tkw1 custom op should stay registered for FP8 apply-router-weight
    paths."""
    _assert_aiter_supported()
    import vllm._aiter_ops as aiter_ops  # noqa: F401

    assert hasattr(torch.ops.vllm, "rocm_aiter_asm_moe_tkw1")
    assert callable(torch.ops.vllm.rocm_aiter_asm_moe_tkw1)


def test_aiter_fused_moe_fake_tensor_support():
    """The fused-MoE op should preserve fake-tensor compatibility for
    torch.compile-style tracing."""
    _assert_aiter_supported()
    import vllm._aiter_ops  # noqa: F401

    num_tokens = 16
    hidden_dim = 1024
    intermediate_dim = 2048
    num_experts = 8
    topk = 2

    hidden_states = torch.randn(
        num_tokens, hidden_dim, dtype=torch.bfloat16, device="cuda"
    )
    w1 = torch.randn(
        num_experts,
        intermediate_dim * 2,
        hidden_dim,
        dtype=torch.bfloat16,
        device="cuda",
    )
    w2 = torch.randn(
        num_experts,
        hidden_dim,
        intermediate_dim,
        dtype=torch.bfloat16,
        device="cuda",
    )
    topk_weights = torch.rand(num_tokens, topk, dtype=torch.float32, device="cuda")
    topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
    topk_ids = _make_topk_ids(num_tokens, num_experts, topk, device="cuda")

    torch.library.opcheck(
        torch.ops.vllm.rocm_aiter_fused_moe,
        (hidden_states, w1, w2, topk_weights, topk_ids),
        kwargs={
            "expert_mask": None,
            "activation_method": 0,
            "quant_method": 0,
            "doweight_stage1": False,
        },
        test_utils=("test_faketensor",),
    )


# Env gating tests --------------------------------------------------------


@pytest.mark.parametrize(
    ("use_aiter", "use_moe", "expected"),
    [
        (True, True, True),
        (True, False, False),
        (False, True, False),
        (False, False, False),
    ],
)
def test_aiter_moe_enablement_follows_env(
    use_aiter: bool,
    use_moe: bool,
    expected: bool,
    monkeypatch: pytest.MonkeyPatch,
):
    """The fused-MoE gate should depend only on the main AITER toggle and the
    MoE-specific toggle."""
    from vllm._aiter_ops import rocm_aiter_ops

    _assert_aiter_supported()

    with monkeypatch.context() as mp:
        mp.setenv("VLLM_ROCM_USE_AITER", "1" if use_aiter else "0")
        mp.setenv("VLLM_ROCM_USE_AITER_MOE", "1" if use_moe else "0")
        _reload_envs()
        rocm_aiter_ops.refresh_env_variables()

        assert rocm_aiter_ops.is_fused_moe_enabled() is expected


@pytest.mark.parametrize(
    ("use_aiter", "use_moe", "use_shared", "expected"),
    [
        (True, True, True, True),
        (True, True, False, False),
        (True, False, True, False),
        (False, True, True, False),
    ],
)
def test_aiter_moe_shared_experts_enablement_follows_env(
    use_aiter: bool,
    use_moe: bool,
    use_shared: bool,
    expected: bool,
    monkeypatch: pytest.MonkeyPatch,
):
    """Shared-expert fusion should only be enabled when the fused-MoE path is
    enabled too."""
    from vllm._aiter_ops import rocm_aiter_ops

    _assert_aiter_supported()

    with monkeypatch.context() as mp:
        mp.setenv("VLLM_ROCM_USE_AITER", "1" if use_aiter else "0")
        mp.setenv("VLLM_ROCM_USE_AITER_MOE", "1" if use_moe else "0")
        mp.setenv(
            "VLLM_ROCM_USE_AITER_FUSION_SHARED_EXPERTS",
            "1" if use_shared else "0",
        )
        _reload_envs()
        rocm_aiter_ops.refresh_env_variables()

        assert rocm_aiter_ops.is_fusion_moe_shared_experts_enabled() is expected


@pytest.mark.parametrize(
    "value,expected_act,expected_env",
    [
        # Unset / auto / legacy "1" all mean the a4w4 default.
        (None, "a4w4", {"AITER_SITUV2_A4W4": "1"}),
        ("auto", "a4w4", {"AITER_SITUV2_A4W4": "1"}),
        ("1", "a4w4", {"AITER_SITUV2_A4W4": "1"}),
        ("a4w4", "a4w4", {"AITER_SITUV2_A4W4": "1"}),
        ("0", "a16w4", {}),
        ("a16w4", "a16w4", {}),
        ("A8W4", "a8w4", {"AITER_SITUV2_A8W4": "1"}),
    ],
)
def test_aiter_moe_situv2_activation_syncs_aiter_env(
    value: str | None,
    expected_act: str,
    expected_env: dict[str, str],
    monkeypatch: pytest.MonkeyPatch,
):
    """Each SiTUv2 activation choice sets exactly its AITER dispatch env."""
    import os

    from vllm._aiter_ops import rocm_aiter_ops

    _assert_aiter_supported()

    with monkeypatch.context() as mp:
        # Stale AITER overrides must be replaced by the vLLM selection.
        mp.setenv("AITER_SITUV2_A8W4", "1")
        mp.setenv("AITER_SITUV2_A4W4", "1")
        mp.setenv("VLLM_ROCM_USE_AITER", "1")
        mp.setenv("VLLM_ROCM_USE_AITER_MOE", "1")
        if value is None:
            mp.delenv("VLLM_ROCM_USE_AITER_MOE_SITUV2", raising=False)
        else:
            mp.setenv("VLLM_ROCM_USE_AITER_MOE_SITUV2", value)
        _reload_envs()
        rocm_aiter_ops.refresh_env_variables()

        assert rocm_aiter_ops.get_fused_moe_situv2_activation() == expected_act
        assert rocm_aiter_ops.is_fused_moe_situv2_gate_up_interleaved() is (
            expected_act == "a8w4"
        )
        assert rocm_aiter_ops.is_fused_moe_situv2_enabled() is (expected_act != "a16w4")
        for name in ("AITER_SITUV2_A8W4", "AITER_SITUV2_A4W4"):
            assert os.environ.get(name) == expected_env.get(name)


def test_aiter_moe_situv2_rejects_unknown_activation(
    monkeypatch: pytest.MonkeyPatch,
):
    with monkeypatch.context() as mp:
        mp.setenv("VLLM_ROCM_USE_AITER_MOE_SITUV2", "a8w8")
        _reload_envs()
        import vllm.envs as envs

        with pytest.raises(ValueError, match="VLLM_ROCM_USE_AITER_MOE_SITUV2"):
            _ = envs.VLLM_ROCM_USE_AITER_MOE_SITUV2


@pytest.mark.parametrize("moe_padding", [True, False])
def test_aiter_moe_padding_env_var(
    moe_padding: bool,
    monkeypatch: pytest.MonkeyPatch,
):
    """The ROCm MoE padding env var should keep its exact parse contract."""
    with monkeypatch.context() as mp:
        mp.setenv("VLLM_ROCM_MOE_PADDING", "1" if moe_padding else "0")
        envs = _reload_envs()
        assert envs.VLLM_ROCM_MOE_PADDING is moe_padding


# a4w4 (FP4 activation) opt-in gating test ---------------------------------


@pytest.mark.parametrize(
    ("env_value", "model_type", "expected"),
    [
        # Unset/"0" is opt-in-only: a4w4 never turns on by itself, even on
        # a validated model type.
        (None, "deepseek_v41", False),
        ("0", "deepseek_v41", False),
        ("1", "deepseek_v41", True),
        ("1", "deepseek_v41_text", True),
    ],
)
def test_aiter_moe_a4w4_dsv4_is_opt_in(
    env_value: str | None,
    model_type: str,
    expected: bool,
    monkeypatch: pytest.MonkeyPatch,
):
    """vllm-project/vllm#58819's a4w4 MoE activation path must stay a plain
    explicit opt-in (VLLM_ROCM_USE_AITER_MOE_A4W4_DSV4=1), with no
    default-on behavior, but is still only permitted on the model types
    this flag has been validated for."""
    import types

    from tests.kernels.moe.utils import make_dummy_moe_config
    from vllm._aiter_ops import rocm_aiter_ops

    _assert_aiter_supported()

    fake_vllm_config = types.SimpleNamespace(
        model_config=types.SimpleNamespace(
            hf_config=types.SimpleNamespace(model_type=model_type)
        )
    )

    with monkeypatch.context() as mp:
        mp.setenv("VLLM_ROCM_USE_AITER", "1")
        mp.setenv("VLLM_ROCM_USE_AITER_MOE", "1")
        if env_value is None:
            mp.delenv("VLLM_ROCM_USE_AITER_MOE_A4W4_DSV4", raising=False)
        else:
            mp.setenv("VLLM_ROCM_USE_AITER_MOE_A4W4_DSV4", env_value)
        _reload_envs()
        rocm_aiter_ops.refresh_env_variables()

        with mock.patch(
            "vllm.config.get_current_vllm_config_or_none",
            return_value=fake_vllm_config,
        ):
            moe_config = make_dummy_moe_config(
                num_experts=4,
                experts_per_token=2,
                hidden_dim=256,
                intermediate_size=512,
                in_dtype=torch.bfloat16,
            )

        assert moe_config.use_mxfp4_w4a4_dsv4 is expected


@pytest.mark.parametrize("model_type", ["deepseek_v4", "gpt_oss", None])
def test_aiter_moe_a4w4_dsv4_rejects_unvalidated_model_type(
    model_type: str | None,
    monkeypatch: pytest.MonkeyPatch,
):
    """Forcing a4w4 on a model type this flag hasn't been GSM8K/AgentX
    validated for must raise: it also flips the MXFP4 weight-shuffle layout
    (ATOM SEPARATED vs INTERLEAVE), which silently produces garbled output
    on the wrong model rather than just being a slower path — the same bug
    class vllm-project/vllm#58819's own second commit fixed for DeepSeek
    V4.1 itself."""
    import types

    from tests.kernels.moe.utils import make_dummy_moe_config
    from vllm._aiter_ops import rocm_aiter_ops

    _assert_aiter_supported()

    fake_vllm_config = types.SimpleNamespace(
        model_config=types.SimpleNamespace(
            hf_config=types.SimpleNamespace(model_type=model_type)
        )
    )

    with monkeypatch.context() as mp:
        mp.setenv("VLLM_ROCM_USE_AITER", "1")
        mp.setenv("VLLM_ROCM_USE_AITER_MOE", "1")
        mp.setenv("VLLM_ROCM_USE_AITER_MOE_A4W4_DSV4", "1")
        _reload_envs()
        rocm_aiter_ops.refresh_env_variables()

        with (
            mock.patch(
                "vllm.config.get_current_vllm_config_or_none",
                return_value=fake_vllm_config,
            ),
            pytest.raises(ValueError, match="VLLM_ROCM_USE_AITER_MOE_A4W4_DSV4"),
        ):
            make_dummy_moe_config(
                num_experts=4,
                experts_per_token=2,
                hidden_dim=256,
                intermediate_size=512,
                in_dtype=torch.bfloat16,
            )


# Dispatch-policy forwarding test ------------------------------------------


@pytest.mark.parametrize("dispatch_policy", [0, 1, 2])
def test_aiter_moe_dispatch_policy_forwarded_to_fused_moe(
    dispatch_policy: int,
    monkeypatch: pytest.MonkeyPatch,
):
    """VLLM_ROCM_AITER_MOE_DISPATCH_POLICY should reach rocm_aiter_ops.fused_moe
    unchanged, the same forwarding AiterExperts.apply does via
    rocm_aiter_fused_experts, for every documented policy value (0=auto,
    1=always single-pass, 2=always multi-pass). See vllm-project/vllm#54966.
    """
    from tests.kernels.moe.utils import make_dummy_moe_config
    from vllm._aiter_ops import rocm_aiter_ops
    from vllm.model_executor.layers.fused_moe.activation import MoEActivation
    from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
        rocm_aiter_fused_experts,
    )

    _assert_aiter_supported()
    case = _make_moe_case(
        num_tokens=8,
        hidden_dim=256,
        intermediate_dim=512,
        num_experts=4,
        topk=2,
        seed=5,
    )
    w1_shuffled, w2_shuffled = _shuffle_moe_weights(case["w1"], case["w2"])
    moe_config = make_dummy_moe_config(
        num_experts=4,
        experts_per_token=2,
        hidden_dim=256,
        intermediate_size=512,
        in_dtype=torch.bfloat16,
    )

    with monkeypatch.context() as mp:
        mp.setenv("VLLM_ROCM_AITER_MOE_DISPATCH_POLICY", str(dispatch_policy))
        _reload_envs()
        rocm_aiter_ops.refresh_env_variables()

        assert rocm_aiter_ops.get_moe_dispatch_policy() == dispatch_policy, (
            "rocm_aiter_ops cached a stale dispatch policy after the env var changed."
        )

        with mock.patch.object(
            rocm_aiter_ops, "fused_moe", wraps=rocm_aiter_ops.fused_moe
        ) as fused_moe_mock:
            rocm_aiter_fused_experts(
                hidden_states=case["hidden_states"],
                w1=w1_shuffled,
                w2=w2_shuffled,
                topk_weights=case["topk_weights"],
                topk_ids=case["topk_ids"],
                moe_config=moe_config,
                activation=MoEActivation.SILU,
                moe_sorting_dispatch_policy=rocm_aiter_ops.get_moe_dispatch_policy(),
            )

        assert (
            fused_moe_mock.call_args.kwargs["moe_sorting_dispatch_policy"]
            == dispatch_policy
        ), (
            "VLLM_ROCM_AITER_MOE_DISPATCH_POLICY was not forwarded to "
            "rocm_aiter_ops.fused_moe."
        )


# Enum tests --------------------------------------------------------------


def test_quant_method_enum_values():
    """The AITER quant-method bridge enum should keep its wire values."""
    from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import QuantMethod

    assert QuantMethod.NO == 0
    assert QuantMethod.PER_TENSOR == 1
    assert QuantMethod.PER_TOKEN == 2
    assert QuantMethod.BLOCK_1X32 == 3
    assert QuantMethod.BLOCK_1X128 == 4
    assert QuantMethod.BLOCK_128x128 == 5


def test_activation_method_enum_values():
    """The AITER activation-method bridge enum should keep its wire values."""
    from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
        ActivationMethod,
    )

    assert ActivationMethod.SILU == 0
    assert ActivationMethod.GELU == 1


# MXFP4 kernel tests ------------------------------------------------------


def test_aiter_mxfp4_quant_scheme_support_matches_gfx950():
    """AITER MXFP4 MoE support should stay gfx950-only."""
    from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
        AiterExperts,
    )
    from vllm.model_executor.layers.quantization.utils.quant_utils import (
        kMxfp4Static,
    )

    assert AiterExperts._supports_quant_scheme(kMxfp4Static, None) is on_gfx950()


@pytest.mark.skipif(not on_gfx950(), reason="gfx950 ROCm only")
def test_aiter_fused_moe_mi350_mxfp4_w4a16_accuracy():
    """The gfx950 AITER MXFP4 W4A16 MoE path should match the dequantized
    MXFP4 reference."""
    from vllm.model_executor.layers.fused_moe.activation import MoEActivation
    from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
        rocm_aiter_fused_experts,
    )

    _assert_aiter_supported()
    case = _make_aiter_mxfp4_moe_case(
        num_tokens=32,
        hidden_dim=512,
        intermediate_dim=1024,
        num_experts=4,
        topk=2,
        seed=11,
    )
    ref_out = ref_moe_forward(
        case.hidden_states,
        case.w1_ref,
        case.w2_ref,
        case.topk_weights,
        case.topk_ids,
    )
    out = rocm_aiter_fused_experts(
        hidden_states=case.hidden_states,
        w1=case.w1_kernel,
        w2=case.w2_kernel,
        topk_weights=case.topk_weights,
        topk_ids=case.topk_ids,
        activation=MoEActivation.SWIGLUOAI,
        quant_config=case.quant_config,
        moe_config=case.moe_config,
        expert_mask=None,
    )

    assert out.shape == case.hidden_states.shape
    _assert_close_budget(
        out.float(),
        ref_out.float(),
        label="mi350_mxfp4_w4a16_accuracy",
        atol=0.1,
        rtol=0.1,
        pass_rate=0.99,
        max_violation_factor=2.0,
    )


@pytest.mark.skipif(not on_gfx950(), reason="gfx950 ROCm only")
def test_aiter_fused_moe_mi350_mxfp4_w4a16_determinism():
    """The gfx950 AITER MXFP4 W4A16 MoE path should stay bitwise
    deterministic."""
    from vllm.model_executor.layers.fused_moe.activation import MoEActivation
    from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
        rocm_aiter_fused_experts,
    )

    _assert_aiter_supported()
    case = _make_aiter_mxfp4_moe_case(
        num_tokens=8,
        hidden_dim=512,
        intermediate_dim=1024,
        num_experts=4,
        topk=2,
        seed=13,
    )

    def run_mxfp4_moe():
        return rocm_aiter_fused_experts(
            hidden_states=case.hidden_states,
            w1=case.w1_kernel,
            w2=case.w2_kernel,
            topk_weights=case.topk_weights,
            topk_ids=case.topk_ids,
            activation=MoEActivation.SWIGLUOAI,
            quant_config=case.quant_config,
            moe_config=case.moe_config,
            expert_mask=None,
        )

    _assert_deterministic(run_mxfp4_moe, n_runs=4)


# FP8 group-quant tests ---------------------------------------------------


@pytest.mark.skipif(
    not (on_gfx942() or on_gfx950()),
    reason="gfx942/gfx950 ROCm only",
)
@pytest.mark.parametrize("num_tokens,hidden_dim", [(16, 2048), (64, 4096), (128, 8192)])
def test_aiter_moe_group_fp8_quant_reconstructs_hidden_states(
    num_tokens: int,
    hidden_dim: int,
):
    """The MoE-facing FP8 group-quant helper should reconstruct hidden states
    within the expected FP8 error budget.

    The raw op's shape and standalone roundtrip contracts are covered in
    ``tests/kernels/core/test_rocm_aiter_ops.py``. This test keeps the
    representative hidden-state sizes that the fused-MoE path actually cares
    about.
    """
    from vllm._aiter_ops import rocm_aiter_ops

    _assert_aiter_supported()
    torch.set_default_device("cuda")
    set_random_seed(1)

    group_size = 128
    hidden_states = torch.randn(num_tokens, hidden_dim, dtype=torch.bfloat16)
    x_fp8, scales = rocm_aiter_ops.group_fp8_quant(hidden_states, group_size)

    assert x_fp8.shape == hidden_states.shape
    assert scales.shape == (num_tokens, (hidden_dim + group_size - 1) // group_size)
    assert scales.dtype == torch.float32

    scales_expanded = scales.repeat_interleave(group_size, dim=1)[:, :hidden_dim]
    dequantized = x_fp8.float() * scales_expanded

    _assert_group_quant_quality(
        dequantized,
        hidden_states,
        label=f"group_fp8_quant shape=({num_tokens}, {hidden_dim})",
    )


# Kernel accuracy tests ---------------------------------------------------


@pytest.mark.parametrize(
    "num_tokens,hidden_dim,intermediate_dim",
    [
        (16, 512, 1024),
        (128, 2048, 4096),
        (2048, 4096, 11008),
    ],
)
def test_aiter_fused_moe_bf16_accuracy(
    num_tokens: int,
    hidden_dim: int,
    intermediate_dim: int,
):
    """The ROCm AITER fused-MoE BF16 path should match the float32 reference
    on representative shapes."""
    from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
        ActivationMethod,
        QuantMethod,
    )

    _assert_aiter_supported()
    case = _make_moe_case(
        num_tokens=num_tokens,
        hidden_dim=hidden_dim,
        intermediate_dim=intermediate_dim,
        num_experts=8,
        topk=2,
        seed=0,
    )
    ref_out = ref_moe_forward(
        case["hidden_states"],
        case["w1"],
        case["w2"],
        case["topk_weights"],
        case["topk_ids"],
        activation="silu",
    )
    out = _run_fused_moe(
        case["hidden_states"],
        case["w1"],
        case["w2"],
        case["topk_weights"],
        case["topk_ids"],
        activation_method=int(ActivationMethod.SILU),
        quant_method=int(QuantMethod.NO),
    )

    assert out.shape == (num_tokens, hidden_dim)
    _assert_close_budget(
        out.float(),
        ref_out,
        label=f"bf16_accuracy shape=({num_tokens}, {hidden_dim}, {intermediate_dim})",
        atol=0.05,
        rtol=0.0,
    )


def test_aiter_fused_moe_gelu_accuracy():
    """The GELU activation variant should stay aligned with the float32
    reference."""
    from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
        ActivationMethod,
        QuantMethod,
    )

    _assert_aiter_supported()
    case = _make_moe_case(
        num_tokens=32,
        hidden_dim=512,
        intermediate_dim=1024,
        num_experts=4,
        topk=2,
        seed=42,
    )
    ref_out = ref_moe_forward(
        case["hidden_states"],
        case["w1"],
        case["w2"],
        case["topk_weights"],
        case["topk_ids"],
        activation="gelu",
    )
    out = _run_fused_moe(
        case["hidden_states"],
        case["w1"],
        case["w2"],
        case["topk_weights"],
        case["topk_ids"],
        activation_method=int(ActivationMethod.GELU),
        quant_method=int(QuantMethod.NO),
    )

    assert out.shape == case["hidden_states"].shape
    _assert_close_budget(
        out.float(),
        ref_out,
        label="gelu_accuracy",
        atol=0.05,
        rtol=0.0,
    )


def test_aiter_fused_moe_determinism():
    """The BF16 fused-MoE kernel should stay bitwise deterministic for the
    same inputs."""
    from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
        ActivationMethod,
        QuantMethod,
    )

    _assert_aiter_supported()
    case = _make_moe_case(
        num_tokens=8,
        hidden_dim=256,
        intermediate_dim=512,
        num_experts=4,
        topk=2,
        seed=2,
    )

    def run_moe():
        return _run_fused_moe(
            case["hidden_states"],
            case["w1"],
            case["w2"],
            case["topk_weights"],
            case["topk_ids"],
            activation_method=int(ActivationMethod.SILU),
            quant_method=int(QuantMethod.NO),
        )

    _assert_deterministic(run_moe, n_runs=4)


# Routed end-to-end tests -------------------------------------------------


@pytest.mark.parametrize(
    "num_tokens,num_experts,topk",
    [
        (1, 8, 2),
        (16, 8, 2),
        (64, 16, 4),
    ],
)
def test_aiter_fused_moe_end_to_end(
    num_tokens: int,
    num_experts: int,
    topk: int,
):
    """The full router-logits to top-k to fused-MoE path should stay accurate."""
    from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
        ActivationMethod,
        QuantMethod,
    )

    _assert_aiter_supported()
    torch.set_default_device("cuda")
    set_random_seed(7)

    hidden_dim = 512
    intermediate_dim = 1024
    hidden_states = torch.randn(num_tokens, hidden_dim, dtype=torch.bfloat16)
    w1 = torch.randn(
        num_experts,
        intermediate_dim * 2,
        hidden_dim,
        dtype=torch.bfloat16,
    ) / math.sqrt(hidden_dim)
    w2 = torch.randn(
        num_experts,
        hidden_dim,
        intermediate_dim,
        dtype=torch.bfloat16,
    ) / math.sqrt(intermediate_dim)

    router_logits = torch.randn(num_tokens, num_experts, device="cuda")
    router_probs = torch.softmax(router_logits, dim=-1)
    topk_weights, topk_ids = torch.topk(router_probs, k=topk, dim=-1)
    topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
    topk_weights = topk_weights.float()
    topk_ids = topk_ids.to(torch.int32)

    ref_out = ref_moe_forward(hidden_states, w1, w2, topk_weights, topk_ids)
    out = _run_fused_moe(
        hidden_states,
        w1,
        w2,
        topk_weights,
        topk_ids,
        activation_method=int(ActivationMethod.SILU),
        quant_method=int(QuantMethod.NO),
    )

    assert out.shape == (num_tokens, hidden_dim)
    assert out.dtype == torch.bfloat16
    _assert_close_budget(
        out.float(),
        ref_out,
        label=f"end_to_end tokens={num_tokens} experts={num_experts} topk={topk}",
        atol=0.05,
        rtol=0.0,
    )


# Arch-specific tests -----------------------------------------------------


@pytest.mark.skipif(
    not (on_gfx942() or on_gfx950()),
    reason="gfx942/gfx950 ROCm only",
)
def test_aiter_fused_moe_mi3xx_bf16_accuracy():
    from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
        ActivationMethod,
        QuantMethod,
    )

    _assert_aiter_supported()
    case = _make_moe_case(
        num_tokens=64,
        hidden_dim=4096,
        intermediate_dim=11008,
        num_experts=8,
        topk=2,
        seed=42,
    )
    ref_out = ref_moe_forward(
        case["hidden_states"],
        case["w1"],
        case["w2"],
        case["topk_weights"],
        case["topk_ids"],
    )
    out = _run_fused_moe(
        case["hidden_states"],
        case["w1"],
        case["w2"],
        case["topk_weights"],
        case["topk_ids"],
        activation_method=int(ActivationMethod.SILU),
        quant_method=int(QuantMethod.NO),
    )

    _assert_close_budget(
        out.float(),
        ref_out,
        label=(f"mi3xx_bf16_accuracy arch={'gfx942' if on_gfx942() else 'gfx950'}"),
        atol=0.05,
        rtol=0.0,
    )


@pytest.mark.skipif(
    not (on_gfx942() or on_gfx950()),
    reason="gfx942/gfx950 ROCm only",
)
def test_aiter_fused_moe_mi3xx_fp8_accuracy():
    """The MI3xx FP8 per-tensor MoE path should stay within the measured FP8
    error budget."""
    from tests.kernels.moe.utils import make_test_weights
    from vllm._aiter_ops import rocm_aiter_ops
    from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
        ActivationMethod,
        QuantMethod,
    )

    _assert_aiter_supported()
    torch.set_default_device("cuda")
    set_random_seed(99)

    num_tokens = 32
    hidden_dim = 512
    intermediate_dim = 1024
    num_experts = 4
    topk = 2

    hidden_states = torch.randn(num_tokens, hidden_dim, dtype=torch.bfloat16) / 10
    fp8_dtype = current_platform.fp8_dtype()
    (w1_bf16, w1_fp8, w1_scale, _), (w2_bf16, w2_fp8, w2_scale, _) = make_test_weights(
        num_experts,
        intermediate_dim * 2,
        hidden_dim,
        torch.bfloat16,
        torch.float8_e4m3fn,
        per_out_ch_quant=False,
    )
    _, a1_scale = rocm_aiter_ops.per_tensor_quant(hidden_states, fp8_dtype)
    router_logits = torch.randn(num_tokens, num_experts, device="cuda")
    router_probs = torch.softmax(router_logits, dim=-1)
    topk_weights, topk_ids = torch.topk(router_probs, k=topk, dim=-1)
    topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
    topk_weights = topk_weights.float()
    topk_ids = topk_ids.to(torch.int32)

    ref_out = ref_moe_forward(
        hidden_states,
        w1_bf16,
        w2_bf16,
        topk_weights,
        topk_ids,
    )

    w1_shuffled, w2_shuffled = _shuffle_moe_weights(w1_fp8, w2_fp8)
    out = torch.ops.vllm.rocm_aiter_fused_moe(
        hidden_states,
        w1_shuffled,
        w2_shuffled,
        topk_weights,
        topk_ids,
        expert_mask=None,
        activation_method=int(ActivationMethod.SILU),
        quant_method=int(QuantMethod.PER_TENSOR),
        doweight_stage1=False,
        w1_scale=w1_scale,
        w2_scale=w2_scale,
        a1_scale=a1_scale,
    )

    assert out.shape == hidden_states.shape
    _assert_close_budget(
        out.float(),
        ref_out,
        label=(f"mi3xx_fp8_accuracy arch={'gfx942' if on_gfx942() else 'gfx950'}"),
        atol=0.02,
        rtol=0.0,
        pass_rate=1.0,
        max_violation_factor=1.5,
    )


# Weight alignment tests --------------------------------------------------
#
# AITER's CK 2stages MoE kernel rejects an intermediate size not divisible by
# its tile width (64 at or below inter_dim 192, 128 above). Some model + TP
# splits land on an unaligned size (e.g. 1792 / TP=8 = 224), so the AITER path
# rounds the intermediate dim up in ``maybe_roundup_sizes`` and allocates the
# weights at the padded size.

ALIGNMENT_HIDDEN = 64
ALIGNMENT_NUM_EXPERTS = 2


def _make_alignment_moe_config(intermediate: int):
    from tests.kernels.moe.utils import make_dummy_moe_config

    return make_dummy_moe_config(
        num_experts=ALIGNMENT_NUM_EXPERTS,
        hidden_dim=ALIGNMENT_HIDDEN,
        intermediate_size=intermediate,
    )


def _make_aiter_method(moe_config):
    """Build the unquantized method with the backend pinned to AITER.

    ``select_unquantized_moe_backend`` needs a real ROCm + AITER runtime, so
    stub it out and set the backend directly.
    """
    from unittest.mock import patch

    from vllm.model_executor.layers.fused_moe.oracle.unquantized import (
        UnquantizedMoeBackend,
    )
    from vllm.model_executor.layers.fused_moe.unquantized_fused_moe_method import (
        UnquantizedFusedMoEMethod,
    )

    with patch(
        "vllm.model_executor.layers.fused_moe.unquantized_fused_moe_method"
        ".select_unquantized_moe_backend",
        return_value=(UnquantizedMoeBackend.AITER, None),
    ):
        return UnquantizedFusedMoEMethod(moe_config)


def _roundup(method, moe_config, intermediate):
    from vllm.model_executor.layers.fused_moe.config import FusedMoEParallelConfig

    return method.maybe_roundup_sizes(
        hidden_size=ALIGNMENT_HIDDEN,
        intermediate_size_per_partition=intermediate,
        act_dtype=moe_config.in_dtype,
        moe_parallel_config=FusedMoEParallelConfig.make_no_parallel(),
    )


@pytest.mark.parametrize(
    ("intermediate", "expected_padded"),
    [
        (64, 64),  # already aligned, untouched
        (96, 128),
        (160, 192),
        (192, 192),  # must stay 192: it is valid and has tuned configs
        (224, 256),  # K2-Horizon-375B at TP=8
        (256, 256),
        (448, 512),  # K2-Horizon-375B at TP=4
        (4096, 4096),
    ],
)
def test_aiter_moe_roundup_pads_intermediate(
    intermediate,
    expected_padded,
    default_vllm_config,
):
    from vllm.model_executor.layers.fused_moe.oracle.unquantized import (
        aiter_moe_intermediate_alignment,
    )

    moe_config = _make_alignment_moe_config(intermediate)
    method = _make_aiter_method(moe_config)

    hidden, padded = _roundup(method, moe_config, intermediate)

    assert padded == expected_padded
    assert hidden == ALIGNMENT_HIDDEN
    # Padding must not move a size across the <= 192 threshold, or the
    # alignment we picked is not the one the kernel dispatches the padded
    # shape to.
    assert aiter_moe_intermediate_alignment(padded) == (
        aiter_moe_intermediate_alignment(intermediate)
    )


@pytest.mark.parametrize("backend_name", ["TRITON", "FLASHINFER_CUTLASS"])
def test_aiter_moe_roundup_is_not_applied_to_other_backends(
    backend_name,
    default_vllm_config,
):
    """Backends with no alignment requirement keep the unaligned size."""
    from vllm.model_executor.layers.fused_moe.oracle.unquantized import (
        UnquantizedMoeBackend,
    )

    intermediate = 224
    moe_config = _make_alignment_moe_config(intermediate)
    method = _make_aiter_method(moe_config)
    method.unquantized_backend = UnquantizedMoeBackend[backend_name]

    _, padded = _roundup(method, moe_config, intermediate)

    assert padded == intermediate


@pytest.mark.parametrize("intermediate", [224, 448])
def test_aiter_moe_reload_zeroes_intermediate_padding(
    intermediate,
    monkeypatch,
    default_vllm_config,
):
    """A reload must re-zero pad lanes left dirty in the weight storage.

    The reload writes back only the logical slices, so the pad lanes hold
    garbage (NaN here) that the AITER conversion must zero. Two passes also
    cover reloading into storage that already holds a previous conversion.
    """
    from unittest.mock import MagicMock

    from vllm.model_executor.layers.fused_moe.oracle.unquantized import (
        aiter_moe_intermediate_alignment,
    )

    alignment = aiter_moe_intermediate_alignment(intermediate)
    padded = -(-intermediate // alignment) * alignment
    assert padded != intermediate

    moe_config = _make_alignment_moe_config(intermediate)
    assert moe_config.is_act_and_mul
    method = _make_aiter_method(moe_config)
    # RoutedExperts records the rounded-up size after maybe_roundup_sizes.
    moe_config.intermediate_size_per_partition = padded
    layer = torch.nn.Module()
    layer.moe_config = moe_config

    method.create_weights(
        layer=layer,
        num_experts=ALIGNMENT_NUM_EXPERTS,
        hidden_size=ALIGNMENT_HIDDEN,
        intermediate_size_per_partition=padded,
        params_dtype=torch.float32,
    )

    # Identity shuffle keeps pad lanes sliceable; both stubs need real ROCm.
    monkeypatch.setattr(
        "vllm.model_executor.layers.fused_moe.oracle.unquantized"
        ".rocm_aiter_ops.shuffle_weights",
        lambda w13, w2: (w13, w2),
    )
    monkeypatch.setattr(
        method,
        "_init_moe_kernel",
        lambda _layer: setattr(method, "moe_kernel", MagicMock()),
    )

    up = padded  # up projection starts one padded block into the fused rows

    for _ in range(2):
        gate = torch.randn(ALIGNMENT_NUM_EXPERTS, intermediate, ALIGNMENT_HIDDEN)
        up_proj = torch.randn(ALIGNMENT_NUM_EXPERTS, intermediate, ALIGNMENT_HIDDEN)
        down = torch.randn(ALIGNMENT_NUM_EXPERTS, ALIGNMENT_HIDDEN, intermediate)

        # NaN the pads, then write only the logical slices like the loader.
        layer.w13_weight.data.fill_(float("nan"))
        layer.w2_weight.data.fill_(float("nan"))
        layer.w13_weight.data[:, :intermediate].copy_(gate)
        layer.w13_weight.data[:, up : up + intermediate].copy_(up_proj)
        layer.w2_weight.data[:, :, :intermediate].copy_(down)

        method.process_weights_after_loading(layer)

        w13 = layer.w13_weight
        w2 = layer.w2_weight
        assert torch.equal(w13[:, :intermediate], gate)
        assert torch.equal(w13[:, up : up + intermediate], up_proj)
        assert torch.equal(w2[:, :, :intermediate], down)
        assert torch.all(w13[:, intermediate:up] == 0)
        assert torch.all(w13[:, up + intermediate :] == 0)
        assert torch.all(w2[:, :, intermediate:] == 0)


def _aiter_accepts_intermediate(intermediate: int, num_tokens: int) -> bool:
    """Run the real AITER CK MoE kernel and report whether it dispatched."""
    from vllm.model_executor.layers.fused_moe.experts.rocm_aiter_moe import (
        ActivationMethod,
        QuantMethod,
    )

    case = _make_moe_case(
        num_tokens=num_tokens,
        hidden_dim=1024,
        intermediate_dim=intermediate,
        num_experts=8,
        topk=2,
        seed=0,
    )
    try:
        _run_fused_moe(
            case["hidden_states"],
            case["w1"],
            case["w2"],
            case["topk_weights"],
            case["topk_ids"],
            activation_method=int(ActivationMethod.SILU),
            quant_method=int(QuantMethod.NO),
        )
    except RuntimeError:
        return False
    return True


@pytest.mark.skipif(
    not on_cdna(),
    reason="CDNA ROCm only",
)
@pytest.mark.parametrize("intermediate", [192, 224, 256])
# With topk=2 and 8 experts, these token counts select block_m 32 / 64 / 128
# respectively (AITER's get_block_size_M, on both 256- and 304-CU parts).
@pytest.mark.parametrize("num_tokens", [1024, 2560, 5120])
def test_aiter_moe_alignment_rule_holds_across_block_m(intermediate, num_tokens):
    """The alignment rule must hold at every reachable ``block_m``.

    Stage 1 checks ``NPerBlock`` and stage 2 ``KPerBlock``; both vary with
    ``block_m``, so a size validated at one tile size says nothing about the
    others.
    """
    from vllm.model_executor.layers.fused_moe.oracle.unquantized import (
        aiter_moe_intermediate_alignment,
    )

    _assert_aiter_supported()

    alignment = aiter_moe_intermediate_alignment(intermediate)
    padded = -(-intermediate // alignment) * alignment

    assert _aiter_accepts_intermediate(intermediate, num_tokens) == (
        intermediate % alignment == 0
    )
    assert _aiter_accepts_intermediate(padded, num_tokens)
