# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.
import inspect
from dataclasses import dataclass

import pytest
import spmd_types as spmd
import torch
import torch.distributed.checkpoint as dcp
import torchtitan.config.transform.quantization as quantization_transform
from spmd_types import SpmdType

from torchtitan.components.data import (
    FirstFitPackingConfig,
    GrainDataLoader,
    SingleDatasetConfig,
)
from torchtitan.components.data.sources import HuggingFaceRandomAccessSource
from torchtitan.config import ConfigLoader
from torchtitan.config.transform import (
    MXFP8LinearConverter,
    NVFP4GroupedLinearConverter,
    NVFP4LinearConverter,
)
from torchtitan.models.common.activation import Sigmoid
from torchtitan.models.common.attention import QKVLinear
from torchtitan.models.common.config_utils import make_router_config
from torchtitan.models.common.decoder_sharding import (
    colwise_config,
    dense_sequence_parallel_placement,
    rowwise_config,
)
from torchtitan.models.common.feed_forward import FeedForward
from torchtitan.models.common.hi_mid_lo_linear import HiMidLoLinear
from torchtitan.models.common.linear import (
    ColumnParallelLinear,
    GroupedLinear,
    Linear,
    RowParallelLinear,
    SharedExpertRowParallelLinear,
)
from torchtitan.models.common.moe import RoutedExperts
from torchtitan.models.common.token_dispatcher import AllToAllTokenDispatcher
from torchtitan.models.common.vision_encoder import InvariantRowParallelLinear
from torchtitan.models.gpt_oss.moe import GptOssGroupedLinear
from torchtitan.quantization import MXFP8Linear, NVFP4Linear
from torchtitan.quantization.mxfp8.experts import _get_mxfp8_grouped_linear_cls
from torchtitan.quantization.nvfp4.experts import _get_nvfp4_grouped_linear_cls
from torchtitan.quantization.utils import get_quantized_linear, has_quantization


class _ScaledLinear(Linear):
    @dataclass(kw_only=True, slots=True)
    class Config(Linear.Config):
        scale: float = 2.0

    def __init__(self, config: Config):
        super().__init__(config)
        self.scale = config.scale

    def _linear(
        self,
        input: torch.Tensor,
        weight: torch.Tensor,
        bias: torch.Tensor | None,
    ) -> torch.Tensor:
        return self.scale * super()._linear(input, weight, bias)


def test_no_quantization_by_default():
    config_loader = ConfigLoader()
    config = config_loader.load(
        [
            "--module",
            "torchtitan_recipes.tests.models.llama3",
            "--config",
            "llama3_debugmodel",
        ]
    )
    model_config = config.model
    assert not has_quantization(model_config)


def _router_config_for_quantization(dim: int):
    return make_router_config(
        dim=dim,
        num_experts=dim,
        score_func=Sigmoid.Config(),
        gate_param_init={"weight": torch.nn.init.zeros_},
    )


@pytest.mark.parametrize(
    "parallel_cls", [InvariantRowParallelLinear, SharedExpertRowParallelLinear]
)
def test_quantization_preserves_specialized_row_parallel_linear(parallel_cls):
    config_cls = get_quantized_linear(_ScaledLinear, parallel_cls).Config
    converted = config_cls(in_features=16, out_features=16, bias=True, scale=3.0)

    assert converted._owner is not None
    assert issubclass(converted._owner, parallel_cls)
    assert issubclass(converted._owner, _ScaledLinear)

    linear = converted.build()
    input = torch.randn(2, 16)
    expected = 3.0 * torch.nn.functional.linear(input, linear.weight, linear.bias)
    torch.testing.assert_close(linear(input), expected)


@pytest.mark.parametrize("config_cls", [HiMidLoLinear.Config])
def test_quantization_rejects_unsupported_linear_wrapper(config_cls):
    config = config_cls(in_features=16, out_features=16)

    with pytest.raises(ValueError, match=f"does not support {config._owner.__name__}"):
        quantization_transform._validate_quantizable_linear(config, "projection")


@pytest.mark.parametrize("parallel_cls", [ColumnParallelLinear, RowParallelLinear])
def test_get_quantized_linear_preserves_compute_and_tp_role(parallel_cls):
    quantized_cls = get_quantized_linear(_ScaledLinear, parallel_cls)
    config = quantized_cls.Config(
        in_features=4, out_features=2, num_linears=2, scale=3.0
    )
    linear = config.build()

    assert quantized_cls is get_quantized_linear(_ScaledLinear, parallel_cls)
    assert issubclass(quantized_cls, parallel_cls)
    assert issubclass(quantized_cls, _ScaledLinear)
    assert issubclass(quantized_cls.Config, _ScaledLinear.Config)

    input = torch.randn(3, 4)
    expected = 3.0 * torch.nn.functional.linear(
        input, linear.weight.flatten(0, -2), linear.bias
    ).unflatten(-1, linear.weight.shape[:-1])
    torch.testing.assert_close(linear(input), expected)


