import tempfile

import pytest
import torch
import torch.multiprocessing as mp
import transformers

# torch.distributed.tensor is a lazy submodule on torch 2.12+; bind it once so
# downstream ``torch.distributed.tensor.distribute_tensor`` / ``.Shard`` access
# doesn't AttributeError before any explicit import has happened.
try:
    import torch.distributed.tensor  # noqa: F401
except Exception:
    pass

from packaging import version
from test.utils import supports_bfloat16
from transformers.models.llama.configuration_llama import LlamaConfig
from transformers.models.llama.modeling_llama import LlamaMLP
from transformers.models.mixtral.configuration_mixtral import MixtralConfig
from transformers.models.phi3.configuration_phi3 import Phi3Config
from transformers.models.phi3.modeling_phi3 import Phi3MLP

import liger_kernel.ops.swiglu as swiglu_ops

from liger_kernel.ops import LigerSiLUMulFunction
from liger_kernel.transformers.functional import liger_swiglu
from liger_kernel.transformers.swiglu import LigerBlockSparseTop2MLP
from liger_kernel.transformers.swiglu import LigerExperts
from liger_kernel.transformers.swiglu import LigerFalconH1SwiGLUMLP
from liger_kernel.transformers.swiglu import LigerPhi3SwiGLUMLP
from liger_kernel.transformers.swiglu import LigerSwiGLUMLP
from liger_kernel.utils import infer_comm_backend
from liger_kernel.utils import infer_device

IS_TRANSFORMERS_V5_OR_LATER = version.parse(transformers.__version__) >= version.parse("5.0.0")
if IS_TRANSFORMERS_V5_OR_LATER:
    from transformers.models.mixtral.modeling_mixtral import MixtralExperts
else:
    from transformers.models.mixtral.modeling_mixtral import MixtralBlockSparseTop2MLP

device = infer_device()

LLAMA_CONFIG = LlamaConfig(
    hidden_size=4096,
    intermediate_size=11008,
    hidden_act="silu",
)
PHI3_CONFIG = Phi3Config(
    hidden_size=4096,
    intermediate_size=11008,
    hidden_act="silu",
)
SLEEP_SECONDS = 0.1


# ---------------------------------------------------------------------------
# Shared plain SiLU-Mul coverage (kept identical across triton/cutedsl/cutile test files)
# ---------------------------------------------------------------------------
_SILU_SHAPES = [
    (4, 2048),  # 2D, power-of-2 aligned
    (2, 256, 512),  # 3D, small
    (6, 42, 431),  # non-aligned / predicated path, odd width
    (1, 1, 7),  # tiny
    (3, 1023),  # odd width
    (4, 11008),  # non-power-of-2 "cliff" (Qwen2.5-7B)
    (2, 3, 13824),  # non-power-of-2 "cliff" (Qwen2.5-3B), 3D
    (2, 18944),  # non-power-of-2 "cliff" (Qwen2.5-14B)
]
_SILU_MULTIPLIERS = [(1.0, 1.0), (1.5, 0.75), (0.7, 1.3)]

# fp32, fp16, bf16 (bf16 guarded) — used by the harmonized SiLU-Mul tests below.
_SILU_DTYPES = [
    pytest.param(torch.float32, id="fp32"),
    pytest.param(torch.float16, id="fp16"),
    pytest.param(
        torch.bfloat16,
        marks=pytest.mark.skipif(not supports_bfloat16(), reason="bfloat16 not supported on this GPU"),
        id="bf16",
    ),
]
# Multiplier sweep runs on the fp32 + bf16 subset (fp16 adds no path coverage here).
_SILU_MULT_DTYPES = [
    pytest.param(torch.float32, id="fp32"),
    pytest.param(
        torch.bfloat16,
        marks=pytest.mark.skipif(not supports_bfloat16(), reason="bfloat16 not supported on this GPU"),
        id="bf16",
    ),
]


def _ref_silu_mul(a, b, gate=1.0, down=1.0):
    """Ground-truth SiLU-Mul reference (silu(a * gate) * b * down) computed in fp32."""
    return torch.nn.functional.silu(a.float() * gate) * b.float() * down


def _silu_tol(dtype):
    """Harmonized (atol, rtol) for the shared plain SiLU-Mul tests."""
    if dtype == torch.float32:
        return 2e-4, 1e-4
    return 1e-2, 1e-2  # fp16 / bf16


