# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Tests for the MOE permute/unpermute kernel.

Run `pytest tests/kernels/test_moe_permute_unpermute.py`.
"""

import numpy as np
import pytest
import torch

from vllm.model_executor.layers.fused_moe import fused_topk
from vllm.model_executor.layers.fused_moe.expert_map_manager import (
    determine_expert_map,
)
from vllm.model_executor.layers.fused_moe.moe_permute_unpermute import (
    MoEPermuteScratch,
    get_moe_permute_scratch,
    moe_permute,
    moe_permute_unpermute_supported,
    moe_prepare_scatter,
    moe_unpermute,
)
from vllm.platforms import current_platform
from vllm.utils.torch_utils import set_random_seed
from vllm.v1.worker.workspace import current_workspace_manager

NUM_EXPERTS = [16, 64, 256]
TOP_KS = [2, 6, 8]
EP_SIZE = [1, 4, 16]
set_random_seed(0)

if current_platform.is_rocm():
    pytest.skip(
        "moe_permute_unpermute_supported is not defined for ROCm",
        allow_module_level=True,
    )


def torch_permute(
    hidden_states: torch.Tensor,
    topk_ids: torch.Tensor,
    #   token_expert_indices: torch.Tensor,
    topk: int,
    n_expert: int,
    n_local_expert: int,
    start_expert: int,
    expert_map: torch.Tensor | None = None,
) -> list[torch.Tensor]:
    n_token = hidden_states.shape[0]
    if expert_map is not None:
        is_local_expert = expert_map[topk_ids] != -1
        not_local_expert = expert_map[topk_ids] == -1
        topk_ids = is_local_expert * (topk_ids - start_expert) + not_local_expert * (
            topk_ids + n_expert
        )
    token_expert_indices = torch.arange(
        0, n_token * topk, dtype=torch.int32, device=hidden_states.device
    ).reshape((n_token, topk))

    sorted_topk_ids, sorted_indices = torch.sort(topk_ids.flatten(), stable=True)
    dst_row_id2src_row_id_map = token_expert_indices.flatten()[sorted_indices]

    expert_first_token_offset = torch.zeros(
        n_local_expert + 1, dtype=torch.int64, device="cuda"
    )
    idx = 0
    for i in range(0, n_local_expert):
        cnt = 0
        while idx < sorted_topk_ids.numel() and sorted_topk_ids[idx] == i:
            cnt += 1
            idx += 1
        expert_first_token_offset[i + 1] = expert_first_token_offset[i] + cnt

    _, src2dst_idx = torch.sort(dst_row_id2src_row_id_map)
    valid_row_idx = []
    permuted_hidden_states = hidden_states[dst_row_id2src_row_id_map // topk, ...]
    src_row_id2dst_row_id_map = torch.arange(
        0, n_token * topk, device="cuda", dtype=torch.int32
    )[src2dst_idx].reshape((n_token, topk))
    valid_row_idx += [i for i in range(expert_first_token_offset[-1])]
    dst_row_id2src_row_id_map[expert_first_token_offset[-1] :] = n_token * topk
    return [
        permuted_hidden_states,
        expert_first_token_offset,
        src_row_id2dst_row_id_map,
        dst_row_id2src_row_id_map,
        valid_row_idx,
    ]


def torch_unpermute(
    permuted_hidden_states: torch.Tensor,
    topk_weights: torch.Tensor,
    topk_ids: torch.Tensor,
    token_expert_indices: torch.Tensor,
    src_row_id2dst_row_id_map: torch.Tensor,
    valid_row_idx: torch.Tensor,
    topk: int,
    n_expert: int,
) -> torch.Tensor:
    # ignore invalid row
    n_hidden = permuted_hidden_states.shape[1]
    mask = torch.zeros(permuted_hidden_states.shape[0], dtype=bool, device="cuda")
    mask[valid_row_idx] = True
    permuted_hidden_states[~mask] = 0

    permuted_hidden_states = permuted_hidden_states[
        src_row_id2dst_row_id_map.flatten(), ...
    ]
    permuted_hidden_states = permuted_hidden_states.view(-1, topk, n_hidden)
    output = (
        (permuted_hidden_states * topk_weights.unsqueeze(2))
        .sum(1)
        .to(permuted_hidden_states.dtype)
    )
    return output


@pytest.mark.parametrize("n_token", [1, 33, 1024, 5000])
@pytest.mark.parametrize("n_hidden", [2048, 7168])
@pytest.mark.parametrize("n_expert", NUM_EXPERTS)
@pytest.mark.parametrize("topk", TOP_KS)
@pytest.mark.parametrize("dtype", [torch.bfloat16])
@pytest.mark.parametrize("ep_size", EP_SIZE)
@pytest.mark.parametrize("use_scratch", [False, True])
def test_moe_permute_unpermute(
    n_token: int,
    n_hidden: int,
    topk: int,
    n_expert: int,
    ep_size: int,
    dtype: torch.dtype,
    use_scratch: bool,
    workspace_init,
):
    if not moe_permute_unpermute_supported():
        pytest.skip("moe_permute_unpermute is not supported on this platform.")
    ep_rank = np.random.randint(0, ep_size)
    expert_map = None
    n_local_expert = n_expert
    if ep_size != 1:
        n_local_expert, expert_map, _ = determine_expert_map(ep_size, ep_rank, n_expert)
        expert_map = expert_map.cuda()
    start_expert = n_local_expert * ep_rank
    set_random_seed(0)
    hidden_states = torch.randn((n_token, n_hidden), device="cuda").to(dtype)
    gating_output = torch.randn((n_token, n_expert), device="cuda").to(dtype)
    topk_weights, topk_ids, token_expert_indices = fused_topk(
        hidden_states, gating_output, topk, False
    )
    (
        gold_permuted_hidden_states,
        gold_expert_first_token_offset,
        gold_inv_permuted_idx,
        gold_permuted_idx,
        valid_row_idx,
    ) = torch_permute(
        hidden_states,
        topk_ids,
        # token_expert_indices,
        topk,
        n_expert,
        n_local_expert,
        start_expert,
        expert_map=expert_map,
    )

    scratch = None
    if use_scratch:
        scratch = get_moe_permute_scratch(
            max_num_tokens=n_token,
            topk=topk,
            num_experts=n_expert,
            num_local_experts=n_local_expert,
            device=hidden_states.device,
            hidden_size=n_hidden,
            hidden_dtype=dtype,
        )

    (
        permuted_hidden_states,
        _,
        expert_first_token_offset,
        inv_permuted_idx,
        _,
    ) = moe_permute(
        hidden_states=hidden_states,
        a1q_scale=None,
        topk_ids=topk_ids,
        n_expert=n_expert,
        n_local_expert=n_local_expert,
        expert_map=expert_map,
        scratch=scratch,
    )

    # check expert_first_token_offset
    torch.testing.assert_close(
        gold_expert_first_token_offset, expert_first_token_offset, atol=0, rtol=0
    )
    # check src_row_id2dst_row_id_map
    torch.testing.assert_close(
        gold_inv_permuted_idx.flatten(), inv_permuted_idx, atol=0, rtol=0
    )

    # check permuted_hidden_states, only valid token
    torch.testing.assert_close(
        gold_permuted_hidden_states[valid_row_idx],
        permuted_hidden_states[valid_row_idx],
        atol=0,
        rtol=0,
    )
    # add a random tensor to simulate group gemm
    result0 = 0.5 * permuted_hidden_states + torch.randn_like(permuted_hidden_states)
    result4 = torch.empty_like(hidden_states)
    moe_unpermute(
        result4, result0, topk_weights, inv_permuted_idx, expert_first_token_offset
    )

    gold4 = torch_unpermute(
        result0,
        topk_weights,
        topk_ids,
        token_expert_indices,
        inv_permuted_idx,
        valid_row_idx,
        topk,
        n_local_expert,
    )
    # check unpermuted hidden
    torch.testing.assert_close(result4, gold4, atol=2e-2, rtol=0)


@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
@pytest.mark.parametrize("n_token", [1, 33, 128])
@pytest.mark.parametrize("topk", [1, 6, 8])
def test_moe_permute_reuses_scratch_buffers(
    dtype: torch.dtype, n_token: int, topk: int, workspace_init
):
    if not moe_permute_unpermute_supported():
        pytest.skip("moe_permute_unpermute is not supported on this platform.")

    n_hidden = 2048
    n_expert = 16
    hidden_states = torch.randn((n_token, n_hidden), device="cuda").to(dtype)
    gating_output = torch.randn((n_token, n_expert), device="cuda").to(dtype)
    _, topk_ids, _ = fused_topk(hidden_states, gating_output, topk, False)

    scratch_config = dict(
        max_num_tokens=n_token,
        topk=topk,
        num_experts=n_expert,
        num_local_experts=n_expert,
        device=hidden_states.device,
        hidden_size=n_hidden,
        hidden_dtype=hidden_states.dtype,
    )
    scratch = get_moe_permute_scratch(**scratch_config)

    first = moe_permute(
        hidden_states=hidden_states,
        a1q_scale=None,
        topk_ids=topk_ids,
        n_expert=n_expert,
        scratch=scratch,
    )
    current_workspace_manager().lock()
    assert get_moe_permute_scratch(**scratch_config) is scratch
    graph = torch.cuda.CUDAGraph()
    with torch.cuda.graph(graph):
        second = moe_permute(
            hidden_states=hidden_states,
            a1q_scale=None,
            topk_ids=topk_ids,
            n_expert=n_expert,
            scratch=get_moe_permute_scratch(**scratch_config),
        )

    for _ in range(2):
        hidden_states.add_(1)
        topk_ids.copy_(topk_ids.roll(1, 0))
        graph.replay()
        expected = moe_permute(hidden_states, None, topk_ids, n_expert)
        for actual, reference in zip(second, expected):
            if actual is not None:
                torch.testing.assert_close(actual, reference)

    (
        permuted_hidden_states_1,
        _,
        expert_first_token_offset_1,
        inv_permuted_idx_1,
        permuted_idx_1,
    ) = first
    (
        permuted_hidden_states_2,
        _,
        expert_first_token_offset_2,
        inv_permuted_idx_2,
        permuted_idx_2,
    ) = second

    torch.testing.assert_close(permuted_hidden_states_1, permuted_hidden_states_2)
    torch.testing.assert_close(expert_first_token_offset_1, expert_first_token_offset_2)
    torch.testing.assert_close(inv_permuted_idx_1, inv_permuted_idx_2)
    torch.testing.assert_close(permuted_idx_1, permuted_idx_2)

    assert (
        permuted_hidden_states_1.untyped_storage().data_ptr()
        == permuted_hidden_states_2.untyped_storage().data_ptr()
    )
    assert (
        expert_first_token_offset_1.untyped_storage().data_ptr()
        == expert_first_token_offset_2.untyped_storage().data_ptr()
    )
    assert (
        inv_permuted_idx_1.untyped_storage().data_ptr()
        == scratch.inv_permuted_idx.untyped_storage().data_ptr()
    )
    assert (
        permuted_idx_1.untyped_storage().data_ptr()
        == scratch.permuted_idx.untyped_storage().data_ptr()
    )


def test_moe_permute_scratch_reused_across_graph_sizes(workspace_init) -> None:
    """Switching graph sizes must not expose metadata left by a larger batch."""
    if not moe_permute_unpermute_supported():
        pytest.skip("moe_permute_unpermute is not supported on this platform.")

    device = torch.device("cuda:0")
    scratch = get_moe_permute_scratch(
        max_num_tokens=128,
        topk=6,
        num_experts=16,
        num_local_experts=16,
        device=device,
        hidden_size=2048,
        hidden_dtype=torch.bfloat16,
    )
    current_workspace_manager().lock()
    runs = []
    for n_token in (128, 1, 33):
        hidden = torch.randn((n_token, 2048), dtype=torch.bfloat16, device=device)
        topk_ids = torch.randint(16, (n_token, 6), device=device)
        moe_permute(hidden, None, topk_ids, 16, scratch=scratch)
        graph = torch.cuda.CUDAGraph()
        with torch.cuda.graph(graph):
            output = moe_permute(hidden, None, topk_ids, 16, scratch=scratch)
        runs.append((graph, hidden, topk_ids, output))

    for graph, hidden, topk_ids, output in runs + runs[::-1]:
        hidden.add_(1)
        topk_ids.copy_((topk_ids + 1) % 16)
        graph.replay()
        expected = moe_permute(hidden, None, topk_ids, 16)
        for actual, reference in zip(output, expected):
            if actual is not None:
                torch.testing.assert_close(actual, reference)


def test_moe_permute_scratch_isolated_across_execution_slots(monkeypatch) -> None:
    """Concurrent graph replays in different ubatches/lanes cannot share scratch."""
    if not moe_permute_unpermute_supported():
        pytest.skip("moe_permute_unpermute is not supported on this platform.")

    from vllm.v1.worker import workspace

    device = torch.device("cuda:0")
    manager = workspace.WorkspaceManager(device, num_ubatches=2, num_lanes=2)
    monkeypatch.setattr(workspace, "_manager", manager)
    ubatch = 0
    monkeypatch.setattr(workspace, "dbo_current_ubatch_id", lambda: ubatch)
    scratch_config = dict(
        max_num_tokens=64,
        topk=6,
        num_experts=16,
        num_local_experts=16,
        device=device,
        hidden_size=2048,
        hidden_dtype=torch.bfloat16,
    )
    runs = []
    for i, n_token in enumerate((1, 7, 33, 64)):
        ubatch = i // 2
        stream = torch.cuda.Stream()
        stream.wait_stream(torch.cuda.current_stream())
        with workspace.use_workspace_lane(i % 2), torch.cuda.stream(stream):
            hidden = torch.randn((n_token, 2048), dtype=torch.bfloat16, device=device)
            topk_ids = torch.randint(16, (n_token, 6), device=device)
            scratch = get_moe_permute_scratch(**scratch_config)
            moe_permute(hidden, None, topk_ids, 16, scratch=scratch)
            graph = torch.cuda.CUDAGraph()
            with torch.cuda.graph(graph, stream=stream):
                output = moe_permute(hidden, None, topk_ids, 16, scratch=scratch)
            runs.append((stream, graph, hidden, topk_ids, output, scratch))
        torch.cuda.current_stream().wait_stream(stream)

    assert len({run[-1].permuted_hidden_states.data_ptr() for run in runs}) == 4
    manager.lock()
    for _ in range(3):
        for i, (stream, graph, hidden, topk_ids, _, scratch) in enumerate(runs):
            ubatch = i // 2
            with workspace.use_workspace_lane(i % 2), torch.cuda.stream(stream):
                assert get_moe_permute_scratch(**scratch_config) is scratch
                hidden.add_(1)
                topk_ids.copy_((topk_ids + 1) % 16)
                graph.replay()
        for stream, _, hidden, topk_ids, output, _ in runs:
            torch.cuda.current_stream().wait_stream(stream)
            expected = moe_permute(hidden, None, topk_ids, 16)
            for actual, reference in zip(output, expected):
                if actual is not None:
                    torch.testing.assert_close(actual, reference)


def test_moe_permute_scratch_without_manager(monkeypatch) -> None:
    """Standalone calls get independent scratch with initialized row indices."""
    if not moe_permute_unpermute_supported():
        pytest.skip("moe_permute_unpermute is not supported on this platform.")

    from vllm.v1.worker import workspace

    monkeypatch.setattr(workspace, "_manager", None)
    config = dict(
        max_num_tokens=4,
        topk=2,
        num_experts=4,
        num_local_experts=4,
        device=torch.device("cuda"),
    )
    first = get_moe_permute_scratch(**config)
    second = get_moe_permute_scratch(**config)

    assert (
        first.token_expert_indices.data_ptr() != second.token_expert_indices.data_ptr()
    )
    torch.testing.assert_close(
        first.token_expert_indices, torch.arange(8, dtype=torch.int32, device="cuda")
    )


def test_moe_permute_ignores_invalid_expert_ids_with_scratch() -> None:
    if not moe_permute_unpermute_supported():
        pytest.skip("moe_permute_unpermute is not supported on this platform.")

    hidden_states = torch.arange(5 * 16, dtype=torch.bfloat16, device="cuda").view(
        5, 16
    )
    topk_ids = torch.tensor([[0], [-1], [1], [4], [2]], device="cuda")
    expert_map = torch.tensor([0, 1, -1, -1], dtype=torch.int32, device="cuda")
    scratch = MoEPermuteScratch(
        max_num_tokens=5,
        topk=1,
        num_experts=4,
        num_local_experts=2,
        device=hidden_states.device,
        hidden_size=16,
        hidden_dtype=hidden_states.dtype,
    )

    permuted, _, expert_offsets, inverse, _ = moe_permute(
        hidden_states=hidden_states,
        a1q_scale=None,
        topk_ids=topk_ids,
        n_expert=4,
        n_local_expert=2,
        expert_map=expert_map,
        scratch=scratch,
    )

    torch.testing.assert_close(
        expert_offsets,
        torch.tensor([0, 1, 2], dtype=torch.int64, device="cuda"),
    )
    torch.testing.assert_close(permuted[:2], hidden_states[[0, 2]])
    assert torch.all(inverse[[1, 3, 4]] >= expert_offsets[-1])
    expected = torch.zeros_like(hidden_states)
    expected[[0, 2]] = hidden_states[[0, 2]]
    output = torch.empty_like(hidden_states)
    moe_unpermute(
        output, permuted, torch.ones(5, 1, device="cuda"), inverse, expert_offsets
    )
    torch.testing.assert_close(output, expected)

    expected_inverse = inverse.clone()
    expected_offsets = expert_offsets.clone()
    expert_offsets, indices = moe_prepare_scatter(topk_ids, expert_map, scratch)
    torch.testing.assert_close(indices.flatten(), expected_inverse)
    torch.testing.assert_close(expert_offsets, expected_offsets)
    graph = torch.cuda.CUDAGraph()
    with torch.cuda.graph(graph):
        expert_offsets, indices = moe_prepare_scatter(topk_ids, expert_map, scratch)
    topk_ids.zero_()
    graph.replay()

    expected_inverse = torch.arange(5, dtype=torch.int32, device="cuda")
    expected_offsets = torch.tensor([0, 5, 5], dtype=torch.int64, device="cuda")
    torch.testing.assert_close(indices.flatten(), expected_inverse)
    torch.testing.assert_close(expert_offsets, expected_offsets)