def test_mxfp8_converter_rejects_router_gate(monkeypatch):
    pytest.importorskip("torchao")
    if MXFP8Linear is None:
        pytest.skip("torchao MXFP8Linear is unavailable")
    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)
    converter = MXFP8LinearConverter.Config().build()
    with pytest.raises(ValueError, match="does not support HiMidLoLinear"):
        converter.convert(_router_config_for_quantization(128))


@pytest.mark.parametrize(
    ("config_cls", "parallel_cls"),
    [
        (ColumnParallelLinear.Config, ColumnParallelLinear),
        (RowParallelLinear.Config, RowParallelLinear),
    ],
)
def test_mxfp8_converter_preserves_tensor_parallel_role(
    monkeypatch, config_cls, parallel_cls
):
    if MXFP8Linear is None:
        pytest.skip("torchao MXFP8Linear is unavailable")
    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)
    converter = MXFP8LinearConverter.Config().build()
    converted = converter.convert(config_cls(in_features=128, out_features=128))

    assert converted._owner is not None
    assert issubclass(converted._owner, MXFP8Linear)
    assert issubclass(converted._owner, parallel_cls)


def test_nvfp4_converter_rejects_router_gate(monkeypatch):
    pytest.importorskip("torchao")
    if NVFP4Linear is None:
        pytest.skip("torchao NVFP4 training prototype not available")
    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)
    converter = NVFP4LinearConverter.Config().build()
    with pytest.raises(ValueError, match="does not support HiMidLoLinear"):
        converter.convert(_router_config_for_quantization(128))


@pytest.mark.parametrize(
    ("config_cls", "parallel_cls"),
    [
        (ColumnParallelLinear.Config, ColumnParallelLinear),
        (RowParallelLinear.Config, RowParallelLinear),
    ],
)
def test_nvfp4_converter_preserves_tensor_parallel_role(
    monkeypatch, config_cls, parallel_cls
):
    if NVFP4Linear is None:
        pytest.skip("torchao NVFP4Linear is unavailable")
    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)
    converter = NVFP4LinearConverter.Config().build()
    converted = converter.convert(config_cls(in_features=128, out_features=128))

    assert converted._owner is not None
    assert issubclass(converted._owner, NVFP4Linear)
    assert issubclass(converted._owner, parallel_cls)


@pytest.mark.parametrize(
    "module, recipe, expected_num_layers",
    [
        (
            "torchtitan_recipes.tests.models.llama3",
            "llama3_debugmodel_nvfp4",
            6,
        ),
        (
            "torchtitan_recipes.tests.models.qwen3",
            "qwen3_debugmodel_nvfp4",
            8,
        ),
    ],
)
def test_nvfp4_converter_targets_layers_not_lm_head(
    monkeypatch, module, recipe, expected_num_layers
):
    pytest.importorskip("torchao")
    from torchtitan.quantization import NVFP4Linear

    if NVFP4Linear is None:
        pytest.skip("torchao NVFP4 training prototype not available")
    # Exercise convert() targeting independent of GPU: bypass the sm100 gate
    # that NVFP4LinearConverter.__init__ enforces (hardware is irrelevant to the
    # config-tree transform under test).
    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)

    config_loader = ConfigLoader()
    config = config_loader.load(["--module", module, "--config", recipe])
    model_config = config.model
    assert has_quantization(model_config)

    converted, stock = [], []
    for fqn, lc, _parent, _attr in model_config.traverse(Linear.Config):
        (converted if isinstance(lc, NVFP4Linear.Config) else stock).append(fqn)

    # Every in-layer linear is swapped; the lm_head stays stock (NVFP4 requires
    # each GEMM dim divisible by 128, which the vocab projection violates).
    assert converted and all("layers" in fqn for fqn in converted)
    assert {int(fqn.split(".")[1]) for fqn in converted} == set(
        range(expected_num_layers)
    )
    assert stock == ["lm_head"]


def test_nvfp4_bf16_tail_fqns():
    from torchtitan.quantization.nvfp4 import nvfp4_bf16_tail_fqns

    # 32 layers, 15% tail -> ceil(4.8)=5 bf16, convert layers 0..26.
    fqns = nvfp4_bf16_tail_fqns(32, 0.15)
    assert fqns == [f"layers.{i}." for i in range(27)]
    # Every fqn is trailing-dot anchored so "layers.2." matches layer 2 only,
    # not "layers.20".."layers.29" (the converter substring-matches).
    assert all(f.startswith("layers.") and f.endswith(".") for f in fqns)
    # Fraction 0 keeps nothing in bf16 -> every layer converted.
    assert nvfp4_bf16_tail_fqns(4, 0.0) == [
        "layers.0.",
        "layers.1.",
        "layers.2.",
        "layers.3.",
    ]
    # A fraction that rounds up to all layers leaves nothing to convert -> raise
    # (an empty fqns list would instead convert *all* Linears).
    with pytest.raises(ValueError, match="nothing to convert"):
        nvfp4_bf16_tail_fqns(4, 1.0)