@pytest.mark.parametrize(
    "bsz, seq_len, hidden_size, intermediate_size",
    [
        (2, 256, 256, 512),
        # weird shapes
        (6, 42, 123, 431),
    ],
)
@pytest.mark.parametrize(
    "dtype, atol, rtol",
    [
        # atol is for small values: they have more difference, so set atol higher
        # rtol is for larger values: they are very close, so set rtol lower
        (torch.float32, 1e-0, 1e-5),
        # TODO: we should find a better way to tune this. 1e4 is too large apparently
        pytest.param(
            torch.bfloat16,
            1e4,
            1e-2,
            marks=pytest.mark.skipif(not supports_bfloat16(), reason="bfloat16 not supported on this GPU"),
        ),
    ],
)
def test_correctness_llamamlp(bsz, seq_len, hidden_size, intermediate_size, dtype, atol, rtol):
    _input = torch.randn(bsz, seq_len, hidden_size, device=device, dtype=dtype)

    x1 = _input.clone().requires_grad_(True)
    x2 = _input.clone().requires_grad_(True)

    # initialize weights
    G = torch.randn(hidden_size, intermediate_size, device=device, dtype=dtype)
    U = torch.randn(hidden_size, intermediate_size, device=device, dtype=dtype)
    D = torch.randn(intermediate_size, hidden_size, device=device, dtype=dtype)

    llama_mlp = LlamaMLP(config=LLAMA_CONFIG).to(device).to(dtype)
    llama_mlp.gate_proj.weight.data = G.T
    llama_mlp.up_proj.weight.data = U.T
    llama_mlp.down_proj.weight.data = D.T

    liger_mlp = LigerSwiGLUMLP(config=LLAMA_CONFIG).to(device).to(dtype)
    liger_mlp.gate_proj.weight.data = G.T
    liger_mlp.up_proj.weight.data = U.T
    liger_mlp.down_proj.weight.data = D.T

    y1 = llama_mlp(x1)
    y2 = liger_mlp(x2)

    assert torch.allclose(y1, y2, atol=atol, rtol=rtol)

    dy = torch.randn_like(y1)

    y1.backward(dy.clone(), retain_graph=True)
    y2.backward(dy.clone(), retain_graph=True)

    assert torch.allclose(
        llama_mlp.gate_proj.weight.grad,
        liger_mlp.gate_proj.weight.grad,
        atol=atol,
        rtol=rtol,
    )
    assert torch.allclose(
        llama_mlp.up_proj.weight.grad,
        liger_mlp.up_proj.weight.grad,
        atol=atol,
        rtol=rtol,
    )
    assert torch.allclose(
        llama_mlp.down_proj.weight.grad,
        liger_mlp.down_proj.weight.grad,
        atol=atol,
        rtol=rtol,
    )

    assert torch.allclose(x1.grad, x2.grad, atol=atol, rtol=rtol)


@pytest.mark.skipif(IS_TRANSFORMERS_V5_OR_LATER, reason="Skip for transformers >= v5.0.0")
@pytest.mark.parametrize(
    "bsz, seq_len, hidden_size, intermediate_size",
    [
        (2, 256, 256, 512),
        # weird shapes
        (6, 42, 123, 431),
    ],
)
@pytest.mark.parametrize(
    "dtype, atol, rtol",
    [
        # atol is for small values: they have more difference, so set atol higher
        # rtol is for larger values: they are very close, so set rtol lower
        (torch.float32, 1e-0, 1e-5),
        # TODO: we should find a better way to tune this. 1e4 is too large apparently
        pytest.param(
            torch.bfloat16,
            1e4,
            1e-2,
            marks=pytest.mark.skipif(not supports_bfloat16(), reason="bfloat16 not supported on this GPU"),
        ),
    ],
)
def test_correctness_mixtralblocksparsetop2mlp(bsz, seq_len, hidden_size, intermediate_size, dtype, atol, rtol):
    MIXTRAL_CONFIG = MixtralConfig(
        num_local_experts=8,
        hidden_size=hidden_size,
        intermediate_size=intermediate_size,
        hidden_act="silu",
        num_experts_per_tok=2,
    )

    _input = torch.randn(bsz, seq_len, hidden_size, device=device, dtype=dtype)
    x1 = _input.clone().requires_grad_(True)
    x2 = _input.clone().requires_grad_(True)

    # initialize weights
    G = torch.randn(hidden_size, intermediate_size, device=device, dtype=dtype)
    U = torch.randn(intermediate_size, hidden_size, device=device, dtype=dtype)
    D = torch.randn(hidden_size, intermediate_size, device=device, dtype=dtype)

    mixtral_blocksparsetop2mlp = MixtralBlockSparseTop2MLP(config=MIXTRAL_CONFIG).to(device).to(dtype)
    mixtral_blocksparsetop2mlp.w1.weight.data = G.T
    mixtral_blocksparsetop2mlp.w2.weight.data = U.T
    mixtral_blocksparsetop2mlp.w3.weight.data = D.T

    liger_blocksparsetop2mlp = LigerBlockSparseTop2MLP(config=MIXTRAL_CONFIG).to(device).to(dtype)
    liger_blocksparsetop2mlp.w1.weight.data = G.T
    liger_blocksparsetop2mlp.w2.weight.data = U.T
    liger_blocksparsetop2mlp.w3.weight.data = D.T

    y1 = mixtral_blocksparsetop2mlp(x1)
    y2 = liger_blocksparsetop2mlp(x2)

    assert torch.allclose(y1, y2, atol=atol, rtol=rtol)

    dy = torch.randn_like(y1)

    y1.backward(dy.clone(), retain_graph=True)
    y2.backward(dy.clone(), retain_graph=True)

    assert torch.allclose(
        mixtral_blocksparsetop2mlp.w1.weight.grad,
        liger_blocksparsetop2mlp.w1.weight.grad,
        atol=atol,
        rtol=rtol,
    )
    assert torch.allclose(
        mixtral_blocksparsetop2mlp.w2.weight.grad,
        liger_blocksparsetop2mlp.w2.weight.grad,
        atol=atol,
        rtol=rtol,
    )
    assert torch.allclose(
        mixtral_blocksparsetop2mlp.w3.weight.grad,
        liger_blocksparsetop2mlp.w3.weight.grad,
        atol=atol,
        rtol=rtol,
    )

    assert torch.allclose(x1.grad, x2.grad, atol=atol, rtol=rtol)


