# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

import random
from unittest.mock import patch

import pytest
import torch

from tests.kernels.allclose_default import get_default_atol, get_default_rtol
from tests.kernels.utils import opcheck
from vllm.model_executor.layers.activation import (
    FastGELU,
    FatreluAndMul,
    GeluAndMul,
    MulAndSilu,
    NewGELU,
    QuickGELU,
    ReLUSquaredActivation,
    SiluAndMul,
    SiluAndMulWithClamp,
    SwigluOAIAndMul,
    SwigluStepAndMul,
    swiglustep_and_mul_triton,
)
from vllm.model_executor.layers.fused_moe.activation import (
    ApplyMoEActivationConfig,
    MoEActivation,
    apply_moe_activation,
)
from vllm.model_executor.layers.fused_moe.utils import swiglu_limit_func
from vllm.platforms import current_platform
from vllm.utils.torch_utils import set_random_seed

DTYPES = [torch.half, torch.bfloat16, torch.float]
NUM_TOKENS = [7, 83, 2048]  # Arbitrary values for testing
D = [512, 13824]  # Arbitrary values for testing
SEEDS = [0]
CUDA_DEVICES = [
    f"cuda:{i}" for i in range(1 if torch.accelerator.device_count() == 1 else 2)
]


def test_masked_moe_activation_rejects_unsupported_activation() -> None:
    input = torch.empty(1, 1, 2)
    output = torch.empty(1, 1, 1)
    valid_token_counts = torch.ones(1, dtype=torch.int32)

    with pytest.raises(NotImplementedError, match="relu2"):
        apply_moe_activation(
            MoEActivation.RELU2,
            output,
            input,
            valid_token_counts=valid_token_counts,
        )


def test_moe_silu_clamp_uses_native_xpu_fallback(
    default_vllm_config, monkeypatch
) -> None:
    monkeypatch.setattr(current_platform, "is_xpu", lambda: True)
    clamp_limit = 3.0
    input = torch.tensor([[12.0, -12.0, 8.0, -8.0], [-2.0, 2.0, -4.0, 4.0]])
    output = torch.empty(2, 2)

    apply_moe_activation(
        MoEActivation.SILU,
        output,
        input,
        activation_config=ApplyMoEActivationConfig(clamp_limit=clamp_limit),
    )

    expected = SiluAndMulWithClamp(clamp_limit, compile_native=False).forward_native(
        input
    )
    torch.testing.assert_close(output, expected)


