# 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 unittest
from dataclasses import dataclass

import spmd_types as spmd
import torch
import torch.nn as nn
from spmd_types import SpmdType

from torchtitan.distributed.parallelism_context import MeshAxisName, ParallelismContext
from torchtitan.models.common.linear import Linear
from torchtitan.protocols.module import Module, ModuleDict, ModuleList, Sequential
from torchtitan.protocols.sharding import ShardingConfig


class TestModuleInitStates(unittest.TestCase):
    """Tests for Module.init_states behavior."""

    def test_default_init_states_no_param_init_raises(self):
        """Subclass with own params but no param_init and no reset_parameters raises."""

        class SimpleModule(Module):
            def __init__(self):
                super().__init__()
                self.weight = nn.Parameter(torch.empty(4))

        m = SimpleModule()
        with self.assertRaises(ValueError):
            m.init_states()

    def test_init_states_falls_back_to_reset_parameters(self):
        """No param_init + reset_parameters available -> reset_parameters runs.

        Wrappers like ``Linear``/``LayerNorm``/``Conv2d`` inherit
        ``reset_parameters`` from their nn base; the base ``Module`` calls
        it when no custom ``param_init`` is provided.
        """
        m = Linear.Config(in_features=4, out_features=8, bias=True).build()
        nn.init.zeros_(m.weight)
        nn.init.ones_(m.bias)
        m.init_states()
        # reset_parameters uses kaiming_uniform; result is non-zero, non-one.
        self.assertFalse(torch.all(m.weight == 0))
        self.assertFalse(torch.all(m.bias == 1))

    def test_init_states_auto_recurses(self):
        """init_states automatically recurses into Module children."""

        class Child(Module):
            def __init__(self):
                super().__init__()
                self.weight = nn.Parameter(torch.empty(4))

        class Parent(Module):
            def __init__(self):
                super().__init__()
                self.child = Child()

        m = Parent()
        m.child._param_init = {"weight": nn.init.zeros_}
        m.init_states()
        self.assertTrue(torch.all(m.child.weight == 0))

    def test__init_self_buffers_called(self):
        """_init_self_buffers is called with kwargs."""

        class BufferModule(Module):
            def __init__(self):
                super().__init__()
                self.register_buffer("buf", torch.zeros(4))
                self.buffer_device_seen = None

            def _init_self_buffers(self, *, buffer_device=None):
                self.buffer_device_seen = buffer_device

        m = BufferModule()
        m.init_states(buffer_device=torch.device("cpu"))
        self.assertEqual(m.buffer_device_seen, torch.device("cpu"))

    def test_no_params_no_error(self):
        """Module with no parameters and no param_init doesn't raise."""

        class NoParams(Module):
            def __init__(self):
                super().__init__()

        m = NoParams()
        m.init_states()  # should not raise


class TestDiamondInheritance(unittest.TestCase):
    """Tests for diamond inheritance: class Foo(nn.SomeModule, Module)."""

    class TestEmbedding(nn.Embedding, Module):
        def __init__(self, num_embeddings, embedding_dim):
            super().__init__(num_embeddings, embedding_dim)

    def test_module_has_no_init(self):
        """Module must not define __init__ in its own __dict__."""
        self.assertNotIn(
            "__init__",
            Module.__dict__,
            "Module must not define __init__. Adding __init__ to Module "
            "may break diamond inheritance (e.g. nn.Embedding + Module). "
            "Please verify all the use cases and ensure the change doesn't "
            "break them. After that, we can consider to remove this test.",
        )

    def test_forward(self):
        """Forward pass works through nn.Embedding's implementation."""
        emb = self.TestEmbedding(100, 32)
        out = emb(torch.tensor([0, 1, 2]))
        self.assertEqual(out.shape, torch.Size([3, 32]))

    def test_isinstance_checks(self):
        """Diamond class is an instance of all parent types."""
        emb = self.TestEmbedding(100, 32)
        self.assertIsInstance(emb, nn.Embedding)
        self.assertIsInstance(emb, nn.Module)
        self.assertIsInstance(emb, Module)

    def test_module_hierarchy_is_flat(self):
        """Diamond embedding adds no extra layer to the module tree."""

        class Model(Module):
            def __init__(self):
                super().__init__()
                self.embed = TestDiamondInheritance.TestEmbedding(100, 32)
                linear_config = Linear.Config(
                    in_features=32, out_features=16, bias=True
                )
                self.linear = linear_config.build()

        model = Model()
        param_names = {name for name, _ in model.named_parameters()}
        self.assertEqual(param_names, {"embed.weight", "linear.weight", "linear.bias"})

    def test_nn_module_init_called_once(self):
        """nn.Module.__init__ is called exactly once (no double init)."""
        call_count = 0
        orig_init = nn.Module.__init__

        def counting_init(self, *args, **kwargs):
            nonlocal call_count
            call_count += 1
            orig_init(self, *args, **kwargs)

        nn.Module.__init__ = counting_init
        try:
            self.TestEmbedding(50, 16)
            self.assertEqual(call_count, 1)
        finally:
            nn.Module.__init__ = orig_init