@pytest.mark.skipif(not IS_TRANSFORMERS_V5_OR_LATER, reason="Skip for transformers < v5.0.0")
@pytest.mark.parametrize(
    "bsz, seq_len, hidden_size, intermediate_size",
    [
        (2, 256, 256, 512),
        # weird shapes
        (6, 42, 123, 431),
    ],
)
@pytest.mark.parametrize(
    "dtype, atol, rtol",
    [
        # TF32 accumulation order differences between Triton and cuBLAS, propagated
        # through 2 GEMMs (error ~ sqrt(I) * gate_proj_error), cause ~3-5% relative
        # error on small-valued outputs and ~5 absolute error near zero.
        (torch.float32, 30, 5e-2),
        # bf16: same Triton-vs-cuBLAS divergence plus bf16 quantization of intermediate
        # activations. Near-zero outputs can differ by ~8-16 absolute due to cancellation
        # at 7-bit mantissa precision.
        pytest.param(
            torch.bfloat16,
            100.0,
            5e-2,
            marks=pytest.mark.skipif(not supports_bfloat16(), reason="bfloat16 not supported on this GPU"),
        ),
    ],
)
def test_correctness_mixtralexperts(bsz, seq_len, hidden_size, intermediate_size, dtype, atol, rtol):
    MIXTRAL_CONFIG = MixtralConfig(
        num_local_experts=8,
        hidden_size=hidden_size,
        intermediate_size=intermediate_size,
        experts_implementation="eager",
        hidden_act="silu",
        num_experts_per_tok=2,
    )

    _input = torch.randn(bsz * seq_len, hidden_size, device=device, dtype=dtype)

    x1 = _input.clone().requires_grad_(True)
    x2 = _input.clone().requires_grad_(True)

    # match shape: (num_experts, 2 * intermediate_dim, hidden_dim)
    GU = torch.randn(
        MIXTRAL_CONFIG.num_local_experts,
        2 * intermediate_size,
        hidden_size,
        device=device,
        dtype=dtype,
        requires_grad=True,
    )
    # match shape: (num_experts, hidden_dim, intermediate_dim)
    D = torch.randn(
        MIXTRAL_CONFIG.num_local_experts, hidden_size, intermediate_size, device=device, dtype=dtype, requires_grad=True
    )

    # Generate random router logits and do topk
    router_logits = torch.randn(bsz * seq_len, MIXTRAL_CONFIG.num_local_experts, device=device, dtype=dtype)
    router_logits = router_logits.softmax(dim=-1)
    top_k_weights, top_k_index = router_logits.topk(k=MIXTRAL_CONFIG.num_experts_per_tok, dim=-1)
    top_k_weights = top_k_weights / (top_k_weights.sum(dim=-1, keepdim=True) + 1e-9)

    mixtral_experts = MixtralExperts(config=MIXTRAL_CONFIG).to(device).to(dtype)
    mixtral_experts.gate_up_proj.data = GU.clone().detach()
    mixtral_experts.down_proj.data = D.clone().detach()

    liger_experts = LigerExperts(config=MIXTRAL_CONFIG).to(device).to(dtype)
    liger_experts.gate_up_proj.data = GU.clone().detach()
    liger_experts.down_proj.data = D.clone().detach()

    mixtral_experts.gate_up_proj.requires_grad_()
    mixtral_experts.down_proj.requires_grad_()
    liger_experts.gate_up_proj.requires_grad_()
    liger_experts.down_proj.requires_grad_()

    y1 = mixtral_experts(x1, top_k_index, top_k_weights)
    y2 = liger_experts(x2, top_k_index, top_k_weights)

    def _assert_close(a, b, label):
        af, bf = a.float(), b.float()
        diff = (af - bf).abs()
        rel = diff / (bf.abs() + 1e-9)
        abs_idx = diff.argmax()
        rel_idx = rel.argmax()
        print(
            f"\n  [{label}]"
            f"\n    max_abs={diff.max():.4g}  ref={af.flatten()[abs_idx]:.4g}  liger={bf.flatten()[abs_idx]:.4g}"
            f"\n    max_rel={rel.max():.4g}   ref={af.flatten()[rel_idx]:.4g}  liger={bf.flatten()[rel_idx]:.4g}"
        )
        torch.testing.assert_close(a, b, atol=float(atol), rtol=float(rtol))

    _assert_close(y1, y2, "forward output")

    dy = torch.randn_like(y1)

    y1.backward(dy.clone(), retain_graph=True)
    y2.backward(dy.clone(), retain_graph=True)

    _assert_close(mixtral_experts.gate_up_proj.grad, liger_experts.gate_up_proj.grad, "gate_up_proj.grad")
    _assert_close(mixtral_experts.down_proj.grad, liger_experts.down_proj.grad, "down_proj.grad")
    _assert_close(x1.grad, x2.grad, "x.grad")