def _assert_masked_moe_activation(
    activation: MoEActivation,
    activation_config: ApplyMoEActivationConfig,
    *,
    dtype: torch.dtype,
    mask_layout: str,
    d: int,
    max_num_tokens: int,
) -> None:
    device = CUDA_DEVICES[0]
    num_experts = 4
    input_dim = 2 * d if activation.is_gated else d
    if mask_layout == "flat":
        input = torch.randn(max_num_tokens, input_dim, dtype=dtype, device=device)
        valid_token_counts = torch.tensor(
            [max_num_tokens // 2], dtype=torch.int32, device=device
        )
        output = torch.full((max_num_tokens, d), 42.0, dtype=dtype, device=device)
    else:
        input = torch.randn(
            num_experts, max_num_tokens, input_dim, dtype=dtype, device=device
        )
        valid_token_counts = torch.tensor(
            [0, 1, max_num_tokens // 2, max_num_tokens],
            dtype=torch.int32,
            device=device,
        )
        output = torch.full(
            (num_experts, max_num_tokens, d), 42.0, dtype=dtype, device=device
        )

    apply_moe_activation(
        activation,
        output,
        input,
        activation_config=activation_config,
        valid_token_counts=valid_token_counts,
    )

    batched_input = input.view(-1, max_num_tokens, input_dim)
    batched_output = output.view(-1, max_num_tokens, d)
    for expert, num_tokens in enumerate(valid_token_counts.cpu().tolist()):
        if num_tokens:
            expected = torch.empty((num_tokens, d), dtype=dtype, device=device)
            apply_moe_activation(
                activation,
                expected,
                batched_input[expert, :num_tokens].clone(),
                activation_config=activation_config,
            )
            torch.testing.assert_close(
                batched_output[expert, :num_tokens],
                expected,
                atol=get_default_atol(output),
                rtol=get_default_rtol(output),
            )
        assert torch.all(batched_output[expert, num_tokens:] == 42.0)


@pytest.mark.parametrize(
    "activation",
    [
        "silu_and_mul",
        "mul_and_silu",
        "gelu",
        "gelu_tanh",
        "fatrelu",
        "swigluoai_and_mul",
        "swiglustep_and_mul",
    ],
)
@pytest.mark.parametrize("num_tokens", NUM_TOKENS)
@pytest.mark.parametrize("d", D)
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("seed", SEEDS)
@pytest.mark.parametrize("device", CUDA_DEVICES)
@torch.inference_mode()
def test_act_and_mul(
    default_vllm_config,
    activation: str,
    num_tokens: int,
    d: int,
    dtype: torch.dtype,
    seed: int,
    device: str,
) -> None:
    set_random_seed(seed)
    torch.set_default_device(device)
    x = torch.randn(num_tokens, 2 * d, dtype=dtype)
    if activation == "silu_and_mul":
        layer = SiluAndMul(compile_native=False)
        fn = torch.ops._C.silu_and_mul
    if activation == "mul_and_silu":
        layer = MulAndSilu()
        fn = torch.ops._C.mul_and_silu
    elif activation == "gelu":
        layer = GeluAndMul(approximate="none")
        fn = torch.ops._C.gelu_and_mul
    elif activation == "gelu_tanh":
        layer = GeluAndMul(approximate="tanh")
        fn = torch.ops._C.gelu_tanh_and_mul
    elif activation == "fatrelu":
        threshold = random.uniform(0, 1)
        layer = FatreluAndMul(threshold)
        fn = torch.ops._C.fatrelu_and_mul
    elif activation == "swigluoai_and_mul":
        layer = SwigluOAIAndMul()
        fn = torch.ops._C.swigluoai_and_mul
    elif activation == "swiglustep_and_mul":
        layer = SwigluStepAndMul()
        fn = swiglustep_and_mul_triton
    out = layer(x)
    ref_out = layer.forward_native(x)
    if activation in ["swigluoai_and_mul", "swiglustep_and_mul"]:
        rtol = {
            # For fp16, change the relative tolerance from 1e-3 to 2e-3
            torch.float16: 2e-3,
            torch.bfloat16: 2e-2,
            torch.float: 1.3e-6,
        }

        def _get_rtol(output) -> float:
            return rtol[output.dtype]

        torch.testing.assert_close(
            out, ref_out, atol=get_default_atol(out), rtol=_get_rtol(out)
        )
    else:
        # The SiluAndMul, MulAndSilu, GELU and FatReLU implementations are
        # equivalent to the native PyTorch implementations, so we can do exact
        # comparison.
        torch.testing.assert_close(out, ref_out, atol=0.0, rtol=0.0)

    d = x.shape[-1] // 2
    output_shape = x.shape[:-1] + (d,)
    out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
    if activation == "fatrelu":
        opcheck(fn, (out, x, threshold))
    elif activation == "swigluoai_and_mul":
        opcheck(fn, (out, x, layer.alpha, layer.limit))
    elif activation != "swiglustep_and_mul":
        opcheck(fn, (out, x))


SWIGLU_LIMITS = [3.0, 7.0, 15.0]


@torch.inference_mode()
def test_swiglu_limit_func_without_routing_uses_output_buffer() -> None:
    x = torch.randn(7, 1024, dtype=torch.bfloat16, device="cuda")
    output = torch.empty(7, 512, dtype=x.dtype, device=x.device)

    swiglu_limit_func(output, x, swiglu_limit=7.0)
    gate, up = x.chunk(2, dim=-1)
    expected = torch.nn.functional.silu(gate.clamp(max=7.0)) * up.clamp(
        min=-7.0, max=7.0
    )

    torch.testing.assert_close(output, expected, atol=2e-2, rtol=2e-2)


@pytest.mark.parametrize(
    ("alpha", "beta"),
    [(1.0, 0.0), (1.702, 0.0), (1.0, 1.0)],
)
@pytest.mark.parametrize("swiglu_limit", SWIGLU_LIMITS)
@pytest.mark.parametrize("num_tokens", NUM_TOKENS)
@pytest.mark.parametrize("d", D)
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("seed", SEEDS)
@pytest.mark.parametrize("device", CUDA_DEVICES)
@torch.inference_mode()
def test_silu_and_mul_with_clamp(
    default_vllm_config,
    alpha: float,
    beta: float,
    swiglu_limit: float,
    num_tokens: int,
    d: int,
    dtype: torch.dtype,
    seed: int,
    device: str,
) -> None:
    """SiluAndMulWithClamp: cuda kernel must match native reference."""
    set_random_seed(seed)
    torch.set_default_device(device)
    # Use large values to ensure clamping is exercised.
    x = torch.randn(num_tokens, 2 * d, dtype=dtype) * swiglu_limit * 2

    default_vllm_config.compilation_config.custom_ops = [
        "none",
        "+silu_and_mul_with_clamp",
    ]
    layer = SiluAndMulWithClamp(
        swiglu_limit,
        alpha=alpha,
        beta=beta,
        compile_native=False,
    )
    if current_platform.is_rocm():
        # forward_hip is always dispatched; the alpha/beta gate is checked
        # inside it at call time rather than picked at construction time, so
        # verify the actual routing by spying on the two candidate methods.
        assert layer._forward_method == layer.forward_hip
        with (
            patch.object(layer, "forward_cuda", wraps=layer.forward_cuda) as cuda_spy,
            patch.object(
                layer, "forward_native", wraps=layer.forward_native
            ) as native_spy,
        ):
            out = layer(x)
        if alpha == 1.0 and beta == 0.0:
            cuda_spy.assert_called_once()
            native_spy.assert_not_called()
        else:
            native_spy.assert_called_once()
            cuda_spy.assert_not_called()
    else:
        assert layer._forward_method == layer.forward_cuda
        out = layer(x)

    ref_out = layer.forward_native(x)

    rtol = {
        torch.float16: 2e-3,
        torch.bfloat16: 2e-2,
        torch.float: 1.3e-6,
    }
    torch.testing.assert_close(
        out, ref_out, atol=get_default_atol(out), rtol=rtol[out.dtype]
    )

    # Verify clamping is actually being applied: the clamped output should
    # differ from the unclamped SiluAndMul output when inputs are large.
    if alpha == 1.0 and beta == 0.0:
        unclamped_out = SiluAndMul.forward_native(x)
        assert not torch.equal(ref_out.float(), unclamped_out.float()), (
            "Input was not large enough to exercise the clamp; increase scale"
        )

    # Verify gate clamping semantics with a controlled scalar case.
    # gate=large_val is clamped to limit first, then silu(limit) * 1.0.
    x_gate = torch.tensor(
        [[swiglu_limit * 20.0, 1.0]], dtype=torch.float32, device=device
    )
    out_gate = SiluAndMulWithClamp(swiglu_limit, compile_native=False)(x_gate)
    expected_gate = torch.nn.functional.silu(
        torch.tensor(swiglu_limit, dtype=torch.float32)
    ).item()
    torch.testing.assert_close(
        out_gate,
        torch.tensor([[expected_gate]], dtype=torch.float32, device=device),
        atol=1e-3,
        rtol=1e-3,
    )

    # Verify up clamping semantics: up >> limit gets clamped to limit.
    x_up = torch.tensor(
        [[1.0, swiglu_limit * 20.0]], dtype=torch.float32, device=device
    )
    out_up = SiluAndMulWithClamp(swiglu_limit, compile_native=False)(x_up)
    silu_1 = torch.nn.functional.silu(torch.tensor(1.0)).item()
    torch.testing.assert_close(
        out_up,
        torch.tensor([[silu_1 * swiglu_limit]], dtype=torch.float32, device=device),
        atol=1e-3,
        rtol=1e-3,
    )

    # opcheck
    out_buf = torch.empty(x.shape[:-1] + (d,), dtype=dtype, device=device)
    opcheck(
        torch.ops._C.silu_and_mul_with_clamp,
        (out_buf, x, swiglu_limit, layer.alpha, layer.beta),
    )


@pytest.mark.parametrize("linear_beta", [-1.0, 2.0])
@pytest.mark.parametrize("dtype", [torch.half, torch.bfloat16])
@torch.inference_mode()
def test_masked_situ_and_mul(
    default_vllm_config,
    linear_beta: float,
    dtype: torch.dtype,
) -> None:
    """Masked SITU computes valid expert rows and preserves padded zeros."""
    device = CUDA_DEVICES[0]
    num_experts, max_num_tokens, d = 4, 7, 512
    beta = 1.5
    input = torch.randn(num_experts, max_num_tokens, 2 * d, dtype=dtype, device=device)
    expert_num_tokens = torch.tensor([0, 1, 4, 7], dtype=torch.int32, device=device)
    output = torch.zeros(num_experts, max_num_tokens, d, dtype=dtype, device=device)

    torch.ops._C.masked_situ_and_mul(
        output, input, expert_num_tokens, beta, linear_beta
    )

    gate, up = input.float().chunk(2, dim=-1)
    expected = beta * torch.tanh(gate / beta) * torch.sigmoid(gate)
    if linear_beta > 0:
        up = linear_beta * torch.tanh(up / linear_beta)
    expected = (expected * up).to(dtype)
    for expert, num_tokens in enumerate(expert_num_tokens.cpu().tolist()):
        torch.testing.assert_close(
            output[expert, :num_tokens],
            expected[expert, :num_tokens],
            atol=get_default_atol(output),
            rtol=get_default_rtol(output),
        )
        assert torch.count_nonzero(output[expert, num_tokens:]) == 0

    opcheck(
        torch.ops._C.masked_situ_and_mul,
        (output, input, expert_num_tokens, beta, linear_beta),
    )


MOE_ACTIVATION_CASES = [
    pytest.param(MoEActivation.SILU, ApplyMoEActivationConfig(), id="silu"),
    pytest.param(
        MoEActivation.SILU, ApplyMoEActivationConfig(clamp_limit=3.0), id="silu_clamp"
    ),
    pytest.param(MoEActivation.GELU, ApplyMoEActivationConfig(), id="gelu"),
    pytest.param(MoEActivation.GELU_TANH, ApplyMoEActivationConfig(), id="gelu_tanh"),
    pytest.param(
        MoEActivation.SITU,
        ApplyMoEActivationConfig(
            activation_situ_beta=1.5,
            activation_situ_linear_beta=2.0,
        ),
        id="situ",
    ),
    pytest.param(MoEActivation.SWIGLUOAI, ApplyMoEActivationConfig(), id="swigluoai"),
    pytest.param(
        MoEActivation.SWIGLUOAI_UNINTERLEAVE,
        ApplyMoEActivationConfig(clamp_limit=3.0, alpha=1.3, beta=0.5),
        id="swigluoai_uninterleave",
    ),
    pytest.param(MoEActivation.SWIGLUSTEP, ApplyMoEActivationConfig(), id="swiglustep"),
    pytest.param(
        MoEActivation.SILU_NO_MUL, ApplyMoEActivationConfig(), id="silu_no_mul"
    ),
    pytest.param(
        MoEActivation.GELU_NO_MUL, ApplyMoEActivationConfig(), id="gelu_no_mul"
    ),
    pytest.param(
        MoEActivation.GELU_TANH_NO_MUL,
        ApplyMoEActivationConfig(),
        id="gelu_tanh_no_mul",
    ),
    pytest.param(
        MoEActivation.RELU2_NO_MUL, ApplyMoEActivationConfig(), id="relu2_no_mul"
    ),
]


@pytest.mark.parametrize(("activation", "activation_config"), MOE_ACTIVATION_CASES)
@torch.inference_mode()
def test_masked_moe_activation_dispatch(
    default_vllm_config,
    activation: MoEActivation,
    activation_config: ApplyMoEActivationConfig,
) -> None:
    _assert_masked_moe_activation(
        activation,
        activation_config,
        dtype=torch.bfloat16,
        mask_layout="batched_experts",
        d=513,
        max_num_tokens=7,
    )


@pytest.mark.parametrize(
    ("activation", "activation_config", "mask_layout"),
    [
        pytest.param(
            MoEActivation.SILU,
            ApplyMoEActivationConfig(),
            "flat",
            id="flat",
        ),
        pytest.param(
            MoEActivation.SITU,
            ApplyMoEActivationConfig(
                activation_situ_beta=1.5,
                activation_situ_linear_beta=2.0,
            ),
            "batched_experts",
            id="batched-experts",
        ),
    ],
)
@torch.inference_mode()
def test_masked_moe_activation_grid_stride(
    default_vllm_config,
    activation: MoEActivation,
    activation_config: ApplyMoEActivationConfig,
    mask_layout: str,
) -> None:
    _assert_masked_moe_activation(
        activation,
        activation_config,
        dtype=torch.half,
        mask_layout=mask_layout,
        d=513,
        max_num_tokens=67,
    )


@torch.inference_mode()
def test_masked_moe_activation_opcheck(default_vllm_config) -> None:
    device = CUDA_DEVICES[0]
    input = torch.randn(2, 3, 64, dtype=torch.half, device=device)
    output = torch.empty(2, 3, 32, dtype=torch.half, device=device)
    valid_token_counts = torch.tensor([1, 3], dtype=torch.int32, device=device)
    opcheck(
        torch.ops._C.masked_moe_activation,
        (output, input, valid_token_counts, "silu", 0.0, 1.0, 0.0, 1.0, -1.0),
    )


@pytest.mark.parametrize(
    "activation",
    [
        (FastGELU, torch.ops._C.gelu_fast),
        (NewGELU, torch.ops._C.gelu_new),
        (QuickGELU, torch.ops._C.gelu_quick),
        (ReLUSquaredActivation, torch.ops._C.relu_squared),
    ],
)
@pytest.mark.parametrize("num_tokens", NUM_TOKENS)
@pytest.mark.parametrize("d", D)
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize("seed", SEEDS)
@pytest.mark.parametrize("device", CUDA_DEVICES)
@torch.inference_mode()
def test_activation(
    default_vllm_config,
    activation: type[torch.nn.Module],
    num_tokens: int,
    d: int,
    dtype: torch.dtype,
    seed: int,
    device: str,
) -> None:
    set_random_seed(seed)
    torch.set_default_device(device)
    x = torch.randn(num_tokens, d, dtype=dtype)
    layer = activation[0]()
    fn = activation[1]
    out = layer(x)
    ref_out = layer.forward_native(x)
    torch.testing.assert_close(
        out, ref_out, atol=get_default_atol(out), rtol=get_default_rtol(out)
    )

    out = torch.empty_like(x)
    opcheck(fn, (out, x))


HUMMING_ACTIVATION_CASES = MOE_ACTIVATION_CASES + [
    pytest.param(MoEActivation.RELU2, ApplyMoEActivationConfig(), id="relu2"),
]
HUMMING_ACTIVATION_CASES += [
    pytest.param(
        MoEActivation.SITU,
        ApplyMoEActivationConfig(
            activation_situ_beta=1.5,
            activation_situ_linear_beta=linear_beta,
        ),
        id=f"situ-linear-beta-{linear_beta}",
    )
    for linear_beta in (None, 0.0, -1.0)
]


@pytest.mark.skipif(not current_platform.is_cuda(), reason="Humming requires CUDA")
@pytest.mark.parametrize(("activation", "activation_config"), HUMMING_ACTIVATION_CASES)
@pytest.mark.parametrize("dtype", DTYPES)
@pytest.mark.parametrize(("num_tokens", "d"), [(1, 512), (7, 768), (83, 512)])
@torch.inference_mode()
def test_humming_activation_matches_framework(
    activation: MoEActivation,
    activation_config: ApplyMoEActivationConfig,
    dtype: torch.dtype,
    num_tokens: int,
    d: int,
) -> None:
    """Compare activation math/layouts without quantization or Hadamard error."""
    pytest.importorskip("humming")
    from humming.ops import process_input

    from vllm.model_executor.layers.quantization.utils.humming.activation import (
        get_humming_activation,
    )

    set_random_seed(0)
    width = d * 2 if activation.is_gated else d
    x = 4 * torch.randn(num_tokens, width, dtype=dtype, device=CUDA_DEVICES[0])
    # Exercise zero, saturation, and both sides of the clamp limits (3 and 7).
    edges = x.new_tensor([0, 0.001, 1, 2.99, 3, 3.01, 6.99, 7, 7.01, 8, 16])
    edges = torch.cat((-edges[1:].flip(0), edges))
    x[0] = edges[torch.arange(width, device=x.device) % edges.numel()]

    actual, _, _ = process_input(
        x,
        quant_mode="none",
        hadamard_block_size=0,
        **get_humming_activation(activation, activation_config),
    )
    expected = torch.empty(num_tokens, d, dtype=dtype, device=x.device)
    if activation == MoEActivation.RELU2:
        # apply_moe_activation has only the non-gated ReLU2 variant.
        gate, up = x.float().chunk(2, dim=-1)
        activated_gate = torch.empty_like(gate)
        apply_moe_activation(MoEActivation.RELU2_NO_MUL, activated_gate, gate.clone())
        expected.copy_(activated_gate * up)
    else:
        # The framework's non-gated ReLU2 path modifies its input in place.
        apply_moe_activation(
            activation, expected, x.clone(), activation_config=activation_config
        )
    torch.testing.assert_close(
        actual,
        expected,
        atol=get_default_atol(expected),
        rtol=get_default_rtol(expected),
    )