@pytest.mark.parametrize(
    "module, recipe, expected_cutoff",
    [
        (
            "torchtitan_recipes.tests.models.llama3",
            "llama3_debugmodel_first_85_pct_layers_nvfp4",
            5,
        ),
        (
            "torchtitan_recipes.tests.models.llama3",
            "llama3_8b_first_85_pct_layers_nvfp4",
            27,
        ),
        (
            "torchtitan_recipes.tests.models.qwen3",
            "qwen3_debugmodel_first_85_pct_layers_nvfp4",
            6,
        ),
        (
            "torchtitan_recipes.tests.models.qwen3",
            "qwen3_8b_first_85_pct_layers_nvfp4",
            30,
        ),
    ],
)
def test_nvfp4_first_85_pct_layers_converts_only_leading_layers(
    monkeypatch, module, recipe, expected_cutoff
):
    pytest.importorskip("torchao")
    from torchtitan.quantization import NVFP4Linear

    if NVFP4Linear is None:
        pytest.skip("torchao NVFP4 training prototype not available")
    import math

    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)

    config = ConfigLoader().load(["--module", module, "--config", recipe])
    model_config = config.model
    n_layers = len(model_config.layers)
    cutoff = n_layers - math.ceil(n_layers * 0.15)
    assert cutoff == expected_cutoff
    assert 0 < cutoff < n_layers  # a real split: some NVFP4, some bf16

    converted_layers, stock = set(), []
    for fqn, lc, _parent, _attr in model_config.traverse(Linear.Config):
        if isinstance(lc, NVFP4Linear.Config):
            converted_layers.add(int(fqn.split(".")[1]))
        else:
            stock.append(fqn)

    # Only the leading layers are NVFP4; the bf16 tail + lm_head stay stock.
    assert converted_layers == set(range(cutoff))
    assert "lm_head" in stock
    assert all(
        not fqn.startswith("layers.") or int(fqn.split(".")[1]) >= cutoff
        for fqn in stock
    )


def _nvfp4_linear_cls():
    pytest.importorskip("torchao")
    from torchtitan.quantization import NVFP4Linear

    if NVFP4Linear is None:
        pytest.skip("torchao NVFP4 training prototype not available")
    return NVFP4Linear


@pytest.mark.parametrize("in_features, out_features", [(512, 300), (300, 512)])
def test_nvfp4_config_rejects_non_128_dims(in_features, out_features):
    # The model dims are known at config-build time, so a non-128 in/out_features
    # (e.g. the LM head) is rejected in Config.__post_init__ before any TP.
    NVFP4Linear = _nvfp4_linear_cls()
    with pytest.raises(ValueError, match="divisible by 128"):
        NVFP4Linear.Config(in_features=in_features, out_features=out_features)


@pytest.mark.parametrize(
    "sharding_config_factory, input_tp",
    [
        pytest.param(
            lambda: colwise_config(input_layout=dense_sequence_parallel_placement()),
            spmd.R,
            id="colwise",
        ),
        pytest.param(
            lambda: rowwise_config(output_layout=dense_sequence_parallel_placement()),
            spmd.S(-1),
            id="rowwise",
        ),
    ],
)
def test_nvfp4_build_configures_local_spmd_sharding(sharding_config_factory, input_tp):
    # Config.build() folds the stock colwise/rowwise sharding into the local
    # SPMD region for the opaque NVFP4 GEMM.
    NVFP4Linear = _nvfp4_linear_cls()
    from torchtitan.distributed.parallelism_context import MeshAxisName
    from torchtitan.models.common.decoder_sharding import dense_activation_placement

    module = NVFP4Linear.Config(
        in_features=512,
        out_features=1024,
        sharding_config=sharding_config_factory(),
    ).build()
    sc = module._sharding_config
    assert sc.local_spmd
    input_layout = dense_activation_placement(tp=input_tp, cp=spmd.S(0))
    assert sc.in_src_shardings == {"input": input_layout}
    assert sc.in_dst_shardings == {"input": input_layout}
    assert list(inspect.signature(module.forward).parameters) == ["input"]
    assert "weight" in sc.state_shardings
    assert sc.state_shardings["_sr_seed"] == SpmdType(
        {
            MeshAxisName.DP: spmd.V,
            MeshAxisName.CP: spmd.V,
            MeshAxisName.TP: spmd.V,
        }
    )


@pytest.mark.parametrize("parallel_cls", [ColumnParallelLinear, RowParallelLinear])
def test_nvfp4_parallel_build_preserves_collective_boundary(parallel_cls):
    if NVFP4Linear is None:
        pytest.skip("torchao NVFP4 training prototype not available")

    boundary_layout = dense_sequence_parallel_placement()
    linear_cls = get_quantized_linear(NVFP4Linear, parallel_cls)
    sharding_config = (
        colwise_config(input_layout=boundary_layout)
        if parallel_cls is ColumnParallelLinear
        else rowwise_config(output_layout=boundary_layout)
    )
    module = linear_cls.Config(
        in_features=512,
        out_features=1024,
        sharding_config=sharding_config,
    ).build()

    assert not module._sharding_config.local_spmd
    assert module._sharding_config.in_src_shardings == (
        sharding_config.in_src_shardings
    )
    assert module._sharding_config.out_src_shardings == (
        sharding_config.out_src_shardings
    )
    assert "_sr_seed" in module._sharding_config.state_shardings