def test_transformers_swiglu_preserves_legacy_function(monkeypatch):
    import liger_kernel.transformers.swiglu as swiglu_module

    observed = []
    expected = object()

    def fake_apply(*args):
        observed.append(args)
        return expected

    monkeypatch.setattr(swiglu_module.LigerSiLUMulFunction, "apply", fake_apply)
    a = torch.randn(2, 8)
    b = torch.randn(2, 8)

    actual = swiglu_module._swiglu_dispatch(a, b, 2.0, 3.0)

    assert actual is expected
    assert observed == [(a, b, 2.0, 3.0)]


@pytest.mark.parametrize(
    "bsz, seq_len, hidden_size, intermediate_size",
    [
        (2, 256, 256, 512),
        # weird shapes
        (6, 42, 123, 431),
    ],
)
@pytest.mark.parametrize(
    "dtype, atol, rtol",
    [
        # atol is for small values: they have more difference, so set atol higher
        # rtol is for larger values: they are very close, so set rtol lower
        (torch.float32, 1e-0, 1e-5),
        # TODO: we should find a better way to tune this. 1e4 is too large apparently
        pytest.param(
            torch.bfloat16,
            1e4,
            1e-2,
            marks=pytest.mark.skipif(not supports_bfloat16(), reason="bfloat16 not supported on this GPU"),
        ),
    ],
)
def test_correctness_phi3mlp(bsz, seq_len, hidden_size, intermediate_size, dtype, atol, rtol):
    _input = torch.randn(bsz, seq_len, hidden_size, device=device, dtype=dtype)

    x1 = _input.clone().requires_grad_(True)
    x2 = _input.clone().requires_grad_(True)

    # initialize weights
    GU = torch.randn(hidden_size, intermediate_size * 2, device=device, dtype=dtype)
    D = torch.randn(intermediate_size, hidden_size, device=device, dtype=dtype)

    phi3_mlp = Phi3MLP(config=PHI3_CONFIG).to(device).to(dtype)
    phi3_mlp.gate_up_proj.weight.data = GU.T
    phi3_mlp.down_proj.weight.data = D.T

    liger_mlp = LigerPhi3SwiGLUMLP(config=PHI3_CONFIG).to(device).to(dtype)
    liger_mlp.gate_up_proj.weight.data = GU.T
    liger_mlp.down_proj.weight.data = D.T

    y1 = phi3_mlp(x1)
    y2 = liger_mlp(x2)

    assert torch.allclose(y1, y2, atol=atol, rtol=rtol)

    dy = torch.randn_like(y1)

    y1.backward(dy.clone(), retain_graph=True)
    y2.backward(dy.clone(), retain_graph=True)

    assert torch.allclose(
        phi3_mlp.gate_up_proj.weight.grad,
        liger_mlp.gate_up_proj.weight.grad,
        atol=atol,
        rtol=rtol,
    )
    assert torch.allclose(
        phi3_mlp.down_proj.weight.grad,
        liger_mlp.down_proj.weight.grad,
        atol=atol,
        rtol=rtol,
    )

    assert torch.allclose(x1.grad, x2.grad, atol=atol, rtol=rtol)


@pytest.mark.parametrize(
    "bsz, seq_len, size",
    [
        (2, 8, 8),
        (9, 7, 41),
    ],
)
@pytest.mark.parametrize(
    "dtype, atol, rtol",
    [
        # atol is for small values: they have more difference, so set atol higher
        # rtol is for larger values: they are very close, so set rtol lower
        (torch.float32, 1e-0, 1e-5),
        # TODO: we should find a better way to tune this. 1e4 is too large apparently
        (torch.bfloat16, 1e4, 1e-2),
    ],
)
def test_correctness_functional(bsz, seq_len, size, dtype, atol, rtol):
    _input = torch.randn(bsz, seq_len, size, device=device, dtype=dtype)
    _b = torch.randn(bsz, seq_len, size, device=device, dtype=dtype)

    x1 = _input.clone().requires_grad_(True)
    x2 = _input.clone().requires_grad_(True)

    b1 = _b.clone().requires_grad_(True)
    b2 = _b.clone().requires_grad_(True)

    y1 = liger_swiglu(a=x1, b=b1)
    y2 = LigerSiLUMulFunction.apply(x2, b2)

    assert torch.allclose(y1, y2, atol=atol, rtol=rtol)

    # Test backward pass
    grad_output = torch.randn_like(y1)

    y1.backward(grad_output)
    y2.backward(grad_output)

    # Check if gradients are close for x
    assert torch.allclose(x1.grad, x2.grad, atol=atol, rtol=rtol)
    assert torch.allclose(b1.grad, b2.grad, atol=atol, rtol=rtol)


