# 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 copy
import json
import os
import socket
import tempfile
import unittest
from unittest.mock import patch

import spmd_types as spmd
import torch
import torch.distributed as dist
from spmd_types import SpmdType
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.tensor import Replicate
from torch.distributed.tensor.debug import CommDebugMode
from torch.testing._internal.distributed._tensor.common_dtensor import (
    DTensorTestBase,
    with_comms,
)
from torchtitan.config.parallelism import ParallelismConfig
from torchtitan.distributed.fsdp import apply_fsdp_to_decoder
from torchtitan.distributed.parallelism_context import (
    DistributedTopology,
    MeshAxisName,
    ParallelismContext,
    unfold_dp_axes,
)
from torchtitan.distributed.spmd_types import (
    _per_axis_types,
    spmd_distribute_tensor,
    spmd_redistribute_per_axis,
    spmd_validate_redistributions,
)
from torchtitan.models.common.decoder_sharding import (
    attention_activation_placement,
    dense_activation_placement,
    dense_sequence_parallel_placement,
    token_id_placement,
)
from torchtitan.models.llama3 import build_model_config
from torchtitan.protocols.sharding import resolve_placements, ShardingConfig


class TestParallelismContextValidation(unittest.TestCase):
    """Test ParallelismContext validation logic without mesh building."""

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_basic_initialization(self):
        """Test basic initialization with valid parameters."""
        parallelism_context = ParallelismContext(
            dp_replicate=2,
            dp_shard=2,
            cp=1,
            tp=2,
            pp=1,
            ep=1,
            world_size=8,
            enable_sequence_parallel=True,
        )
        self.assertEqual(parallelism_context.dp_replicate, 2)
        self.assertEqual(parallelism_context.dp_shard, 2)
        self.assertEqual(parallelism_context.cp, 1)
        self.assertEqual(parallelism_context.tp, 2)
        self.assertEqual(parallelism_context.pp, 1)
        self.assertEqual(parallelism_context.ep, 1)
        self.assertEqual(parallelism_context.world_size, 8)
        self.assertTrue(parallelism_context.sp_enabled)

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_from_config(self):
        """Test constructing ParallelismContext from a ParallelismConfig."""
        config = ParallelismConfig(
            data_parallel_replicate_degree=2,
            data_parallel_shard_degree=-1,
            context_parallel_degree=1,
            tensor_parallel_degree=2,
            pipeline_parallel_degree=1,
            expert_parallel_degree=1,
            enable_sequence_parallel=False,
        )
        parallelism_context = ParallelismContext.from_config(
            config, DistributedTopology(world_size=8), dump_folder=""
        )
        self.assertEqual(parallelism_context.dp_replicate, 2)
        self.assertEqual(
            parallelism_context.dp_shard, 2
        )  # auto-calculated: 8 / (2*1*2*1)
        self.assertEqual(parallelism_context.cp, 1)
        self.assertEqual(parallelism_context.tp, 2)
        self.assertEqual(parallelism_context.pp, 1)
        self.assertEqual(parallelism_context.ep, 1)
        self.assertEqual(parallelism_context.world_size, 8)
        self.assertFalse(parallelism_context.sp_enabled)

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_auto_calculate_dp_shard(self):
        """Test automatic calculation of dp_shard when set to -1."""
        parallelism_context = ParallelismContext(
            dp_replicate=2,
            dp_shard=-1,
            cp=1,
            tp=2,
            pp=1,
            ep=1,
            world_size=8,
            enable_sequence_parallel=False,
        )
        self.assertEqual(parallelism_context.dp_shard, 2)

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_validation_invalid_world_size(self):
        """Test validation fails when parallelism degrees don't match world_size."""
        with self.assertRaisesRegex(ValueError, r"!= WORLD_SIZE"):
            ParallelismContext(
                dp_replicate=2,
                dp_shard=2,
                cp=1,
                tp=2,
                pp=1,
                ep=1,
                world_size=10,  # Invalid: 2*2*1*2*1 = 8, not 10
                enable_sequence_parallel=False,
            )

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_validation_zero_parallelism(self):
        """Test validation fails when parallelism degree is 0."""
        with self.assertRaisesRegex(ValueError, r"degree should be >= 1"):
            ParallelismContext(
                dp_replicate=0,  # Invalid: must be >= 1
                dp_shard=1,
                cp=1,
                tp=1,
                pp=1,
                ep=1,
                world_size=1,
                enable_sequence_parallel=False,
            )

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_validation_invalid_dp_shard(self):
        """Test validation fails when dp_shard is invalid (not -1 and not >=1)."""
        with self.assertRaisesRegex(ValueError, r"dp_shard must"):
            ParallelismContext(
                dp_replicate=1,
                dp_shard=0,  # Invalid: must be -1 or >= 1
                cp=1,
                tp=1,
                pp=1,
                ep=1,
                world_size=1,
                enable_sequence_parallel=False,
            )

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_enabled_properties(self):
        """Test all enabled properties."""
        # Test with DP enabled
        parallelism_context = ParallelismContext(
            dp_replicate=2,
            dp_shard=2,
            cp=1,
            tp=2,
            pp=1,
            ep=1,
            world_size=8,
            enable_sequence_parallel=True,
        )
        self.assertTrue(parallelism_context.dp_enabled)
        self.assertTrue(parallelism_context.dp_replicate_enabled)
        self.assertTrue(parallelism_context.dp_shard_enabled)
        self.assertFalse(parallelism_context.cp_enabled)
        self.assertTrue(parallelism_context.tp_enabled)
        self.assertTrue(parallelism_context.sp_enabled)
        self.assertFalse(parallelism_context.pp_enabled)
        self.assertFalse(parallelism_context.ep_enabled)
        self.assertTrue(parallelism_context.fsdp_enabled)

        # Test with CP enabled
        parallelism_context = ParallelismContext(
            dp_replicate=1,
            dp_shard=1,
            cp=2,
            tp=1,
            pp=1,
            ep=1,
            world_size=2,
            enable_sequence_parallel=True,
        )
        self.assertFalse(parallelism_context.dp_enabled)
        self.assertTrue(parallelism_context.cp_enabled)
        self.assertTrue(parallelism_context.dp_cp_enabled)
        self.assertTrue(parallelism_context.fsdp_enabled)
        self.assertFalse(parallelism_context.sp_enabled)

        # Test with EP enabled (EP must not contribute to world_size)
        parallelism_context = ParallelismContext(
            dp_replicate=1,
            dp_shard=2,
            cp=1,
            tp=1,
            pp=1,
            ep=2,
            world_size=2,
            enable_sequence_parallel=False,
        )
        self.assertTrue(parallelism_context.ep_enabled)

        # Test with PP enabled
        parallelism_context = ParallelismContext(
            dp_replicate=1,
            dp_shard=1,
            cp=1,
            tp=1,
            pp=2,
            ep=1,
            world_size=2,
            enable_sequence_parallel=False,
        )
        self.assertTrue(parallelism_context.pp_enabled)

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_non_data_parallel_size(self):
        """Test non_data_parallel_size calculation."""
        parallelism_context = ParallelismContext(
            dp_replicate=2,
            dp_shard=2,
            cp=2,
            tp=3,
            pp=2,
            ep=1,
            world_size=48,
            enable_sequence_parallel=False,
        )
        # Should be cp * tp * pp = 2 * 3 * 2 = 12
        self.assertEqual(parallelism_context.non_data_parallel_size, 12)

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_seq_len_divisor(self):
        """Test seq_len_divisor calculation."""
        parallelism_context = ParallelismContext(
            dp_replicate=2,
            dp_shard=1,
            cp=2,
            tp=4,
            pp=1,
            ep=1,
            world_size=16,
            enable_sequence_parallel=False,
        )
        # Should be tp * (cp * 2) = 4 * 4 = 16
        self.assertEqual(parallelism_context.seq_len_divisor, 16)