class TestNnModuleWrappers(unittest.TestCase):
    """Tests for native Module wrappers in nn_modules.py."""

    def test_is_subclass(self):
        """Wrapper class is subclass of both original and Module."""
        from torchtitan.models.common.nn_modules import Conv2d

        self.assertTrue(issubclass(Conv2d, nn.Conv2d))
        self.assertTrue(issubclass(Conv2d, Module))

    def test_isinstance(self):
        """Instance satisfies isinstance checks."""
        from torchtitan.models.common.nn_modules import Conv2d

        m = Conv2d.Config(
            in_channels=3,
            out_channels=16,
            kernel_size=3,
        ).build()
        self.assertIsInstance(m, nn.Conv2d)
        self.assertIsInstance(m, Module)

    def test_init_states_calls_reset_parameters(self):
        """init_states delegates to reset_parameters when no param_init."""
        from torchtitan.models.common.nn_modules import LayerNorm

        m = LayerNorm.Config(normalized_shape=32).build()
        nn.init.zeros_(m.weight)
        m.init_states()
        self.assertTrue(torch.allclose(m.weight, torch.ones(32)))

    def test_init_states_uses_custom_param_init(self):
        """Custom param_init runs instead of reset_parameters fallback.

        Guards against the old ``from_nn_module`` regression where the
        injected ``_init_self_parameters`` unconditionally called
        ``reset_parameters`` and silently dropped ``param_init``.
        """
        from torchtitan.models.common.nn_modules import LayerNorm

        m = LayerNorm.Config(
            normalized_shape=8,
            param_init={"weight": nn.init.zeros_, "bias": nn.init.zeros_},
        ).build()
        m.init_states()
        self.assertTrue(torch.all(m.weight == 0))
        self.assertTrue(torch.all(m.bias == 0))

    def test_init_states_noop_for_parameterless(self):
        """Parameterless modules handle init_states without error."""
        from torchtitan.models.common.nn_modules import GELU

        m = GELU.Config(approximate="tanh").build()
        m.init_states()  # should not raise

    def test_forward_unchanged(self):
        """Forward output is identical to plain nn module."""
        from torchtitan.models.common.nn_modules import LayerNorm

        torch.manual_seed(42)
        orig = nn.LayerNorm(16)
        wrapped = LayerNorm.Config(normalized_shape=16).build()
        wrapped.load_state_dict(orig.state_dict())
        x = torch.randn(2, 16)
        torch.testing.assert_close(orig(x), wrapped(x))

    def test_state_dict_unchanged(self):
        """state_dict keys and values match the plain nn module."""
        from torchtitan.models.common.nn_modules import Conv2d

        orig = nn.Conv2d(3, 16, 3)
        wrapped = Conv2d.Config(
            in_channels=3,
            out_channels=16,
            kernel_size=3,
        ).build()
        wrapped.load_state_dict(orig.state_dict())
        for key in orig.state_dict():
            self.assertIn(key, wrapped.state_dict())
            torch.testing.assert_close(
                orig.state_dict()[key], wrapped.state_dict()[key]
            )

    def test_config_has_typed_fields(self):
        """Config classes have proper typed fields."""
        from torchtitan.models.common.nn_modules import LayerNorm
        from torchtitan.protocols.sharding import ShardingConfig

        self.assertTrue(issubclass(LayerNorm.Config, Module.Config))
        cfg = LayerNorm.Config(normalized_shape=16, eps=1e-5)
        self.assertIsNone(cfg.sharding_config)
        self.assertIsNone(cfg.param_init)

        cfg = LayerNorm.Config(
            normalized_shape=16,
            sharding_config=ShardingConfig(),
        )
        self.assertIsInstance(cfg.sharding_config, ShardingConfig)

    def test_config_build_propagates_sharding(self):
        """Config.build propagates sharding_config to instance."""
        from torchtitan.models.common.nn_modules import LayerNorm
        from torchtitan.protocols.sharding import ShardingConfig

        sc = ShardingConfig()
        instance = LayerNorm.Config(
            normalized_shape=16,
            eps=1e-5,
            sharding_config=sc,
        ).build()
        self.assertIsInstance(instance, LayerNorm)
        self.assertEqual(instance.normalized_shape, (16,))
        self.assertEqual(instance.eps, 1e-5)
        self.assertIs(instance._sharding_config, sc)