def _torch_silu_mul_ref(a, b, gate_multiplier, down_multiplier):
    """Pure-PyTorch reference for silu(a * gate_mult) * b * down_mult."""
    scaled = a * gate_multiplier
    return torch.nn.functional.silu(scaled) * b * down_multiplier


@pytest.mark.parametrize(
    "bsz, seq_len, size",
    [
        (2, 8, 8),
        (9, 7, 41),
    ],
)
@pytest.mark.parametrize(
    "gate_multiplier, down_multiplier",
    [
        (0.7, 1.3),
        (1.5, 0.5),
        (1.0, 1.0),  # degenerate case — must match the no-multiplier path
    ],
)
@pytest.mark.parametrize(
    "dtype, atol, rtol",
    [
        (torch.float32, 1e-3, 1e-5),
        pytest.param(
            torch.bfloat16,
            1e-1,
            1e-2,
            marks=pytest.mark.skipif(not supports_bfloat16(), reason="bfloat16 not supported on this GPU"),
        ),
    ],
)
def test_correctness_silumul_with_multipliers(bsz, seq_len, size, gate_multiplier, down_multiplier, dtype, atol, rtol):
    _a = torch.randn(bsz, seq_len, size, device=device, dtype=dtype)
    _b = torch.randn(bsz, seq_len, size, device=device, dtype=dtype)

    a1 = _a.clone().detach().requires_grad_(True)
    b1 = _b.clone().detach().requires_grad_(True)
    a2 = _a.clone().detach().requires_grad_(True)
    b2 = _b.clone().detach().requires_grad_(True)

    y_ref = _torch_silu_mul_ref(a1, b1, gate_multiplier, down_multiplier)
    y_liger = LigerSiLUMulFunction.apply(a2, b2, gate_multiplier, down_multiplier)

    torch.testing.assert_close(y_ref, y_liger, atol=atol, rtol=rtol)

    grad = torch.randn_like(y_ref)
    y_ref.backward(grad.clone())
    y_liger.backward(grad.clone())

    torch.testing.assert_close(a1.grad, a2.grad, atol=atol, rtol=rtol)
    torch.testing.assert_close(b1.grad, b2.grad, atol=atol, rtol=rtol)


def test_silumul_default_multipliers_backward_compat():
    """Calling LigerSiLUMulFunction.apply(a, b) without multipliers must behave exactly as before."""
    _a = torch.randn(4, 16, 32, device=device, dtype=torch.float32)
    _b = torch.randn(4, 16, 32, device=device, dtype=torch.float32)

    a1 = _a.clone().detach().requires_grad_(True)
    b1 = _b.clone().detach().requires_grad_(True)
    a2 = _a.clone().detach().requires_grad_(True)
    b2 = _b.clone().detach().requires_grad_(True)

    y_default = LigerSiLUMulFunction.apply(a1, b1)
    y_explicit = LigerSiLUMulFunction.apply(a2, b2, 1.0, 1.0)

    torch.testing.assert_close(y_default, y_explicit)

    grad = torch.randn_like(y_default)
    y_default.backward(grad.clone())
    y_explicit.backward(grad.clone())

    torch.testing.assert_close(a1.grad, a2.grad)
    torch.testing.assert_close(b1.grad, b2.grad)


class _FalconH1MLPRef(torch.nn.Module):
    """Pure-PyTorch reference matching Falcon H1's MLP forward from issue #936."""

    def __init__(self, hidden_size, intermediate_size, gate_multiplier, down_multiplier):
        super().__init__()
        self.gate_proj = torch.nn.Linear(hidden_size, intermediate_size, bias=False)
        self.up_proj = torch.nn.Linear(hidden_size, intermediate_size, bias=False)
        self.down_proj = torch.nn.Linear(intermediate_size, hidden_size, bias=False)
        self.gate_multiplier = gate_multiplier
        self.down_multiplier = down_multiplier

    def forward(self, x):
        gate = self.gate_proj(x)
        up = self.up_proj(x)
        activated = torch.nn.functional.silu(gate * self.gate_multiplier) * up
        return self.down_proj(activated) * self.down_multiplier