class TestSpmdLayout(DTensorTestBase):
    @property
    def world_size(self):
        return 4

    def test_seq_parallel_activation_per_axis_spmd_types(self):
        """PartitionSpec can map multiple mesh axes to one tensor dim."""
        layout = SpmdType(
            {
                MeshAxisName.DP: spmd.V,
                MeshAxisName.CP: spmd.V,
                MeshAxisName.TP: spmd.V,
            },
            partition_spec=spmd.PartitionSpec(
                MeshAxisName.DP,
                (MeshAxisName.CP, MeshAxisName.TP),
                None,
            ),
        )

        self.assertEqual(
            _per_axis_types(layout),
            {
                MeshAxisName.DP: spmd.S(0),
                MeshAxisName.CP: spmd.S(1),
                MeshAxisName.TP: spmd.S(1),
            },
        )

    def test_decoder_layout_partition_spec_ranks(self):
        self.assertEqual(
            token_id_placement().partition_spec,
            ((MeshAxisName.DP, MeshAxisName.CP),),
        )
        self.assertEqual(
            token_id_placement(enable_sp=True).partition_spec,
            ((MeshAxisName.DP, MeshAxisName.CP, MeshAxisName.TP),),
        )
        self.assertEqual(
            dense_activation_placement(tp=spmd.R, cp=spmd.S(0)).partition_spec,
            ((MeshAxisName.DP, MeshAxisName.CP), None),
        )
        self.assertEqual(
            attention_activation_placement().partition_spec,
            ((MeshAxisName.DP, MeshAxisName.CP), MeshAxisName.TP, None),
        )

    def test_unfold_dp_axes(self):
        """Logical DP expands only when resolving concrete mesh axes."""
        self.assertEqual(
            unfold_dp_axes([MeshAxisName.DP, MeshAxisName.CP, MeshAxisName.TP]),
            ["dp_replicate", "dp_shard", "cp", "tp"],
        )

    @with_comms
    def test_resolve_placements_ignores_extra_untranslatable_axes(self):
        """Extra layout axes are ignored before converting to DTensor placements."""
        mesh = init_device_mesh(
            self.device_type, (self.world_size,), mesh_dim_names=("tp",)
        )
        layout = SpmdType(
            {
                MeshAxisName.DP: spmd.V,
                MeshAxisName.TP: spmd.I,
            }
        )

        self.assertEqual(resolve_placements(layout, mesh), (Replicate(),))

    def test_rejects_partition_spec_reorder_redistribute(self):
        """((DP, CP), None) -> ((CP, DP), None) not supported by a single redistribute call."""
        with self.assertRaises(ValueError) as cm:
            spmd_validate_redistributions(
                ShardingConfig(
                    out_src_shardings=SpmdType(
                        {
                            MeshAxisName.DP: spmd.V,
                            MeshAxisName.CP: spmd.V,
                        },
                        partition_spec=spmd.PartitionSpec(
                            (MeshAxisName.DP, MeshAxisName.CP), None
                        ),
                    ),
                    out_dst_shardings=SpmdType(
                        {
                            MeshAxisName.DP: spmd.V,
                            MeshAxisName.CP: spmd.V,
                        },
                        partition_spec=spmd.PartitionSpec(
                            (MeshAxisName.CP, MeshAxisName.DP), None
                        ),
                    ),
                )
            )

    def test_rejects_multi_axis_redistribute(self):
        """Redistributing multiple mesh axes is unsupported."""
        with self.assertRaises(ValueError) as cm:
            spmd_validate_redistributions(
                ShardingConfig(
                    in_src_shardings={
                        "x": SpmdType(
                            {
                                MeshAxisName.DP: spmd.S(0),
                                MeshAxisName.CP: spmd.S(1),
                                MeshAxisName.TP: spmd.R,
                            }
                        )
                    },
                    in_dst_shardings={
                        "x": SpmdType(
                            {
                                MeshAxisName.DP: spmd.R,
                                MeshAxisName.CP: spmd.R,
                                MeshAxisName.TP: spmd.R,
                            }
                        )
                    },
                )
            )

    def test_rejects_redistribute_from_varying(self):
        for src_dp, dst_dp in ((spmd.V, spmd.R), (spmd.R, spmd.V)):
            with self.subTest(src_dp=src_dp, dst_dp=dst_dp):
                with self.assertRaisesRegex(
                    ValueError,
                    "output: SpmdType-based redistribution changes mesh axis "
                    "'dp' with spmd.V as the source or destination type",
                ):
                    spmd_validate_redistributions(
                        ShardingConfig(
                            out_src_shardings=SpmdType(
                                {
                                    MeshAxisName.DP: src_dp,
                                    MeshAxisName.TP: spmd.I,
                                }
                            ),
                            out_dst_shardings=SpmdType(
                                {
                                    MeshAxisName.DP: dst_dp,
                                    MeshAxisName.TP: spmd.I,
                                }
                            ),
                        )
                    )

    @with_comms
    def test_partition_spec_order_controls_state_shard(self):
        """Test spmd_distribute_tensor follows PartitionSpec order.

        ``(DP, CP)`` and ``(CP, DP)`` both shard dim 0 across the same 2x2
        mesh, but assign different global slices to ranks. This verifies local
        state sharding preserves that ordering instead of relying on unordered
        per-axis shard types.
        """
        mesh = init_device_mesh(self.device_type, (2, 2), mesh_dim_names=("dp", "cp"))
        global_weight = torch.arange(
            8, dtype=torch.float32, device=self.device_type
        ).reshape(8, 1)

        for axis_order in (
            (MeshAxisName.DP, MeshAxisName.CP),
            (MeshAxisName.CP, MeshAxisName.DP),
        ):
            with self.subTest(axis_order=axis_order):
                layout = SpmdType(
                    {
                        MeshAxisName.DP: spmd.V,
                        MeshAxisName.CP: spmd.V,
                    },
                    partition_spec=spmd.PartitionSpec(axis_order, None),
                )

                local_weight = spmd_distribute_tensor(
                    global_weight.clone(), mesh, layout
                )

                axis_ranks = {
                    MeshAxisName.DP: mesh.get_local_rank("dp"),
                    MeshAxisName.CP: mesh.get_local_rank("cp"),
                }
                axis_sizes = {
                    MeshAxisName.DP: mesh.size(0),
                    MeshAxisName.CP: mesh.size(1),
                }
                shard_idx = 0
                for axis_name in axis_order:
                    shard_idx = (
                        shard_idx * axis_sizes[axis_name] + axis_ranks[axis_name]
                    )
                local_rows = global_weight.shape[0] // self.world_size
                expected = global_weight.narrow(0, shard_idx * local_rows, local_rows)
                torch.testing.assert_close(local_weight, expected)

    @with_comms
    def test_spmd_redistribute_per_axis_allgather(self):
        """
        Test spmd_redistribute_per_axis performs seq-dim allgather.
        src: dense SP placement (V + PartitionSpec)
        dst: dense activation w/ I@TP
        """
        mesh = init_device_mesh(
            self.device_type,
            (1, 1, 4),
            mesh_dim_names=("dp", "cp", "tp"),
        )
        x = torch.ones(2, 2, device=self.device_type)
        src = dense_sequence_parallel_placement()
        dst = dense_activation_placement(tp=spmd.I, cp=spmd.S(0))

        comm_mode = CommDebugMode()
        with comm_mode:
            result = spmd_redistribute_per_axis(
                x,
                mesh,
                src,
                dst,
            )

        self.assertEqual(comm_mode.get_total_counts(), 1)
        self.assertTrue(torch.equal(result, torch.ones(8, 2, device=self.device_type)))

    @with_comms
    def test_spmd_redistribute_per_axis_shards_token_ids(self):
        """R@TP token IDs chunk T when moving to the 1D sequence-parallel layout."""
        mesh = init_device_mesh(
            self.device_type,
            (1, 1, 4),
            mesh_dim_names=("dp", "cp", "tp"),
        )
        x = torch.arange(8, device=self.device_type)
        src = token_id_placement()
        dst = token_id_placement(enable_sp=True)
        spmd_validate_redistributions(
            ShardingConfig(
                in_src_shardings={"input_ids_T": src},
                in_dst_shardings={"input_ids_T": dst},
            )
        )

        comm_mode = CommDebugMode()
        with comm_mode:
            result = spmd_redistribute_per_axis(x, mesh, src, dst)

        # convert(R, S(0)) is a local chunk, not a collective.
        self.assertEqual(comm_mode.get_total_counts(), 0)
        tp_rank = mesh.get_local_rank("tp")
        expected = torch.arange(tp_rank * 2, tp_rank * 2 + 2, device=self.device_type)
        self.assertTrue(torch.equal(result, expected))