class TestContainerInitStates(unittest.TestCase):
    """Tests for ModuleList, ModuleDict, Sequential init_states."""

    def test_module_list_init_states(self):
        """ModuleList.init_states initializes children."""
        from torchtitan.models.common.nn_modules import LayerNorm

        ln_cfg = LayerNorm.Config(normalized_shape=8)
        norms = ModuleList([ln_cfg.build() for _ in range(3)])
        for n in norms:
            nn.init.zeros_(n.weight)
        norms.init_states()
        for n in norms:
            self.assertTrue(torch.allclose(n.weight, torch.ones(8)))

    def test_module_dict_init_states(self):
        """ModuleDict.init_states initializes children."""
        from torchtitan.models.common.nn_modules import LayerNorm

        ln_cfg = LayerNorm.Config(normalized_shape=8)
        norms = ModuleDict({"a": ln_cfg.build(), "b": ln_cfg.build()})
        for n in norms.values():
            nn.init.zeros_(n.weight)
        norms.init_states()
        for n in norms.values():
            self.assertTrue(torch.allclose(n.weight, torch.ones(8)))

    def test_sequential_init_states(self):
        """Sequential.init_states recurses into children."""
        from torchtitan.models.common.nn_modules import GELU

        seq = Sequential(GELU.Config().build())
        seq.init_states()  # should not raise

    def test_containers_are_module(self):
        """Container instances satisfy Module protocol."""
        self.assertIsInstance(ModuleList(), Module)
        self.assertIsInstance(ModuleDict(), Module)
        self.assertIsInstance(Sequential(), Module)


class TestConfigBuildPropagatesParamInit(unittest.TestCase):
    """Tests for Config.build() propagating param_init to the instance."""

    def test_param_init_on_instance(self):
        """build() sets _param_init on the constructed instance."""
        param_init = {"weight": nn.init.zeros_}
        config = Linear.Config(in_features=4, out_features=4, param_init=param_init)
        linear = config.build()
        self.assertTrue(hasattr(linear, "_param_init"))
        self.assertIs(linear._param_init, param_init)

    def test_no_param_init_by_default(self):
        """build() without param_init leaves it as None."""
        config = Linear.Config(in_features=4, out_features=4)
        linear = config.build()
        self.assertIsNone(linear._param_init)

    def test_init_states_uses_config_param_init(self):
        """init_states uses param_init from config when available."""

        class Parent(Module):
            def __init__(self):
                super().__init__()
                linear_config = Linear.Config(
                    in_features=4, out_features=4, param_init={"weight": nn.init.ones_}
                )
                self.linear = linear_config.build()

        m = Parent()
        nn.init.zeros_(m.linear.weight)
        m.init_states()
        self.assertTrue(torch.all(m.linear.weight == 1))