@pytest.mark.parametrize(
    "bsz, seq_len, hidden_size, intermediate_size",
    [
        (2, 256, 256, 512),
        (6, 42, 123, 431),
    ],
)
@pytest.mark.parametrize(
    "gate_multiplier, down_multiplier",
    [
        (0.7, 1.3),
        (1.5, 0.5),
    ],
)
@pytest.mark.parametrize(
    "dtype, atol, rtol",
    [
        (torch.float32, 1e-0, 1e-5),
        pytest.param(
            torch.bfloat16,
            1e4,
            1e-2,
            marks=pytest.mark.skipif(not supports_bfloat16(), reason="bfloat16 not supported on this GPU"),
        ),
    ],
)
def test_correctness_falcon_h1_mlp(
    bsz, seq_len, hidden_size, intermediate_size, gate_multiplier, down_multiplier, dtype, atol, rtol
):
    """Parity test for LigerFalconH1SwiGLUMLP against a pure-PyTorch reference.

    A pure-PyTorch reference is used rather than HF's FalconH1MLP so the test
    doesn't depend on transformers exposing FalconH1MLP at module scope.
    """

    class _FakeConfig:
        def __init__(self):
            self.hidden_size = hidden_size
            self.intermediate_size = intermediate_size
            self.hidden_act = "silu"
            self.mlp_bias = False
            self.mlp_multipliers = (gate_multiplier, down_multiplier)

    config = _FakeConfig()

    _input = torch.randn(bsz, seq_len, hidden_size, device=device, dtype=dtype)
    x1 = _input.clone().requires_grad_(True)
    x2 = _input.clone().requires_grad_(True)

    G = torch.randn(hidden_size, intermediate_size, device=device, dtype=dtype)
    U = torch.randn(hidden_size, intermediate_size, device=device, dtype=dtype)
    D = torch.randn(intermediate_size, hidden_size, device=device, dtype=dtype)

    ref_mlp = _FalconH1MLPRef(hidden_size, intermediate_size, gate_multiplier, down_multiplier).to(device).to(dtype)
    ref_mlp.gate_proj.weight.data = G.T.contiguous()
    ref_mlp.up_proj.weight.data = U.T.contiguous()
    ref_mlp.down_proj.weight.data = D.T.contiguous()

    liger_mlp = LigerFalconH1SwiGLUMLP(config=config).to(device).to(dtype)
    liger_mlp.gate_proj.weight.data = G.T.contiguous()
    liger_mlp.up_proj.weight.data = U.T.contiguous()
    liger_mlp.down_proj.weight.data = D.T.contiguous()

    y1 = ref_mlp(x1)
    y2 = liger_mlp(x2)

    torch.testing.assert_close(y1, y2, atol=atol, rtol=rtol)

    dy = torch.randn_like(y1)
    y1.backward(dy.clone(), retain_graph=True)
    y2.backward(dy.clone(), retain_graph=True)

    torch.testing.assert_close(ref_mlp.gate_proj.weight.grad, liger_mlp.gate_proj.weight.grad, atol=atol, rtol=rtol)
    torch.testing.assert_close(ref_mlp.up_proj.weight.grad, liger_mlp.up_proj.weight.grad, atol=atol, rtol=rtol)
    torch.testing.assert_close(ref_mlp.down_proj.weight.grad, liger_mlp.down_proj.weight.grad, atol=atol, rtol=rtol)
    torch.testing.assert_close(x1.grad, x2.grad, atol=atol, rtol=rtol)


def _test_dtensor_liger_silumul(
    rank,
    world_size,
    bsz,
    seq_len,
    hidden_size,
    dtype,
    atol,
    rtol,
    gate_multiplier,
    down_multiplier,
    file_name,
):
    torch.distributed.init_process_group(
        backend=infer_comm_backend(),
        init_method=f"file://{file_name}",
        rank=rank,
        world_size=world_size,
    )
    device = f"{infer_device()}:{rank}" if infer_device() != "cpu" else "cpu"
    device_mesh = torch.distributed.device_mesh.init_device_mesh(
        infer_device(), mesh_shape=(world_size,), mesh_dim_names=("tp",)
    )

    _a = torch.randn(bsz, seq_len, hidden_size, device=device, dtype=dtype)
    _b = torch.randn(bsz, seq_len, hidden_size, device=device, dtype=dtype)

    # Broadcast from rank 0 so all ranks operate on identical tensors
    torch.distributed.broadcast(_a, src=0)
    torch.distributed.broadcast(_b, src=0)

    assert hidden_size % world_size == 0, f"hidden_size ({hidden_size}) must be divisible by world_size ({world_size})"

    # DTensor path: shard inputs along the hidden dim
    a1 = _a.clone().detach().requires_grad_(True)
    b1 = _b.clone().detach().requires_grad_(True)
    da = torch.distributed.tensor.distribute_tensor(
        a1, device_mesh=device_mesh, placements=[torch.distributed.tensor.Shard(2)]
    )
    db = torch.distributed.tensor.distribute_tensor(
        b1, device_mesh=device_mesh, placements=[torch.distributed.tensor.Shard(2)]
    )

    # Regular tensor path
    a2 = _a.clone().detach().requires_grad_(True)
    b2 = _b.clone().detach().requires_grad_(True)

    c1 = LigerSiLUMulFunction.apply(da, db, gate_multiplier, down_multiplier)
    c2 = LigerSiLUMulFunction.apply(a2, b2, gate_multiplier, down_multiplier)

    torch.testing.assert_close(c1.full_tensor(), c2, atol=atol, rtol=rtol)

    grad = torch.randn_like(c2)
    torch.distributed.broadcast(grad, src=0)
    dgrad = torch.distributed.tensor.distribute_tensor(
        grad, device_mesh=device_mesh, placements=[torch.distributed.tensor.Shard(2)]
    )

    c1.backward(dgrad)
    c2.backward(grad)

    torch.testing.assert_close(da.grad.full_tensor(), a2.grad, atol=atol, rtol=rtol)
    torch.testing.assert_close(db.grad.full_tensor(), b2.grad, atol=atol, rtol=rtol)