class TestParallelismContextMeshOperations(unittest.TestCase):
    """Test ParallelismContext mesh operations with single-rank distributed environment."""

    def setUp(self):
        """Initialize distributed environment for CPU testing."""
        if not dist.is_initialized():
            dist.init_process_group(
                backend="gloo",
                init_method="tcp://localhost:12356",
                world_size=1,
                rank=0,
            )

    def tearDown(self):
        """Clean up distributed environment."""
        if dist.is_initialized():
            dist.destroy_process_group()

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_from_config_saves_parallelism_folder(self):
        with tempfile.TemporaryDirectory() as dump_folder, patch.dict(
            os.environ, {"LOCAL_RANK": "0"}
        ):
            parallelism_context = ParallelismContext.from_config(
                ParallelismConfig(save_parallelism_folder="sub/parallelism"),
                DistributedTopology(world_size=1),
                dump_folder=dump_folder,
            )
            with open(
                os.path.join(dump_folder, "sub", "parallelism", "rank_0.json")
            ) as f:
                layout = json.load(f)

        self.assertEqual(
            {k: layout[k] for k in ("host", "local_rank", "global_rank")},
            {"host": socket.gethostname(), "local_rank": 0, "global_rank": 0},
        )
        self.assertEqual(layout["world_size"], 1)
        self.assertEqual(
            layout["meshes"],
            {
                name: {
                    "axis_names": list(mesh.mesh_dim_names),
                    "mesh": mesh.mesh.tolist(),
                }
                for name, mesh in parallelism_context._global_meshes.items()
            },
        )
        self.assertEqual(
            layout["meshes"]["dense"]["axis_names"],
            ["pp", "dp_replicate", "dp_shard", "cp", "tp"],
        )
        # The meshes built for the file are the ones later lookups return.
        self.assertIs(
            parallelism_context.get_mesh("dp_shard"),
            parallelism_context._single_axis_meshes["dp_shard"],
        )

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_from_config_skips_mesh_build_without_parallelism_folder(self):
        parallelism_context = ParallelismContext.from_config(
            ParallelismConfig(), DistributedTopology(world_size=1), dump_folder=""
        )
        self.assertEqual(parallelism_context._single_axis_meshes, {})

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_real_pp_group_for_fake_spmd_is_used_during_mesh_construction(self):
        group = dist.distributed_c10d._get_default_group()
        topology = DistributedTopology(
            world_size=1,
            real_pp_group_for_fake_spmd=group,
        )
        parallelism_context = ParallelismContext.from_config(
            ParallelismConfig(), topology, dump_folder=""
        )

        parallelism_context.build_mesh()

        pp_mesh = parallelism_context.get_optional_mesh(
            "pp", include_singleton_axes=True
        )
        assert pp_mesh is not None
        self.assertIs(pp_mesh.get_group(), group)

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_get_mesh_invalid_name(self):
        """Test getting mesh with invalid name raises error."""
        parallelism_context = ParallelismContext(
            dp_replicate=1,
            dp_shard=1,
            cp=1,
            tp=1,
            pp=1,
            ep=1,
            world_size=1,
            enable_sequence_parallel=False,
        )
        parallelism_context.build_mesh()

        with self.assertRaises(ValueError) as context:
            parallelism_context.get_mesh("invalid_mesh")
        self.assertIn("Invalid mesh axis", str(context.exception))

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_get_mesh_lazy_initialization(self):
        """Test that get_optional_mesh triggers build_mesh if not built yet."""
        parallelism_context = ParallelismContext(
            dp_replicate=1,
            dp_shard=1,
            cp=1,
            tp=1,
            pp=1,
            ep=1,
            world_size=1,
            enable_sequence_parallel=False,
        )
        # Don't call build_mesh explicitly
        self.assertEqual(len(parallelism_context._single_axis_meshes), 0)

        # get_optional_mesh should trigger build_mesh
        # Result is None because tp has size 1, but build_mesh should have been called
        self.assertIsNone(parallelism_context.get_optional_mesh("tp"))
        self.assertGreater(len(parallelism_context._single_axis_meshes), 0)

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_single_rank_mesh_operations(self):
        """Comprehensive test for all single-rank mesh operations.

        This test verifies mesh building, mesh retrieval, mesh sizes, and property
        access when all parallelism dimensions are set to 1 (single rank).
        """
        parallelism_context = ParallelismContext(
            dp_replicate=1,
            dp_shard=1,
            cp=1,
            tp=1,
            pp=1,
            ep=1,
            world_size=1,
            enable_sequence_parallel=False,
        )

        # Test mesh building
        world_mesh = parallelism_context.build_mesh()
        self.assertIsNotNone(world_mesh)
        self.assertEqual(world_mesh.size(), 1)

        # Verify all expected meshes are created
        self.assertIsNotNone(parallelism_context._single_axis_meshes)
        self.assertIn("pp", parallelism_context._single_axis_meshes)
        self.assertIn("loss", parallelism_context._single_axis_meshes)
        self.assertIn("dp_replicate", parallelism_context._single_axis_meshes)
        self.assertIn("dp", parallelism_context._single_axis_meshes)
        self.assertIn("dp_shard", parallelism_context._single_axis_meshes)
        self.assertIn("cp", parallelism_context._single_axis_meshes)
        self.assertIn("tp", parallelism_context._single_axis_meshes)

        # Validate 1D mesh sizes - all should be 1 for single rank
        self.assertEqual(
            parallelism_context._single_axis_meshes["dp_replicate"].size(), 1
        )
        self.assertEqual(parallelism_context._single_axis_meshes["dp"].size(), 1)
        self.assertEqual(parallelism_context._single_axis_meshes["dp_shard"].size(), 1)
        self.assertEqual(parallelism_context._single_axis_meshes["tp"].size(), 1)
        self.assertEqual(parallelism_context._single_axis_meshes["loss"].size(), 1)
        self.assertEqual(parallelism_context._single_axis_meshes["pp"].size(), 1)
        self.assertEqual(parallelism_context._single_axis_meshes["cp"].size(), 1)
        self.assertEqual(parallelism_context._single_axis_meshes["ep"].size(), 1)
        self.assertEqual(parallelism_context._single_axis_meshes["edp_shard"].size(), 1)

        # Validate 2D mesh shapes
        dp_replicate_fsdp_mesh = parallelism_context.get_optional_mesh(
            ["dp_replicate", "dp_shard"]
        )
        self.assertIsNone(dp_replicate_fsdp_mesh)  # Both dimensions have size 1
        dp_replicate_edp_shard_mesh = parallelism_context.get_optional_mesh(
            ["dp_replicate", "edp_shard"]
        )
        self.assertIsNone(dp_replicate_edp_shard_mesh)  # Both dimensions have size 1

        # Test get_optional_mesh returns None when all dimensions have size 1
        self.assertIsNone(parallelism_context.get_optional_mesh("tp"))
        self.assertIsNone(parallelism_context.get_optional_mesh("dp_replicate"))
        self.assertIsNone(parallelism_context.get_optional_mesh("pp"))
        self.assertIsNone(parallelism_context.get_optional_mesh("cp"))

        # Test get_optional_mesh with list input
        self.assertIsNone(
            parallelism_context.get_optional_mesh(["dp_replicate", "dp_shard"])
        )

        # Test get_all_one_dimensional_meshes returns empty when all dimensions have size 1
        one_d_meshes = parallelism_context.get_all_one_dimensional_meshes()
        self.assertEqual(len(one_d_meshes), 0)

        # Test world_mesh property
        world_mesh_property = parallelism_context.world_mesh
        self.assertIsNotNone(world_mesh_property)
        self.assertEqual(world_mesh_property.size(), 1)

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_get_mesh_with_list_input(self):
        """Test get_optional_mesh accepts both string and list inputs."""
        parallelism_context = ParallelismContext(
            dp_replicate=1,
            dp_shard=1,
            cp=1,
            tp=1,
            pp=1,
            ep=1,
            world_size=1,
            enable_sequence_parallel=False,
        )
        parallelism_context.build_mesh()

        # Should accept list input
        result = parallelism_context.get_optional_mesh(["dp_replicate", "dp_shard"])
        # Returns None because both dimensions have size 1
        self.assertIsNone(result)

    @patch("torchtitan.distributed.parallelism_context.device_type", "cpu")
    def test_expert_parallelism_validation(self):
        """Test expert parallelism configurations."""
        # EP enabled (valid) - world_size = dp_replicate * dp_shard * cp * tp * pp
        parallelism_context = ParallelismContext(
            dp_replicate=1,
            dp_shard=2,
            cp=1,
            tp=1,
            pp=1,
            ep=2,
            world_size=2,  # 1 * 2 * 1 * 1 * 1 = 2
            enable_sequence_parallel=False,
        )
        self.assertTrue(parallelism_context.ep_enabled)

        # Test with larger configuration
        parallelism_context = ParallelismContext(
            dp_replicate=2,
            dp_shard=2,
            cp=1,
            tp=1,
            pp=1,
            ep=2,
            world_size=4,  # 2 * 2 * 1 * 1 * 1 = 4
            enable_sequence_parallel=False,
        )
        self.assertTrue(parallelism_context.ep_enabled)
        self.assertTrue(parallelism_context.dp_replicate_enabled)
        self.assertTrue(parallelism_context.dp_shard_enabled)

        with self.assertRaisesRegex(
            ValueError,
            r"expert_parallel_degree \(3\) must divide dp_shard \* cp \* tp \(2\)",
        ):
            ParallelismContext(
                dp_replicate=2,
                dp_shard=2,
                cp=1,
                tp=1,
                pp=1,
                ep=3,
                world_size=4,
                enable_sequence_parallel=False,
            )