@pytest.mark.parametrize(
    "module, recipe",
    [
        (
            "torchtitan_recipes.tests.models.llama3",
            "llama3_debugmodel_nvfp4",
        ),
        (
            "torchtitan_recipes.tests.models.llama3",
            "llama3_debugmodel_first_85_pct_layers_nvfp4",
        ),
        (
            "torchtitan_recipes.tests.models.llama3",
            "llama3_8b_first_85_pct_layers_nvfp4",
        ),
        (
            "torchtitan_recipes.tests.models.qwen3",
            "qwen3_debugmodel_nvfp4",
        ),
        (
            "torchtitan_recipes.tests.models.qwen3",
            "qwen3_debugmodel_first_85_pct_layers_nvfp4",
        ),
        (
            "torchtitan_recipes.tests.models.qwen3",
            "qwen3_8b_first_85_pct_layers_nvfp4",
        ),
    ],
)
def test_nvfp4_recipes_parse(monkeypatch, module, recipe):
    _nvfp4_linear_cls()
    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)
    base_args = ["--module", module, "--config", recipe]

    ConfigLoader().load(base_args)


@pytest.mark.parametrize(
    "recipe",
    [
        "qwen3_debugmodel_nvfp4",
        "qwen3_debugmodel_first_85_pct_layers_nvfp4",
        "qwen3_8b_first_85_pct_layers_nvfp4",
    ],
)
def test_qwen3_recipes_resolve(monkeypatch, recipe):
    _nvfp4_linear_cls()
    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)
    config = ConfigLoader().load(
        ["--module", "torchtitan_recipes.tests.models.qwen3", "--config", recipe]
    )
    assert type(config.model).__qualname__ == "Qwen3Model.Config"
    if recipe == "qwen3_8b_first_85_pct_layers_nvfp4":
        assert isinstance(config.dataloader, GrainDataLoader.Config)
        packed_dataset = config.dataloader.dataset
        assert isinstance(packed_dataset, FirstFitPackingConfig)
        dataset = packed_dataset.dataset
        assert isinstance(dataset, SingleDatasetConfig)
        assert isinstance(dataset.source, HuggingFaceRandomAccessSource.Config)
        assert dataset.source.path == "openai/gsm8k"
        assert config.checkpointer.initial_load_in_hf
        assert config.model.local_compile_regions == [
            "loss",
            "fused_binary_activation",
            "cos_sin_rope",
            "fp32_to_bf16_split",
        ]


def test_nvfp4_module_buffers_and_native_checkpoint():
    """Built module has the stock weight param plus the two NVFP4 runtime
    buffers, and both buffers are non-persistent -- the RHT vector is a fixed
    constant and the SR seed is per-rank -- so a native checkpoint carries only
    the stock weight."""
    NVFP4Linear = _nvfp4_linear_cls()
    from torchtitan.quantization.nvfp4 import _HARDCODED_SIGN_VECTOR

    module = NVFP4Linear.Config(in_features=512, out_features=1024).build()
    assert {name for name, _ in module.named_parameters()} == {"weight"}
    module.init_states()
    buffers = dict(module.named_buffers())
    assert set(buffers) == {"_sr_seed", "_rht_sign_vector"}
    assert buffers["_sr_seed"].dtype == torch.int64
    assert tuple(buffers["_rht_sign_vector"].shape) == (16,)
    # The RHT vector is the fixed v1-recipe constant, identical on every rank.
    assert tuple(int(v) for v in buffers["_rht_sign_vector"]) == _HARDCODED_SIGN_VECTOR
    # Both runtime buffers are non-persistent, so a native checkpoint carries
    # only the stock weight.
    assert set(module.state_dict()) == {"weight"}


def test_nvfp4_stock_checkpoint_loads_before_init_states():
    """A stock bf16 checkpoint (no NVFP4 buffers) loads; buffers stay unmaterialized
    until init_states creates them."""
    NVFP4Linear = _nvfp4_linear_cls()
    stock = Linear.Config(in_features=512, out_features=1024).build()
    nvfp4 = NVFP4Linear.Config(in_features=512, out_features=1024).build()

    nvfp4.load_state_dict(stock.state_dict(), strict=False)
    assert nvfp4._rht_sign_vector is None
    assert nvfp4._rht_sign_vector_tuple is None

    nvfp4.init_states()
    assert nvfp4._rht_sign_vector is not None
    assert nvfp4._rht_sign_vector_tuple is not None