class TestModuleRedistribution(unittest.TestCase):
    class WeightModule(Module):
        def __init__(self, shape: tuple[int, ...], layout: SpmdType):
            super().__init__()
            self.weight = nn.Parameter(torch.empty(shape))
            self._sharding_config = ShardingConfig(state_shardings={"weight": layout})

    def test_rejects_uneven_tp_parameter_sharding(self):
        module = self.WeightModule(
            (4, 5),
            SpmdType(
                {MeshAxisName.TP: spmd.V},
                partition_spec=spmd.PartitionSpec(None, MeshAxisName.TP),
            ),
        )
        parallelism_context = ParallelismContext(
            dp_replicate=1,
            dp_shard=1,
            cp=1,
            tp=2,
            pp=1,
            ep=1,
            world_size=2,
            enable_sequence_parallel=False,
        )

        with self.assertRaisesRegex(
            ValueError,
            r"WeightModule\.weight.*tensor dimension 1.*mesh axis tp with size 2",
        ):
            module._parallelize(parallelism_context)

    def test_rejects_uneven_ep_parameter_sharding(self):
        module = self.WeightModule(
            (3, 4),
            SpmdType({MeshAxisName.EP: spmd.S(0)}),
        )
        parallelism_context = ParallelismContext(
            dp_replicate=1,
            dp_shard=2,
            cp=1,
            tp=1,
            pp=1,
            ep=2,
            world_size=2,
            enable_sequence_parallel=False,
        )

        with self.assertRaisesRegex(
            ValueError,
            r"WeightModule\.weight.*tensor dimension 0.*mesh axis ep with size 2",
        ):
            module._parallelize(parallelism_context)


class TestParallelizeModuleProtocol(unittest.TestCase):
    """Tests for protocol validation during Module._parallelize."""

    def test_passes_for_all_module(self):
        """No error when all submodules are Module instances."""
        from torchtitan.protocols.model import BaseModel

        class GoodModel(BaseModel):
            @dataclass(kw_only=True, slots=True)
            class Config(BaseModel.Config):
                def get_nparams_and_flops(self, model, seq_len):
                    return (0, 0)

            def __init__(self):
                super().__init__()
                linear_config = Linear.Config(in_features=4, out_features=4)
                self.linear = linear_config.build()

            def _apply_fsdp(self, **kwargs):
                pass

        model = GoodModel()
        model._parallelize(None)

    def test_default_raises_for_plain_nn_module(self):
        """A plain stateful nn.Module is rejected during parallelization."""
        from torchtitan.protocols.model import BaseModel

        class BadModel(BaseModel):
            @dataclass(kw_only=True, slots=True)
            class Config(BaseModel.Config):
                def get_nparams_and_flops(self, model, seq_len):
                    return (0, 0)

            def __init__(self):
                super().__init__()
                self.plain = nn.Linear(4, 4)

            def _apply_fsdp(self, **kwargs):
                pass

        model = BadModel()
        with self.assertRaises(RuntimeError):
            model._parallelize(None)

    def test_stateless_plain_module_is_allowed(self):
        """Stateless PyTorch modules do not need the TorchTitan protocol."""
        from torchtitan.protocols.model import BaseModel

        class ThirdPartyModel(BaseModel):
            @dataclass(kw_only=True, slots=True)
            class Config(BaseModel.Config):
                def get_nparams_and_flops(self, model, seq_len):
                    return (0, 0)

            def __init__(self):
                super().__init__()
                self.plain = nn.ReLU()

            def _apply_fsdp(self, **kwargs):
                pass

        model = ThirdPartyModel()
        model._parallelize(None)

    def test_explicitly_exempt_stateful_child_is_allowed(self):
        """A protocol module may own an opaque third-party implementation."""
        from torchtitan.protocols.model import BaseModel

        class ThirdPartyModel(BaseModel):
            _module_protocol_exempt_children = frozenset({"plain"})

            @dataclass(kw_only=True, slots=True)
            class Config(BaseModel.Config):
                def get_nparams_and_flops(self, model, seq_len):
                    return (0, 0)

            def __init__(self):
                super().__init__()
                self.plain = nn.Linear(4, 4)

            def _apply_fsdp(self, **kwargs):
                pass

        model = ThirdPartyModel()
        model._parallelize(None)


if __name__ == "__main__":
    unittest.main()