class TestDenseStorageAxes(DTensorTestBase):
    """Dense storage mesh axes exposed by ParallelismContext."""

    @property
    def world_size(self):
        return 8

    def _build(self) -> ParallelismContext:
        pd = ParallelismContext(
            dp_replicate=2,
            dp_shard=2,
            cp=1,
            tp=2,
            pp=1,
            ep=1,
            world_size=8,
            enable_sequence_parallel=False,
        )
        pd.build_mesh()
        return pd

    @with_comms
    def test_keeps_dp_shard_separate(self):
        with patch(
            "torchtitan.distributed.parallelism_context.device_type", self.device_type
        ):
            axes = self._build().get_all_one_dimensional_meshes()
            self.assertNotIn("fsdp", axes)
            self.assertIn("dp", axes)
            self.assertIn("dp_shard", axes)


class TestOneDimensionalMeshesSkipFakeAxes(DTensorTestBase):
    """get_all_one_dimensional_meshes() must not report fake-backed axes."""

    @property
    def world_size(self):
        return 8

    @with_comms
    def test_edp_shard_excluded_when_ep_disabled(self):
        """With ep=1, edp_shard is fake-backed even though its size is > 1."""
        with patch(
            "torchtitan.distributed.parallelism_context.device_type", self.device_type
        ):
            pd = ParallelismContext(
                dp_replicate=1,
                dp_shard=4,
                cp=1,
                tp=2,
                pp=1,
                ep=1,
                world_size=8,
                enable_sequence_parallel=False,
            )
            pd.build_mesh()

            # edp_shard = dp_shard * cp * tp / ep = 4 * 1 * 2 / 1 = 8, so the
            # size > 1 filter alone would let this fake-backed axis through.
            self.assertEqual(pd._single_axis_meshes["edp_shard"].size(), 8)
            self.assertIsNone(pd.get_optional_mesh("edp_shard"))

            one_d_meshes = pd.get_all_one_dimensional_meshes()
            self.assertNotIn("edp_shard", one_d_meshes)
            self.assertIn("dp", one_d_meshes)
            self.assertIn("dp_shard", one_d_meshes)
            self.assertIn("tp", one_d_meshes)
            # Every reported axis must own a usable process group.
            for name, mesh in one_d_meshes.items():
                self.assertNotEqual(
                    dist.get_backend(mesh.get_group()), "fake", f"axis {name}"
                )

    @with_comms
    def test_edp_shard_reported_when_ep_enabled(self):
        """With ep>1, edp_shard is real and must still be reported."""
        with patch(
            "torchtitan.distributed.parallelism_context.device_type", self.device_type
        ):
            pd = ParallelismContext(
                dp_replicate=1,
                dp_shard=4,
                cp=1,
                tp=2,
                pp=1,
                ep=2,
                world_size=8,
                enable_sequence_parallel=False,
            )
            pd.build_mesh()

            one_d_meshes = pd.get_all_one_dimensional_meshes()
            self.assertIn("edp_shard", one_d_meshes)
            self.assertIn("ep", one_d_meshes)
            for name, mesh in one_d_meshes.items():
                self.assertNotEqual(
                    dist.get_backend(mesh.get_group()), "fake", f"axis {name}"
                )