def test_nvfp4_hf_export_strips_buffers(monkeypatch):
    """The HF export boundary contains only stock keys -- no NVFP4 runtime buffers."""
    NVFP4Linear = _nvfp4_linear_cls()
    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)
    from torchtitan.models.llama3.state_dict_adapter import Llama3StateDictAdapter

    config = ConfigLoader().load(
        [
            "--module",
            "torchtitan_recipes.tests.models.llama3",
            "--config",
            "llama3_debugmodel_nvfp4",
        ]
    )
    model_config = config.model
    model = model_config.build()
    model.init_states()
    assert isinstance(model.get_submodule("layers.0.feed_forward.w13"), NVFP4Linear)

    sd = model.state_dict()
    # Both NVFP4 runtime buffers are non-persistent, so neither the RHT vector
    # nor the per-rank SR seed appears in the native state dict.
    assert not any("_rht_sign_vector" in k for k in sd)
    assert not any("_sr_seed" in k for k in sd)

    hf_sd = Llama3StateDictAdapter(model_config, hf_assets_path=None).to_hf(sd)
    assert "model.layers.0.mlp.gate_proj.weight" in hf_sd
    assert not any("_rht_sign_vector" in k for k in hf_sd)


def test_quantized_grouped_linear():
    """Quantized grouped linears preserve base and specialized module types."""
    MXFP8GroupedLinear = _get_mxfp8_grouped_linear_cls(GroupedLinear)

    assert MXFP8GroupedLinear.Config._owner is MXFP8GroupedLinear

    mxfp8_cls = _get_mxfp8_grouped_linear_cls(GptOssGroupedLinear)

    assert mxfp8_cls.Config._owner is mxfp8_cls
    assert issubclass(mxfp8_cls, GptOssGroupedLinear)


@pytest.mark.parametrize("parent_cls", [GroupedLinear, GptOssGroupedLinear])
@pytest.mark.parametrize(
    "make_quantized_cls",
    [_get_mxfp8_grouped_linear_cls, _get_nvfp4_grouped_linear_cls],
    ids=["mxfp8", "nvfp4"],
)
def test_grouped_mm_overrides_keep_the_seam_signature(make_quantized_cls, parent_cls):
    """Every ``_grouped_mm`` override must accept the base class's keywords.

    ``MoE.forward`` calls the seam by keyword, so an override whose parameter
    names drift raises TypeError at the first expert GEMM rather than at import
    time -- and only in a MoE training run, which no other unit test reaches.
    """
    base = inspect.signature(parent_cls._grouped_mm)
    override = inspect.signature(make_quantized_cls(parent_cls)._grouped_mm)

    assert list(override.parameters) == list(base.parameters)
    for name, parameter in base.parameters.items():
        assert override.parameters[name].kind == parameter.kind


def test_mxfp8_grouped_linear_flattens_structured_w13(monkeypatch):
    """MXFP8 receives flattened W13 and returns its structured output."""
    from torchao.prototype.moe_training import utils as moe_training_utils

    captured = {}

    def grouped_mm(input_RI, weight_EIO, *, config, offs):
        del config, offs
        captured["weight_shape"] = weight_EIO.shape
        return input_RI.new_zeros(input_RI.shape[0], weight_EIO.shape[-1])

    monkeypatch.setattr(
        moe_training_utils,
        "_quantize_then_scaled_grouped_mm",
        grouped_mm,
    )
    grouped_linear_cls = _get_mxfp8_grouped_linear_cls(GroupedLinear)
    grouped_linear = grouped_linear_cls.Config(
        group_size=4,
        in_features=128,
        out_features=64,
        num_linears=2,
    ).build()

    output_R2O = grouped_linear(
        torch.zeros(8, 128),
        torch.tensor([2, 4, 6, 8], dtype=torch.int32),
    )

    assert captured["weight_shape"] == torch.Size([4, 128, 128])
    assert output_R2O.shape == torch.Size([8, 2, 64])


@pytest.mark.filterwarnings("ignore:torch.distributed is disabled")
def test_mxfp8_linear_dcp_round_trip_needs_no_safe_globals(tmp_path):
    pytest.importorskip("torchao")
    if MXFP8Linear is None:
        pytest.skip("torchao MXFP8Linear is unavailable")

    config = MXFP8Linear.Config(
        in_features=128,
        out_features=128,
        bias=False,
    )
    source = config.build()
    target = config.build()

    with torch.no_grad():
        source.weight._tensor.copy_(
            torch.arange(source.weight.numel()).reshape(source.weight.shape)
        )
        target.weight._tensor.zero_()

    # DCP reads with torch.load(weights_only=True). Clearing the safe globals
    # makes the load fail if the wrapper subclass was pickled into the shard.
    saved_safe_globals = torch.serialization.get_safe_globals()
    try:
        torch.serialization.clear_safe_globals()
        dcp.save(source.state_dict(), checkpoint_id=tmp_path, no_dist=True)
        dcp.load(target.state_dict(), checkpoint_id=tmp_path, no_dist=True)
    finally:
        torch.serialization.clear_safe_globals()
        torch.serialization.add_safe_globals(saved_safe_globals)

    assert torch.equal(
        target.weight._tensor.view(torch.uint8),
        source.weight._tensor.view(torch.uint8),
    )


