# 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 unittest.mock import patch

import torch
import torch.distributed as dist
import torch.nn as nn

from torchtitan.config.configs import TrainingConfig
from torchtitan.distributed import ParallelismContext
from torchtitan.experiments.graph_trainer.common_utils import apply_simple_fsdp
from torchtitan.models.common.attention import ScaledDotProductInnerAttention


class TestApplySimpleFSDPSingleRank(unittest.TestCase):
    """Verify simple_fsdp's MixedPrecisionPolicy actually casts params at NGPU=1."""

    def setUp(self):
        if not dist.is_initialized():
            dist.init_process_group(
                backend="gloo",
                init_method="tcp://localhost:12358",
                world_size=1,
                rank=0,
            )

    def tearDown(self):
        if dist.is_initialized():
            dist.destroy_process_group()

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_uses_dtensor_storage_and_local_compute(self):
        parallelism_context = ParallelismContext(
            dp_replicate=1,
            dp_shard=1,
            cp=1,
            tp=1,
            pp=1,
            ep=1,
            world_size=1,
            enable_sequence_parallel=False,
        )
        training = TrainingConfig(
            mixed_precision_param="bfloat16",
            mixed_precision_reduce="float32",
        )

        model = apply_simple_fsdp(
            nn.Linear(8, 8),
            parallelism_context=parallelism_context,
            training=training,
        )

        self.assertIsInstance(
            model._parameters["weight"], torch.distributed.tensor.DTensor
        )
        self.assertEqual(model._parameters["weight"].dtype, torch.float32)
        self.assertNotIsInstance(model.weight, torch.distributed.tensor.DTensor)
        self.assertEqual(model.weight.dtype, torch.bfloat16)
        self.assertEqual(
            model(torch.randn(2, 8, dtype=torch.bfloat16)).dtype, torch.bfloat16
        )

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_preserves_inner_attention_metadata_key(self):
        parallelism_context = ParallelismContext(
            dp_replicate=1,
            dp_shard=1,
            cp=1,
            tp=1,
            pp=1,
            ep=1,
            world_size=1,
            enable_sequence_parallel=False,
        )
        training = TrainingConfig(
            mixed_precision_param="bfloat16",
            mixed_precision_reduce="float32",
        )
        inner_attention = ScaledDotProductInnerAttention(
            ScaledDotProductInnerAttention.Config()
        )

        model = apply_simple_fsdp(
            inner_attention,
            parallelism_context=parallelism_context,
            training=training,
        )

        self.assertIs(model.attention_metadata_key, ScaledDotProductInnerAttention)


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