class TestParallelismContextWorld8MeshOperations(DTensorTestBase):
    """Test ParallelismContext mesh operations with 8-rank distributed environment."""

    @property
    def world_size(self):
        return 8

    @with_comms
    def test_world_size_8_mesh_operations(self):
        """Comprehensive test for world_size=8 mesh operations.

        This test validates mesh building, mesh retrieval, mesh sizes, and properties
        for a world_size=8 configuration with multiple parallelism dimensions enabled.
        Configuration: dp_replicate=2, dp_shard=2, cp=1, tp=2, pp=1 (2*2*1*2*1 = 8)
        """
        with patch(
            "torchtitan.distributed.parallelism_context.device_type", self.device_type
        ):
            parallelism_context = ParallelismContext(
                dp_replicate=2,
                dp_shard=2,
                cp=1,
                tp=2,
                pp=1,
                ep=1,
                world_size=8,
                enable_sequence_parallel=False,
            )

            # Test mesh building
            world_mesh = parallelism_context.build_mesh()
            self.assertIsNotNone(world_mesh)
            self.assertEqual(world_mesh.size(), 8)

            # Verify all expected meshes are created
            self.assertIsNotNone(parallelism_context._single_axis_meshes)
            self.assertIn("pp", parallelism_context._single_axis_meshes)
            self.assertIn("loss", parallelism_context._single_axis_meshes)
            self.assertIn("dp_replicate", parallelism_context._single_axis_meshes)
            self.assertIn("dp", parallelism_context._single_axis_meshes)
            self.assertIn("dp_shard", parallelism_context._single_axis_meshes)
            self.assertIn("cp", parallelism_context._single_axis_meshes)
            self.assertIn("tp", parallelism_context._single_axis_meshes)
            self.assertIn("ep", parallelism_context._single_axis_meshes)
            self.assertIn("edp_shard", parallelism_context._single_axis_meshes)

            # Validate 1D mesh sizes match parallelism configuration
            self.assertEqual(parallelism_context._single_axis_meshes["pp"].size(), 1)
            self.assertEqual(
                parallelism_context._single_axis_meshes["loss"].size(), 4
            )  # dp_replicate * dp_shard * cp = 2 * 2 * 1
            self.assertEqual(
                parallelism_context._single_axis_meshes["dp_replicate"].size(), 2
            )
            self.assertEqual(parallelism_context._single_axis_meshes["dp"].size(), 4)
            self.assertEqual(
                parallelism_context._single_axis_meshes["dp_shard"].size(), 2
            )
            self.assertEqual(parallelism_context._single_axis_meshes["cp"].size(), 1)
            self.assertEqual(parallelism_context._single_axis_meshes["tp"].size(), 2)
            self.assertEqual(parallelism_context._single_axis_meshes["ep"].size(), 1)
            self.assertEqual(
                parallelism_context._single_axis_meshes["edp_shard"].size(), 4
            )  # fsdp * tp / ep = 2 * 2 / 1 = 4

            # Validate 2D mesh shapes
            dp_replicate_fsdp_mesh = parallelism_context.get_mesh(
                ["dp_replicate", "dp_shard"]
            )
            self.assertIsNotNone(dp_replicate_fsdp_mesh)
            self.assertEqual(
                dp_replicate_fsdp_mesh.shape, (2, 2)
            )  # (dp_replicate, dp_shard)
            # edp_shard mesh only exists when ep > 1, so dp_replicate_edp_shard should be None when ep=1
            dp_replicate_edp_shard_mesh = parallelism_context.get_optional_mesh(
                ["dp_replicate", "edp_shard"]
            )
            self.assertIsNone(
                dp_replicate_edp_shard_mesh
            )  # edp_shard disabled when ep=1
            # Test get_mesh returns valid meshes for enabled dimensions (size > 1)
            self.assertIsNotNone(parallelism_context.get_mesh("tp"))
            self.assertIsNotNone(parallelism_context.get_mesh("dp_replicate"))
            self.assertIsNotNone(parallelism_context.get_mesh("dp"))
            self.assertIsNotNone(parallelism_context.get_mesh("dp_shard"))
            self.assertIsNotNone(parallelism_context.get_mesh("loss"))

            # Test get_optional_mesh returns None for disabled dimensions (size = 1)
            self.assertIsNone(parallelism_context.get_optional_mesh("pp"))
            self.assertIsNone(parallelism_context.get_optional_mesh("cp"))
            self.assertIsNone(parallelism_context.get_optional_mesh("ep"))

            # Test get_mesh with 2D mesh names
            self.assertIsNotNone(
                parallelism_context.get_mesh(["dp_replicate", "dp_shard"])
            )
            hsdp_mesh = parallelism_context.get_mesh(["dp_replicate", "dp_shard"])
            self.assertEqual(hsdp_mesh.shape, (2, 2))

            # Test get_all_one_dimensional_meshes returns only enabled meshes
            one_d_meshes = parallelism_context.get_all_one_dimensional_meshes()
            self.assertGreater(len(one_d_meshes), 0)
            # Includes the enabled data- and tensor-parallel axes.
            self.assertIn("dp_replicate", one_d_meshes)
            self.assertIn("dp", one_d_meshes)
            self.assertIn("dp_shard", one_d_meshes)
            self.assertIn("tp", one_d_meshes)
            self.assertIn("loss", one_d_meshes)
            # Should not include: pp, cp, ep (all with size = 1)
            self.assertNotIn("pp", one_d_meshes)
            self.assertNotIn("cp", one_d_meshes)
            self.assertNotIn("ep", one_d_meshes)
            # Should not include edp_shard: with ep=1 it does not exist, so it was
            # unflattened with the fake backend even though its size is 4.
            self.assertNotIn("edp_shard", one_d_meshes)

            # Test that we can get 2D meshes via get_mesh() instead
            dp_replicate_fsdp = parallelism_context.get_mesh(
                ["dp_replicate", "dp_shard"]
            )
            self.assertIsNotNone(dp_replicate_fsdp)
            self.assertEqual(dp_replicate_fsdp.ndim, 2)

            # Test world_mesh property
            world_mesh_property = parallelism_context.world_mesh
            self.assertIsNotNone(world_mesh_property)
            self.assertEqual(world_mesh_property.size(), 8)

            # Validate enabled properties
            self.assertTrue(parallelism_context.dp_enabled)
            self.assertTrue(parallelism_context.dp_replicate_enabled)
            self.assertTrue(parallelism_context.dp_shard_enabled)
            self.assertTrue(parallelism_context.fsdp_enabled)
            self.assertTrue(parallelism_context.tp_enabled)
            self.assertFalse(parallelism_context.cp_enabled)
            self.assertFalse(parallelism_context.pp_enabled)
            self.assertFalse(parallelism_context.ep_enabled)

            # Validate calculated properties
            self.assertEqual(
                parallelism_context.non_data_parallel_size, 2
            )  # cp * tp * pp = 1 * 2 * 1
            self.assertEqual(
                parallelism_context.seq_len_divisor, 4
            )  # tp * (cp * 2) = 2 * (1 * 2) = 2 * 2