def test_mxfp8_linear_validates_config_and_installs_weight_wrapper():
    pytest.importorskip("torchao")
    if MXFP8Linear is None:
        pytest.skip("torchao MXFP8Linear is unavailable")
    from torchtitan.quantization._fsdp_tensor import _UnshardedFSDPTensor
    from torchtitan.quantization.mxfp8.tensor import (
        _LinearShardedTensorWithMXFP8Compute,
    )

    with pytest.raises(ValueError, match="in_features divisible by 32"):
        MXFP8Linear.Config(in_features=127, out_features=128)
    with pytest.raises(ValueError, match="out_features divisible by 32"):
        MXFP8Linear.Config(in_features=128, out_features=127)
    with pytest.raises(
        ValueError,
        match="input_activation_format_for_backward must be one of",
    ):
        MXFP8Linear.Config(
            in_features=128,
            out_features=128,
            input_activation_format_for_backward="missing",
        )
    with pytest.raises(ValueError, match="out_features divisible by 32"):
        MXFP8Linear.Config(
            in_features=128,
            out_features=127,
            num_linears=2,
        )

    local_stacked_weight = _LinearShardedTensorWithMXFP8Compute(
        torch.empty(3, 16, 128, dtype=torch.bfloat16)
    )
    with pytest.raises(ValueError, match="local matrix out_features divisible by 32"):
        local_stacked_weight._build_operands(local_stacked_weight._tensor)

    for sharding_config in (
        colwise_config(input_layout=dense_sequence_parallel_placement()),
        rowwise_config(output_layout=dense_sequence_parallel_placement()),
    ):
        linear = MXFP8Linear.Config(
            in_features=128,
            out_features=128,
            bias=False,
            sharding_config=sharding_config,
        ).build()
        assert linear._sharding_config is not None
        # The wrapper is installed at construction, so no caller has to opt
        # in. Until a data parallel implementation drives its lifecycle it is
        # the sharded state, which holds the BF16 weight; the unsharded tensor
        # is a separate type the post-all-gather hook produces.
        assert isinstance(linear.weight, _LinearShardedTensorWithMXFP8Compute)
        assert not isinstance(linear.weight, _UnshardedFSDPTensor)


def test_mxfp8_converter_replaces_a_root_linear_config(monkeypatch):
    """A Linear.Config with no parent is returned, not mutated in place.

    ``convert`` writes into ``parent`` for nested configs, so the root case is
    the one branch that has to return the replacement. Not covered by the FQN
    test below, which passes a FeedForward and so always has a parent.
    """
    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)
    converter = MXFP8LinearConverter.Config().build()

    converted = converter.convert(
        Linear.Config(in_features=128, out_features=128, bias=False)
    )

    assert isinstance(converted, MXFP8Linear.Config)
    assert converted.input_activation_format_for_backward == "bf16"


def test_mxfp8_converter_rejects_unaligned_fused_qkv_head_dim(monkeypatch):
    if MXFP8Linear is None:
        pytest.skip("torchao MXFP8Linear is unavailable")
    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)
    converter = MXFP8LinearConverter.Config().build()
    head_dim = 48
    n_heads = 4
    n_kv_heads = 2
    qkv_config = QKVLinear.Config(
        head_dim=head_dim,
        n_heads=n_heads,
        n_kv_heads=n_kv_heads,
        wqkv=Linear.Config(
            in_features=128,
            out_features=(n_heads + 2 * n_kv_heads) * head_dim,
        ),
    )

    with pytest.raises(ValueError, match="head_dim divisible by 32"):
        converter.convert(qkv_config)


def test_mxfp8_converter_applies_mxfp8_saved_input_fqns(monkeypatch):
    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)
    converter = MXFP8LinearConverter.Config(
        linears_saving_inputs_for_backward_in_mxfp8=["w2"],
    ).build()
    converted = converter.convert(
        FeedForward.Config(
            w13=Linear.Config(in_features=128, out_features=128, num_linears=2),
            w2=Linear.Config(in_features=128, out_features=128),
        )
    )

    assert isinstance(converted.w13, MXFP8Linear.Config)
    assert isinstance(converted.w2, MXFP8Linear.Config)
    assert converted.w13.input_activation_format_for_backward == "bf16"
    assert converted.w2.input_activation_format_for_backward == "mxfp8"


def test_mxfp8_converter_rejects_unmatched_saved_input_fqns(monkeypatch):
    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)
    converter = MXFP8LinearConverter.Config(
        linears_saving_inputs_for_backward_in_mxfp8=["missing"],
    ).build()
    model_config = FeedForward.Config(
        w13=Linear.Config(in_features=128, out_features=128, num_linears=2),
        w2=Linear.Config(in_features=128, out_features=128),
    )

    with pytest.raises(
        ValueError,
        match="selectors did not match any converted Linear.Config",
    ):
        converter.convert(model_config)


def test_mxfp8_converter_rejects_empty_saved_input_fqn():
    with pytest.raises(ValueError, match="cannot contain an empty FQN selector"):
        MXFP8LinearConverter.Config(
            linears_saving_inputs_for_backward_in_mxfp8=[""],
        )


