# 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 types import SimpleNamespace
from unittest.mock import MagicMock, patch

import torch

from torchtitan.models.deepseek_v4.model import DeepSeekV4Model
from torchtitan.models.deepseek_v4.sharding import set_deepseek_v4_sharding_config
from torchtitan_recipes.tests.models.deepseek_v4 import (
    deepseek_v4_debugmodel,
    deepseek_v4_mtp_debugmodel,
)


class TestDeepSeekV4MTPConfig(unittest.TestCase):
    def test_dataloader_mtp_depth_matches_model(self):
        for config in (deepseek_v4_debugmodel(), deepseek_v4_mtp_debugmodel()):
            self.assertEqual(
                config.dataloader.num_mtp_layers, config.model.n_mtp_layers
            )

    def test_mtp_debugmodel_builds_mtp_layers(self):
        config = deepseek_v4_mtp_debugmodel()
        model_config = config.model
        self.assertEqual(model_config.n_mtp_layers, 1)
        self.assertIsNotNone(model_config.mtp_layers)
        self.assertEqual(len(model_config.mtp_layers), 1)

    def test_mtp_metadata_remains_replicated_at_block_boundary(self):
        config = deepseek_v4_mtp_debugmodel().model
        set_deepseek_v4_sharding_config(config, enable_sp=True, enable_ep=True)

        assert config.mtp_layers is not None
        mtp_config = config.mtp_layers[0].sharding_config
        assert mtp_config is not None
        self.assertIsNotNone(mtp_config.in_src_shardings)
        self.assertIsNone(mtp_config.in_dst_shardings)

    @patch("torchtitan.distributed.fsdp.resolve_sparse_fsdp_mesh")
    @patch("torchtitan.distributed.fsdp.resolve_fsdp_mesh")
    @patch("torchtitan.models.deepseek_v4.model.apply_fsdp_to_mtp_decoder")
    def test_uses_mtp_fsdp_path(
        self,
        apply_fsdp_to_mtp_decoder,
        resolve_fsdp_mesh,
        resolve_sparse_fsdp_mesh,
    ):
        model = MagicMock(spec=DeepSeekV4Model)
        dp_mesh = object()
        dp_mesh_dims = object()
        edp_mesh = object()
        edp_mesh_dims = object()
        resolve_fsdp_mesh.return_value = (dp_mesh, dp_mesh_dims)
        resolve_sparse_fsdp_mesh.return_value = (edp_mesh, edp_mesh_dims)
        parallelism_context = SimpleNamespace(pp_enabled=False, ep=2)
        training = SimpleNamespace(
            mixed_precision_param="bfloat16",
            mixed_precision_reduce="float32",
            enable_cpu_offload=False,
        )
        parallelism = SimpleNamespace(
            fsdp_reshard_after_forward="default",
            fsdp_symm_mem_scope=None,
        )

        DeepSeekV4Model._apply_fsdp(
            model,
            parallelism_context=parallelism_context,
            training=training,
            parallelism=parallelism,
        )

        apply_fsdp_to_mtp_decoder.assert_called_once_with(
            model,
            dp_mesh,
            param_dtype=torch.bfloat16,
            reduce_dtype=torch.float32,
            pp_enabled=False,
            cpu_offload=False,
            reshard_after_forward_policy="default",
            ep_degree=2,
            edp_mesh=edp_mesh,
            dp_mesh_dims=dp_mesh_dims,
            edp_mesh_dims=edp_mesh_dims,
            symm_mem_scope=None,
        )


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