class TestSingleGPUMixedPrecisionFSDP(DTensorTestBase):
    """Verify apply_fsdp on Llama3 debugmodel matches single-device reference.

    Tests that torchtitan's apply_fsdp with MixedPrecisionPolicy at degree 1
    produces numerically identical results to a reference model with manually
    cast bf16 parameters, following the pattern in
    pytorch/test/distributed/_composable/fsdp/test_fully_shard_mixed_precision.py.

    See https://github.com/pytorch/torchtitan/issues/2886
    """

    @property
    def world_size(self):
        return 1

    @with_comms
    def test_apply_fsdp_mixed_precision_single_gpu(self):
        """apply_fsdp with bf16 on Llama3 debugmodel matches manual bf16 reference on a single GPU."""
        torch.manual_seed(42)

        model_config = build_model_config("debugmodel")

        # This test runs forward+backward on self.device_type (CPU in the
        # CPU CI job). The default FlexInnerAttention backend has no CPU backward,
        # so use ScaledDotProductInnerAttention, which runs on CPU without a mask.
        from torchtitan.models.common.attention import ScaledDotProductInnerAttention

        for layer in model_config.layers:
            layer.attention.inner_attention = ScaledDotProductInnerAttention.Config()

        with torch.device("meta"):
            model = model_config.build()
        model.to_empty(device=self.device_type)
        with torch.no_grad():
            model.init_states(buffer_device=None)

        ref_model = copy.deepcopy(model)
        ref_optim = torch.optim.Adam(ref_model.parameters(), lr=1e-4)

        dp_mesh = init_device_mesh(self.device_type, (1,))
        apply_fsdp_to_decoder(
            model,
            dp_mesh,
            param_dtype=torch.bfloat16,
            reduce_dtype=torch.float32,
            pp_enabled=False,
        )
        optim = torch.optim.Adam(model.parameters(), lr=1e-4)

        # Cast only parameters to bf16, matching MixedPrecisionPolicy behavior
        # (buffers like freqs_cis stay fp32)
        ref_model_bf16 = copy.deepcopy(ref_model)
        for p in ref_model_bf16.parameters():
            p.data = p.data.to(torch.bfloat16)

        tokens = torch.randint(
            0, model_config.vocab_size, (64,), device=self.device_type
        )
        for iter_idx in range(10):
            optim.zero_grad(set_to_none=(iter_idx % 2 == 0))
            loss = model(tokens).sum()
            loss.backward()
            optim.step()

            ref_optim.zero_grad(set_to_none=(iter_idx % 2 == 0))
            ref_loss = ref_model_bf16(tokens).sum()
            ref_loss.backward()
            for p_fp32, p_bf16 in zip(
                ref_model.parameters(), ref_model_bf16.parameters()
            ):
                p_fp32.grad = p_bf16.grad.to(p_fp32.dtype)
                p_bf16.grad = None
            ref_optim.step()
            for p_fp32, p_bf16 in zip(
                ref_model.parameters(), ref_model_bf16.parameters()
            ):
                p_bf16.detach().copy_(p_fp32)

            # Validates that apply_fsdp with param_dtype=bf16 matches the manual
            # bf16 reference. Would fail if mp_policy used param_dtype=fp32 instead,
            # since the ref model runs forward/backward in bf16.
            self.assertEqual(loss, ref_loss)


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