@pytest.mark.parametrize(
    "config_factory, mxfp8_fqns",
    [
        (
            "llama3",
            ("attention.qkv_linear.wqkv", "feed_forward.w2"),
        ),
        (
            "llama3_graph",
            ("attention.qkv_linear.wqkv", "feed_forward.w2"),
        ),
        (
            "deepseek_v3",
            ("attention.wkv_b", "feed_forward.w2", "shared_experts.w2"),
        ),
        (
            "deepseek_v3_graph",
            ("attention.wkv_b", "feed_forward.w2", "shared_experts.w2"),
        ),
    ],
)
def test_builtin_mxfp8_configs_assign_input_activation_format_for_backward(
    monkeypatch, config_factory, mxfp8_fqns
):
    if MXFP8Linear is None:
        pytest.skip("torchao MXFP8Linear is unavailable")
    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)
    if config_factory == "llama3":
        from torchtitan_recipes.tests.models.llama3 import (
            llama3_debugmodel_mxfp8 as build_config,
        )
    elif config_factory == "llama3_graph":
        from torchtitan_recipes.tests.graph_trainer.llama3 import (
            graph_trainer_llama3_debugmodel_mxfp8 as build_config,
        )
    elif config_factory == "deepseek_v3":
        from torchtitan_recipes.tests.models.deepseek_v3 import (
            deepseek_v3_debugmodel_mxfp8 as build_config,
        )
    else:
        from torchtitan_recipes.tests.graph_trainer.deepseek_v3 import (
            graph_trainer_deepseek_v3_debugmodel_mxfp8 as build_config,
        )

    trainer_config = build_config()
    model_config = trainer_config.model
    assignments = {
        fqn: config.input_activation_format_for_backward
        for fqn, config, _parent, _attr in model_config.traverse(MXFP8Linear.Config)
    }
    assert assignments
    assert "bf16" in assignments.values()
    assert "mxfp8" in assignments.values()
    for fqn, save_format in assignments.items():
        expected = (
            "mxfp8" if any(selector in fqn for selector in mxfp8_fqns) else "bf16"
        )
        assert save_format == expected, f"Unexpected policy for {fqn}"


def test_mxfp8_linear_loads_stock_checkpoint():
    pytest.importorskip("torchao")
    if MXFP8Linear is None:
        pytest.skip("torchao MXFP8Linear is unavailable")
    from torchtitan.quantization.mxfp8.tensor import (
        _LinearShardedTensorWithMXFP8Compute,
    )

    stock = Linear.Config(in_features=128, out_features=96).build()
    mxfp8 = MXFP8Linear.Config(in_features=128, out_features=96).build()
    with torch.no_grad():
        stock.weight.normal_()

    mxfp8.load_state_dict(stock.state_dict())
    assert isinstance(mxfp8.weight, _LinearShardedTensorWithMXFP8Compute)
    assert torch.equal(mxfp8.weight._tensor, stock.weight)


def test_nvfp4_grouped_converter_selects_both_projections_and_one_dispatcher_swap(
    monkeypatch,
):
    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)
    config = RoutedExperts.Config(
        w13=GroupedLinear.Config(
            group_size=2, in_features=128, out_features=128, num_linears=2
        ),
        w2=GroupedLinear.Config(group_size=2, in_features=128, out_features=128),
        token_dispatcher=AllToAllTokenDispatcher.Config(num_experts=2, top_k=1),
    )
    calls = []
    actual_swap = quantization_transform.swap_token_dispatcher

    def counted_swap(owner, pad_multiple):
        calls.append((owner, pad_multiple))
        actual_swap(owner, pad_multiple)

    monkeypatch.setattr(quantization_transform, "swap_token_dispatcher", counted_swap)
    converter = NVFP4GroupedLinearConverter(
        NVFP4GroupedLinearConverter.Config(fqns=["w"], pad_multiple=256)
    )
    converted = converter.convert(config)
    quantized_cls = _get_nvfp4_grouped_linear_cls(GroupedLinear)
    assert isinstance(converted.w13, quantized_cls.Config)
    assert isinstance(converted.w2, quantized_cls.Config)
    assert has_quantization(converted)
    assert calls == [(converted, 256)]
    assert converted.token_dispatcher.pad_multiple == 256


def test_nvfp4_grouped_converter_checks_both_widths_before_mutation(monkeypatch):
    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)
    config = RoutedExperts.Config(
        w13=GroupedLinear.Config(
            group_size=2, in_features=128, out_features=127, num_linears=2
        ),
        w2=GroupedLinear.Config(group_size=2, in_features=127, out_features=128),
        token_dispatcher=AllToAllTokenDispatcher.Config(num_experts=2, top_k=1),
    )
    converter = NVFP4GroupedLinearConverter(NVFP4GroupedLinearConverter.Config())
    with pytest.raises(ValueError, match="input and output widths"):
        converter.convert(config)
    assert type(config.w13) is GroupedLinear.Config
    assert type(config.w2) is GroupedLinear.Config
    assert type(config.token_dispatcher) is AllToAllTokenDispatcher.Config