@pytest.mark.parametrize(
    "world_size, bsz, seq_len, hidden_size",
    [
        (4, 2, 2, 8),
        (8, 9, 7, 64),
    ],
)
@pytest.mark.parametrize(
    "dtype, atol, rtol",
    [
        (torch.float32, 1e-4, 1e-6),
        (torch.bfloat16, 2e-1, 2e-2),
    ],
)
@pytest.mark.parametrize("gate_multiplier, down_multiplier", [(1.0, 1.0), (0.7, 1.3)])
def test_dtensor_liger_silumul(
    world_size, bsz, seq_len, hidden_size, dtype, atol, rtol, gate_multiplier, down_multiplier
):
    device_type = infer_device()
    device_module = getattr(torch, device_type, None)
    device_count = device_module.device_count() if hasattr(device_module, "device_count") else 0
    if device_count < world_size:
        pytest.xfail(f"Requires {world_size} {device_type.upper()} devices, but only {device_count} are available.")

    with tempfile.NamedTemporaryFile() as f:
        mp.spawn(
            _test_dtensor_liger_silumul,
            args=(
                world_size,
                bsz,
                seq_len,
                hidden_size,
                dtype,
                atol,
                rtol,
                gate_multiplier,
                down_multiplier,
                f.name,
            ),
            nprocs=world_size,
            join=True,
        )


@pytest.mark.skipif(not torch.cuda.is_available(), reason="Blackwell tiled SwiGLU path is CUDA-only")
@pytest.mark.parametrize(
    "n_rows, n_cols",
    [
        (4, 11009),  # wide + NOT a multiple of the 1024 tile -> exercises the column mask
        (3, 14337),  # ragged final tile, odd row count
        (4, 16384),  # exactly tile-aligned (16 full tiles)
    ],
)
@pytest.mark.parametrize("gate_multiplier", [1.0, 1.3])
@pytest.mark.parametrize(
    "dtype",
    [
        torch.float32,
        pytest.param(
            torch.bfloat16,
            marks=pytest.mark.skipif(not supports_bfloat16(), reason="bfloat16 not supported on this GPU"),
        ),
    ],
)
def test_swiglu_blackwell_tiled_matches_original(monkeypatch, n_rows, n_cols, gate_multiplier, dtype):
    """The Blackwell column-tiled path must be bit-for-bit identical to the one-row kernel.

    The dispatch normally only fires on SM 10.x, so CI never reaches it. The tiled kernels
    use no Blackwell-only instructions, so we force the gate on and compare against the
    original kernel on any CUDA GPU. Widths land on a ragged final tile (n_cols % 1024 != 0)
    to cover the column mask.
    """
    torch.manual_seed(0)
    a = torch.randn(n_rows, n_cols, device=device, dtype=dtype)
    b = torch.randn(n_rows, n_cols, device=device, dtype=dtype)
    dc = torch.randn(n_rows, n_cols, device=device, dtype=dtype)

    def run():
        a_ = a.clone().detach().requires_grad_(True)
        b_ = b.clone().detach().requires_grad_(True)
        c = LigerSiLUMulFunction.apply(a_, b_, gate_multiplier)
        c.backward(dc)
        return c.detach(), a_.grad.detach(), b_.grad.detach()

    # Original one-row kernel (gate forced off).
    monkeypatch.setattr(swiglu_ops, "infer_device_arch", lambda: "hopper")
    assert not swiglu_ops._should_tile(n_cols)
    c_ref, da_ref, db_ref = run()

    # Forced Blackwell tiled kernel (gate forced on).
    monkeypatch.setattr(swiglu_ops, "infer_device_arch", lambda: "blackwell")
    assert swiglu_ops._should_tile(n_cols)
    c_tiled, da_tiled, db_tiled = run()

    # Launch-geometry change only -> require exact equality (the PR's headline claim).
    torch.testing.assert_close(c_tiled, c_ref, rtol=0, atol=0)
    torch.testing.assert_close(da_tiled, da_ref, rtol=0, atol=0)
    torch.testing.assert_close(db_tiled, db_ref, rtol=0, atol=0)