def test_nvfp4_grouped_linear_forwards_flattened_w13_and_runtime_state(monkeypatch):
    from torchtitan.quantization.nvfp4 import experts as nvfp4_experts

    captured = {}

    def grouped_mm(input_RI, weight_EOI, sign_vector, sr_seed, **kwargs):
        captured.update(
            input_RI=input_RI,
            weight_EOI=weight_EOI,
            sign_vector=sign_vector,
            sr_seed=sr_seed,
            **kwargs,
        )
        return input_RI.new_zeros(input_RI.shape[0], weight_EOI.shape[1])

    monkeypatch.setattr(
        nvfp4_experts, "_to_nvfp4_rht_rs_then_scaled_grouped_mm", grouped_mm
    )
    module = (
        _get_nvfp4_grouped_linear_cls(GroupedLinear)
        .Config(
            group_size=2,
            in_features=128,
            out_features=128,
            num_linears=2,
            param_init={"weight": torch.nn.init.zeros_},
        )
        .build()
    )
    module.init_states()
    input_RI = torch.zeros(256, 128)
    offsets_E = torch.tensor([128, 256], dtype=torch.int32)
    output = module(input_RI, offsets_E)
    assert captured["weight_EOI"].shape == (2, 256, 128)
    assert captured["offs"] is offsets_E
    assert captured["pad_token_groups_for_grouped_mm"] is False
    assert captured["sign_vector"] == module.rht_sign_vector
    assert captured["sr_seed"] is module._sr_seed
    assert output.shape == (256, 2, 128)
    assert "_sr_seed" not in module.state_dict()
    assert "_rht_sign_vector" not in module.state_dict()


@pytest.mark.parametrize(
    "recipe",
    [
        "deepseek_v3_debugmodel_nvfp4_ffn_mxfp8_attn",
        "deepseek_v3_16b_nvfp4_ffn_mxfp8_attn",
        "deepseek_v3_671b_nvfp4_ffn_mxfp8_attn",
    ],
)
@pytest.mark.parametrize("bf16_tail_fraction", [0.0, 0.5])
def test_deepseek_nvfp4_recipes_preserve_quantization_and_routing(
    recipe, bf16_tail_fraction, monkeypatch
):
    from torchtitan_recipes.models import deepseek_v3 as model_recipes
    from torchtitan_recipes.tests.models import deepseek_v3 as test_recipes

    monkeypatch.setattr(quantization_transform, "has_cuda_capability", lambda *_: True)
    config_registry = model_recipes if "671b" in recipe else test_recipes

    config = getattr(config_registry, recipe)(bf16_tail_fraction=bf16_tail_fraction)
    num_nvfp4_layers = (
        len(config.model.layers)
        if bf16_tail_fraction == 0
        else len(config.model.layers) // 2
    )
    assert config.model.local_compile_regions == ["loss"]
    grouped = list(config.model.traverse(GroupedLinear.Config))
    assert grouped
    quantized_cls = _get_nvfp4_grouped_linear_cls(GroupedLinear)
    assert all(
        isinstance(projection, quantized_cls.Config)
        == (int(fqn.split(".")[1]) < num_nvfp4_layers)
        for fqn, projection, _, _ in grouped
    )
    assert all(fqn.endswith((".w13", ".w2")) for fqn, _, _, _ in grouped)
    from torchtitan.quantization.nvfp4 import NVFP4Linear

    linears = dict(
        (fqn, projection)
        for fqn, projection, _, _ in config.model.traverse(Linear.Config)
    )
    baseline_recipe = recipe.removesuffix("_nvfp4_ffn_mxfp8_attn")
    baseline = getattr(config_registry, baseline_recipe)()
    baseline_linears = {
        fqn: projection
        for fqn, projection, _, _ in baseline.model.traverse(Linear.Config)
    }
    assert all(
        type(projection) is type(baseline_linears[fqn])
        for fqn, projection in linears.items()
        if "router.gate" in fqn or fqn == "lm_head"
    )
    assert any(
        isinstance(projection, NVFP4Linear.Config) for projection in linears.values()
    )
    if recipe == "deepseek_v3_16b_nvfp4_ffn_mxfp8_attn":
        assert all(
            type(projection) is type(baseline_linears[fqn])
            for fqn, projection in linears.items()
            if ".feed_forward." in fqn
        )

    assert all(
        not isinstance(projection, NVFP4Linear.Config)
        for fqn, projection in linears.items()
        if fqn.startswith("layers.") and int(fqn.split(".")[1]) >= num_nvfp4_layers
    )
    for fqn, routed, _, _ in config.model.traverse(RoutedExperts.Config):
        if int(fqn.split(".")[1]) < num_nvfp4_layers:
            assert routed.token_dispatcher.pad_multiple == 128
        if recipe != "deepseek_v3_debugmodel_nvfp4_ffn_mxfp8_attn":
            assert routed.token_dispatcher.non_blocking_capacity_factor == (
                0.1875 if recipe == "deepseek_v3_16b_nvfp4_ffn_mxfp8_attn" else 0.03125
            )