# ---------------------------------------------------------------------------
# Harmonized plain SiLU-Mul tests (parallel across triton/cutedsl/cutile).
# These exercise the Triton LigerSiLUMulFunction directly against the fp32
# reference; the cutedsl/cutile suites mirror them (and additionally compare
# against this Triton Function via ``test_<be>_matches_triton``).
# ---------------------------------------------------------------------------
@pytest.mark.flaky(reruns=3, reruns_delay=2)
@pytest.mark.parametrize("shape", _SILU_SHAPES)
@pytest.mark.parametrize("dtype", _SILU_DTYPES)
def test_triton_silumul_correctness(shape, dtype):
    """Forward + backward vs the fp32 SiLU-Mul reference across shapes/dtypes."""
    torch.manual_seed(0)
    atol, rtol = _silu_tol(dtype)
    a = torch.randn(*shape, device=device, dtype=dtype)
    b = torch.randn(*shape, device=device, dtype=dtype)
    grad = torch.randn(*shape, device=device, dtype=dtype)

    ref_a = a.float().detach().requires_grad_(True)
    ref_b = b.float().detach().requires_grad_(True)
    ref_out = _ref_silu_mul(ref_a, ref_b)
    ref_out.backward(grad.float())

    test_a = a.clone().detach().requires_grad_(True)
    test_b = b.clone().detach().requires_grad_(True)
    test_out = LigerSiLUMulFunction.apply(test_a, test_b)
    test_out.backward(grad.clone())

    torch.testing.assert_close(test_out.float(), ref_out.float(), atol=atol, rtol=rtol)
    torch.testing.assert_close(test_a.grad.float(), ref_a.grad.float(), atol=atol, rtol=rtol)
    torch.testing.assert_close(test_b.grad.float(), ref_b.grad.float(), atol=atol, rtol=rtol)


@pytest.mark.parametrize("shape", [(512,), (123,), (2, 256, 512), (3, 5, 7), (2, 4, 8, 16)])
def test_triton_silumul_shape_flexibility(shape):
    """Arbitrary >=1-D shapes: forward + backward vs the fp32 reference; shape preserved."""
    torch.manual_seed(0)
    dtype = torch.float32
    atol, rtol = _silu_tol(dtype)
    a = torch.randn(*shape, device=device, dtype=dtype)
    b = torch.randn(*shape, device=device, dtype=dtype)
    grad = torch.randn(*shape, device=device, dtype=dtype)

    ref_a = a.clone().detach().requires_grad_(True)
    ref_b = b.clone().detach().requires_grad_(True)
    ref_out = _ref_silu_mul(ref_a, ref_b)
    ref_out.backward(grad)

    test_a = a.clone().detach().requires_grad_(True)
    test_b = b.clone().detach().requires_grad_(True)
    test_out = LigerSiLUMulFunction.apply(test_a, test_b)
    assert test_out.shape == a.shape
    test_out.backward(grad.clone())

    torch.testing.assert_close(test_out.float(), ref_out.float(), atol=atol, rtol=rtol)
    torch.testing.assert_close(test_a.grad.float(), ref_a.grad.float(), atol=atol, rtol=rtol)
    torch.testing.assert_close(test_b.grad.float(), ref_b.grad.float(), atol=atol, rtol=rtol)


@pytest.mark.parametrize("shape, dtype", [((73, 13824), torch.bfloat16), ((4, 2048), torch.float32)])
def test_triton_silumul_cuda_graph(shape, dtype):
    """CUDA-graph capture + replay of the forward must be bit-identical to eager."""
    if dtype == torch.bfloat16 and not supports_bfloat16():
        pytest.skip("bfloat16 not supported on this GPU")
    if not torch.cuda.is_available():
        pytest.skip("CUDA graph capture requires a CUDA device")
    torch.manual_seed(0)
    static_a = torch.randn(*shape, device=device, dtype=dtype)
    static_b = torch.randn(*shape, device=device, dtype=dtype)

    def run():
        with torch.no_grad():
            return LigerSiLUMulFunction.apply(static_a, static_b)

    # Warm up (compile / autotune) on a side stream before capture.
    s = torch.cuda.Stream()
    s.wait_stream(torch.cuda.current_stream())
    with torch.cuda.stream(s):
        for _ in range(3):
            run()
    torch.cuda.current_stream().wait_stream(s)

    g = torch.cuda.CUDAGraph()
    with torch.cuda.graph(g):
        static_out = run()

    # Replay on fresh data copied into the static input buffers.
    fresh_a = torch.randn(*shape, device=device, dtype=dtype)
    fresh_b = torch.randn(*shape, device=device, dtype=dtype)
    static_a.copy_(fresh_a)
    static_b.copy_(fresh_b)
    g.replay()
    torch.cuda.synchronize()
    graph_out = static_out.clone()

    with torch.no_grad():
        eager_out = LigerSiLUMulFunction.apply(fresh_a, fresh_b)

    max_abs_diff = (graph_out.float() - eager_out.float()).abs().max().item()
    assert max_abs_diff == 0.0, f"graph vs eager mismatch: max_abs_diff={max_abs_diff}"
