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

import spmd_types as spmd
import torch
import torch.distributed as dist
import torch.nn as nn
import torch.nn.functional as F
from spmd_types._checker import typecheck
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.tensor import distribute_tensor, DTensor, Replicate, Shard
from torch.distributed.tensor.parallel import loss_parallel
from torch.testing._internal.distributed._tensor.common_dtensor import (
    DTensorTestBase,
    with_comms,
)
from torchtitan.components.loss import (
    _LossParallelCrossEntropy,
    BaseLoss,
    ChunkedLossWrapper,
    compute_logprobs,
    cross_entropy_loss,
    CrossEntropyLoss,
    GradAccumulator,
    IGNORE_INDEX,
)
from torchtitan.distributed.spmd_types import set_current_spmd_mesh
from torchtitan.models.deepseek_v3.mtp import (
    get_mtp_token_counts,
    MTPDecoder,
    MTPLoss,
    roll_mtp_sequence,
)


class TestLoss(unittest.TestCase):
    def test_compute_logprobs_token_major(self):
        torch.manual_seed(42)
        logits = torch.randn(8, 16)
        labels = torch.randint(0, 16, (8,))

        logprobs, entropy = compute_logprobs(
            logits,
            labels,
            vocab_parallel_group=None,
            return_entropy=True,
            global_vocab_size=logits.shape[-1],
        )
        expected_logprobs = -F.cross_entropy(logits, labels, reduction="none")
        expected_entropy = torch.logsumexp(logits, dim=-1) - (
            torch.softmax(logits, dim=-1) * logits
        ).sum(dim=-1)

        torch.testing.assert_close(logprobs, expected_logprobs)
        torch.testing.assert_close(entropy, expected_entropy)

    def test_roll_mtp_sequence_respects_packed_document_boundaries(self):
        tokens = torch.tensor([10, 11, 12, 20, 21, 22, 23, 24])
        positions = torch.tensor([0, 1, 2, 0, 1, 2, 3, 4])

        torch.testing.assert_close(
            roll_mtp_sequence(tokens, shift=1, positions=positions, fill_value=0),
            torch.tensor([11, 12, 0, 21, 22, 23, 24, 0]),
        )
        torch.testing.assert_close(
            roll_mtp_sequence(tokens, shift=2, positions=positions, fill_value=0),
            torch.tensor([12, 0, 0, 22, 23, 24, 0, 0]),
        )

    def test_roll_mtp_sequence_requires_current_and_shifted_tokens_non_padding(self):
        tokens = torch.tensor([10, 11, 12, 13])
        positions = torch.arange(4)
        padding_mask = torch.tensor([True, False, False, False])

        shifted, valid_mask = roll_mtp_sequence(
            tokens,
            shift=1,
            positions=positions,
            padding_mask=padding_mask,
            fill_value=0,
            return_valid_mask=True,
        )

        torch.testing.assert_close(shifted, torch.tensor([0, 12, 13, 0]))
        torch.testing.assert_close(valid_mask, torch.tensor([False, True, True, False]))

    def test_mtp_target_and_routing_masks_have_distinct_semantics(self):
        labels = torch.tensor([IGNORE_INDEX, IGNORE_INDEX, 12, 13])
        positions = torch.arange(4)

        shifted_labels, routing_valid_mask = roll_mtp_sequence(
            labels,
            shift=1,
            positions=positions,
            fill_value=IGNORE_INDEX,
            return_valid_mask=True,
        )

        torch.testing.assert_close(
            shifted_labels,
            torch.tensor([IGNORE_INDEX, 12, 13, IGNORE_INDEX]),
        )
        torch.testing.assert_close(
            routing_valid_mask, torch.tensor([True, True, True, False])
        )
        self.assertEqual(int((shifted_labels != IGNORE_INDEX).sum()), 2)
        self.assertEqual(int(routing_valid_mask.sum()), 3)

    def test_mtp_preprocess_aligns_tokens_and_labels(self):
        # Two packed documents: [A0, A1, A2, B0, B1, B2, B3, B4].
        # TorchTitan positions reset at each document boundary.
        tokens = torch.tensor([10, 11, 12, 20, 21, 22, 23, 24])
        positions = torch.tensor([0, 1, 2, 0, 1, 2, 3, 4])
        labels = torch.arange(8)
        padding_mask = torch.tensor(
            [False, False, False, False, False, False, True, True]
        )
        model = _FakeMTPDecoder(skip_lm_head=True, num_mtp_layers=2)
        with patch(
            "torchtitan.models.deepseek_v3.mtp.annotate_input_spmd_types",
            side_effect=lambda _parallelism_context, batch, _input_sharding: batch,
        ):
            input_tokens, loss_labels, extra_kwargs = model.preprocess_inputs(
                {
                    "input": tokens,
                    "labels": labels,
                    "positions": positions,
                    "padding_mask": padding_mask,
                },
                parallelism_context=SimpleNamespace(cp_enabled=False),
                parallelism=SimpleNamespace(),
                max_num_documents=2,
                max_context_length=8,
            )

        assert isinstance(input_tokens, tuple)
        assert isinstance(loss_labels, tuple)
        torch.testing.assert_close(input_tokens[0], tokens)
        torch.testing.assert_close(
            input_tokens[1], torch.tensor([11, 12, 0, 21, 22, 0, 0, 0])
        )
        torch.testing.assert_close(
            input_tokens[2], torch.tensor([12, 0, 0, 22, 0, 0, 0, 0])
        )
        torch.testing.assert_close(loss_labels[0], labels)
        torch.testing.assert_close(
            loss_labels[1],
            torch.tensor(
                [1, 2, IGNORE_INDEX, 4, 5, IGNORE_INDEX, IGNORE_INDEX, IGNORE_INDEX]
            ),
        )
        torch.testing.assert_close(
            loss_labels[2],
            torch.tensor(
                [
                    2,
                    IGNORE_INDEX,
                    IGNORE_INDEX,
                    5,
                    IGNORE_INDEX,
                    IGNORE_INDEX,
                    IGNORE_INDEX,
                    IGNORE_INDEX,
                ]
            ),
        )
        self.assertEqual(len(extra_kwargs["mtp_input_valid_masks"]), 2)
        torch.testing.assert_close(
            extra_kwargs["mtp_input_valid_masks"][0],
            torch.tensor([True, True, False, True, True, False, False, False]),
        )
        torch.testing.assert_close(
            extra_kwargs["mtp_input_valid_masks"][1],
            torch.tensor([True, False, False, True, False, False, False, False]),
        )
        torch.testing.assert_close(extra_kwargs["padding_mask"], padding_mask)
        loss_token_counts, routing_token_counts = get_mtp_token_counts(
            target_mask=labels != IGNORE_INDEX,
            positions=positions,
            padding_mask=padding_mask,
            num_mtp_layers=2,
        )
        torch.testing.assert_close(loss_token_counts, torch.tensor([8, 4, 2]))
        torch.testing.assert_close(routing_token_counts, torch.tensor([6, 4, 2]))

    def test_mtp_counts_chat_targets_separately_from_routing(self):
        labels = torch.tensor([IGNORE_INDEX, IGNORE_INDEX, 12, 13])
        positions = torch.arange(4)
        padding_mask = torch.zeros(4, dtype=torch.bool)
        loss_token_counts, routing_token_counts = get_mtp_token_counts(
            target_mask=labels != IGNORE_INDEX,
            positions=positions,
            padding_mask=padding_mask,
            num_mtp_layers=1,
        )
        torch.testing.assert_close(loss_token_counts, torch.tensor([2, 2]))
        torch.testing.assert_close(routing_token_counts, torch.tensor([4, 3]))

    def test_mtp_loss_rejects_plain_tensor(self):
        loss_fn = MTPLoss(MTPLoss.Config(global_vocab_size=16))
        pred = torch.zeros(1, 4, 16)
        labels = torch.zeros(1, 4, dtype=torch.long)

        with self.assertRaisesRegex(ValueError, "expects prediction and labels tuples"):
            loss_fn(pred, labels)

    def test_mtp_loss_uses_each_depth_target_count(self):
        loss_fn = MTPLoss(MTPLoss.Config(mtp_scale=0.3))
        logits = tuple(torch.zeros(4, 4, requires_grad=True) for _ in range(3))
        labels = (
            torch.zeros(4, dtype=torch.long),
            torch.tensor([0, 0, 0, IGNORE_INDEX]),
            torch.tensor([0, 0, IGNORE_INDEX, IGNORE_INDEX]),
        )
        loss_token_counts = torch.tensor([4, 3, 2])

        loss, _ = loss_fn(logits, labels, loss_token_counts)

        torch.testing.assert_close(loss, torch.log(torch.tensor(4.0)) * 1.3)
        loss.backward()
        torch.testing.assert_close(logits[0].grad[0, 0], torch.tensor(-0.75 / 4))
        torch.testing.assert_close(logits[1].grad[0, 0], torch.tensor(-0.15 * 0.75 / 3))
        torch.testing.assert_close(logits[2].grad[0, 0], torch.tensor(-0.15 * 0.75 / 2))

    def test_mtp_loss_without_counts_returns_unnormalized_sum(self):
        loss_fn = MTPLoss(MTPLoss.Config(mtp_scale=0.3))
        logits = tuple(torch.zeros(4, 4) for _ in range(2))
        labels = tuple(torch.zeros(4, dtype=torch.long) for _ in range(2))

        loss, _ = loss_fn(logits, labels)

        torch.testing.assert_close(loss, torch.log(torch.tensor(4.0)) * 4 * 1.3)

    def test_ignore_index_equal_per_token_contribution(self):
        """Test that each valid token contributes equally to the loss.

        This test verifies that:
        1. Tokens marked with IGNORE_INDEX don't contribute to the loss
        2. Each valid token contributes the same amount to the total loss
        3. The sum-based loss calculation is correct for token normalization
        """
        torch.manual_seed(42)
        num_tokens = 32
        vocab_size = 100

        # Create predictions (logits) - same for all test cases
        predictions = torch.randn(num_tokens, vocab_size)

        # Create base labels with some tokens as IGNORE_INDEX, others valid
        # This ensures we test on the same subset of valid tokens
        labels = torch.randint(0, vocab_size, (num_tokens,))
        # Mark specific positions as IGNORE_INDEX
        labels[1] = IGNORE_INDEX
        labels[11] = IGNORE_INDEX
        labels[21] = IGNORE_INDEX
        labels[31] = IGNORE_INDEX

        # Test case 1: Compute loss on this label set
        loss1 = cross_entropy_loss(predictions, labels)
        num_loss_tokens1 = (labels != IGNORE_INDEX).sum().item()

        # Test case 2: Use the exact same predictions and labels in multiple microbatches
        # Simulating gradient accumulation with identical data
        loss2 = cross_entropy_loss(predictions, labels) + cross_entropy_loss(
            predictions, labels
        )
        num_loss_tokens2 = num_loss_tokens1 * 2

        # Per-token loss should be identical
        per_token_loss1 = loss1 / num_loss_tokens1
        per_token_loss2 = loss2 / num_loss_tokens2

        self.assertAlmostEqual(
            per_token_loss1.item(),
            per_token_loss2.item(),
            places=6,
            msg="Per-token loss should be the same across gradient accumulation steps",
        )

        # Test case 3: Verify loss scaling with token replication
        # Concatenate the sequence to create 2x the data
        predictions_doubled = torch.cat([predictions, predictions], dim=0)
        labels_doubled = torch.cat([labels, labels], dim=0)

        loss_doubled = cross_entropy_loss(predictions_doubled, labels_doubled)
        num_loss_tokens_doubled = (labels_doubled != IGNORE_INDEX).sum().item()

        per_token_loss_doubled = loss_doubled / num_loss_tokens_doubled

        self.assertAlmostEqual(
            per_token_loss1.item(),
            per_token_loss_doubled.item(),
            places=6,
            msg="Per-token loss should remain constant when scaling token count",
        )

        # Verify that total loss scales linearly with number of valid tokens
        expected_ratio = num_loss_tokens_doubled / num_loss_tokens1
        actual_ratio = loss_doubled / loss1

        self.assertAlmostEqual(
            expected_ratio,
            actual_ratio.item(),
            places=6,
            msg=f"Loss should scale linearly with valid token count. "
            f"Expected ratio: {expected_ratio}, got: {actual_ratio.item()}",
        )

    def test_ignore_index_gradient_accumulation_consistency(self):
        """Test that loss is consistent across gradient accumulation steps.

        This simulates the scenario where we have:
        - Multiple microbatches with different numbers of valid tokens
        - Total loss should equal sum of individual losses
        - Per-token loss should be consistent
        """
        torch.manual_seed(123)
        vocab_size = 100

        # Microbatch 1: eight valid tokens
        pred1 = torch.randn(8, vocab_size)
        labels1 = torch.randint(0, vocab_size, (8,))
        loss1 = cross_entropy_loss(pred1, labels1)
        tokens1 = (labels1 != IGNORE_INDEX).sum()

        # Microbatch 2: half valid tokens
        pred2 = torch.randn(8, vocab_size)
        labels2 = torch.randint(0, vocab_size, (8,))
        labels2[::2] = IGNORE_INDEX  # Mask every other token
        loss2 = cross_entropy_loss(pred2, labels2)
        tokens2 = (labels2 != IGNORE_INDEX).sum()

        # Microbatch 3: two valid tokens
        pred3 = torch.randn(8, vocab_size)
        labels3 = torch.randint(0, vocab_size, (8,))
        labels3[2:] = IGNORE_INDEX
        loss3 = cross_entropy_loss(pred3, labels3)
        tokens3 = (labels3 != IGNORE_INDEX).sum()

        # Simulate gradient accumulation: sum losses, sum tokens, then normalize
        total_loss = loss1 + loss2 + loss3
        total_tokens = tokens1 + tokens2 + tokens3
        global_avg_loss = total_loss / total_tokens

        # Verify this equals the average of individual per-token losses weighted by token count
        weighted_avg = (
            (loss1 / tokens1) * tokens1
            + (loss2 / tokens2) * tokens2
            + (loss3 / tokens3) * tokens3
        ) / total_tokens

        self.assertAlmostEqual(
            global_avg_loss.item(),
            weighted_avg.item(),
            places=5,
            msg="Global averaged loss should equal weighted average of per-token losses",
        )


class TestGradAccumulator(unittest.TestCase):
    def test_accumulate_matches_cat(self):
        """Verify GradAccumulator produces the same result as torch.cat."""
        torch.manual_seed(42)
        T, D = 16, 16
        num_chunks = 4
        reference = torch.randn(T, D)
        chunks = torch.chunk(reference, num_chunks, dim=0)

        acc = GradAccumulator(
            reference,
            num_chunks=num_chunks,
            dtype=reference.dtype,
        )
        for chunk in chunks:
            acc.add(chunk)

        result = acc.buffer
        torch.testing.assert_close(result, reference)

    def test_accumulate_with_dtype_conversion(self):
        """Verify fp32 accumulation from bf16 chunks."""
        torch.manual_seed(42)
        T, D = 16, 16
        num_chunks = 4
        reference = torch.randn(T, D)
        bf16_chunks = [c.bfloat16() for c in torch.chunk(reference, num_chunks, dim=0)]

        acc = GradAccumulator(reference, num_chunks=num_chunks, dtype=torch.float32)
        for chunk in bf16_chunks:
            acc.add(chunk)

        result = acc.buffer
        self.assertEqual(result.dtype, torch.float32)
        # Verify values match (allowing for bf16 precision loss)
        expected = torch.cat([c.float() for c in bf16_chunks], dim=0)
        torch.testing.assert_close(result, expected)

    def test_too_many_adds_raises(self):
        """Verify error when adding more chunks than expected."""
        acc = GradAccumulator(torch.randn(8, 16), num_chunks=2, dtype=torch.float32)
        acc.add(torch.randn(4, 16))
        acc.add(torch.randn(4, 16))
        with self.assertRaises(ValueError):
            acc.add(torch.randn(4, 16))


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

    @with_comms
    def test_loss_parallel_cross_entropy_parity(self):
        """
        Tests loss-parallel cross-entropy loss bitwise parity, with torch.distributed.tensor.parallel.loss_parallel().
        Tests even/uneven vocab sharding, TP, DP+TP, and IGNORE_INDEX labels.

        Runs _LossParallelCrossEntropy under typing checking: logits S(1)@TP, labels I@TP -> I@TP.
        """
        # Ensure the determinism.
        torch.use_deterministic_algorithms(True)
        torch.set_num_threads(1)

        T = 128
        mesh_configs = (
            ((4,), ("tp",), (Shard(1),), (Replicate(),)),
            ((2, 2), ("dp", "tp"), (Shard(0), Shard(1)), (Shard(0), Replicate())),
        )
        cases = ((32000, 109, False), (32003, 211, False), (32000, 307, True))

        for mesh_shape, axis_names, logits_placements, label_placements in mesh_configs:
            mesh = init_device_mesh(
                self.device_type, mesh_shape, mesh_dim_names=axis_names
            )
            tp_group = mesh.get_group("tp")
            for vocab_size, seed, ignore in cases:
                with self.subTest(
                    mesh_shape=mesh_shape, vocab_size=vocab_size, ignore=ignore
                ):
                    # Build global test data once, then let DTensor derive the
                    # exact local shards for TP-only and DP+TP placements.
                    generator = torch.Generator(device=self.device_type).manual_seed(
                        seed
                    )
                    global_logits = torch.randn(
                        T,
                        vocab_size,
                        device=self.device_type,
                        generator=generator,
                    )
                    global_labels = torch.randint(
                        0,
                        vocab_size,
                        (T,),
                        device=self.device_type,
                        dtype=torch.long,
                        generator=generator,
                    )
                    if ignore:
                        mask = torch.rand(
                            T,
                            device=self.device_type,
                            generator=generator,
                        )
                        global_labels[mask < 0.3] = IGNORE_INDEX
                    # distribute as DTensor; loss_parallel wrapper takes DTensors.
                    logits_dtensor = distribute_tensor(
                        global_logits,
                        mesh,
                        logits_placements,
                    ).detach()
                    logits_dtensor.requires_grad_(True)
                    labels_dtensor = distribute_tensor(
                        global_labels,
                        mesh,
                        label_placements,
                    )

                    # pytorch loss_parallel() as ground truth.
                    with loss_parallel():
                        wrapper_loss = F.cross_entropy(
                            logits_dtensor.float(),
                            labels_dtensor,
                            reduction="sum",
                            ignore_index=IGNORE_INDEX,
                        )

                    # typecheck S(1)@TP, I@TP -> I@TP
                    logits_type = {tp_group: spmd.S(1)}
                    labels_type = {tp_group: spmd.I}
                    if "dp" in axis_names:
                        dp_group = mesh.get_group("dp")
                        logits_type[dp_group] = spmd.S(0)
                        labels_type[dp_group] = spmd.S(0)

                    # run custom autograd
                    local_logits = (
                        logits_dtensor.to_local().detach().requires_grad_(True)
                    )
                    local_labels = labels_dtensor.to_local()
                    spmd.assert_type(local_logits, logits_type)
                    spmd.assert_type(local_labels, labels_type)
                    with typecheck(strict_mode="strict"):
                        local_loss = _LossParallelCrossEntropy.apply(
                            local_logits,
                            local_labels,
                            tp_group,
                            vocab_size,
                            "sum",
                        )
                    self.assertIs(
                        spmd.get_axis_local_type(local_loss, tp_group), spmd.I
                    )

                    # Loss and local-shard gradients must be bitwise-equivalent
                    # to torch.distributed.tensor.parallel.loss_parallel().
                    self.assertTrue(torch.equal(local_loss, wrapper_loss.to_local()))

                    # compare grad_logits; easier to wrap DTensor -> full_tensor for check.
                    with loss_parallel():
                        wrapper_loss.backward()
                    local_loss.backward()
                    local_grad_dtensor = DTensor.from_local(
                        local_logits.grad,
                        mesh,
                        logits_placements,
                        shape=torch.Size((T, vocab_size)),
                        stride=(vocab_size, 1),
                    )
                    self.assertTrue(
                        torch.equal(
                            local_grad_dtensor.full_tensor(),
                            logits_dtensor.grad.full_tensor(),
                        )
                    )

    @with_comms
    def test_vocab_parallel_policy_stats_parity(self):
        """Check no-gather logprobs, entropy, gradients, and SPMD types."""
        torch.use_deterministic_algorithms(True)
        torch.set_num_threads(1)

        T = 32
        mesh_configs = (
            ((4,), ("tp",), (Shard(1),), (Replicate(),)),
            ((2, 2), ("dp", "tp"), (Shard(0), Shard(1)), (Shard(0), Replicate())),
        )
        cases = (
            (128, torch.float32, False),
            (131, torch.bfloat16, True),
        )

        for mesh_shape, axis_names, logits_placements, label_placements in mesh_configs:
            mesh = init_device_mesh(
                self.device_type, mesh_shape, mesh_dim_names=axis_names
            )
            tp_group = mesh.get_group("tp")
            for vocab_size, dtype, ignore in cases:
                with self.subTest(
                    mesh_shape=mesh_shape,
                    vocab_size=vocab_size,
                    dtype=dtype,
                    ignore=ignore,
                ):
                    generator = torch.Generator(device=self.device_type).manual_seed(
                        1729 + vocab_size
                    )
                    global_logits = torch.randn(
                        T,
                        vocab_size,
                        device=self.device_type,
                        dtype=dtype,
                        generator=generator,
                    )
                    global_labels = torch.randint(
                        0,
                        vocab_size,
                        (T,),
                        device=self.device_type,
                        generator=generator,
                    )
                    if ignore:
                        global_labels[[1, 17]] = IGNORE_INDEX
                    # Exercise masked logits without allowing 0 * -inf to
                    # contaminate otherwise finite entropy values.
                    global_logits[3, vocab_size // 2 :] = -torch.inf
                    global_labels[3] = 0
                    global_weights = torch.linspace(
                        -1.25, 1.75, T, device=self.device_type
                    )

                    reference_logits = (
                        global_logits.detach().clone().requires_grad_(True)
                    )
                    reference_logits_fp32 = reference_logits.float()
                    reference_logprobs = -F.cross_entropy(
                        reference_logits_fp32,
                        global_labels,
                        reduction="none",
                        ignore_index=IGNORE_INDEX,
                    )
                    with torch.no_grad():
                        reference_probs = torch.softmax(reference_logits_fp32, dim=-1)
                        reference_entropy = -torch.special.xlogy(
                            reference_probs, reference_probs
                        ).sum(dim=-1)
                    (reference_logprobs * global_weights).sum().backward()

                    local_logits = (
                        distribute_tensor(global_logits, mesh, logits_placements)
                        .to_local()
                        .detach()
                        .requires_grad_(True)
                    )
                    local_labels = distribute_tensor(
                        global_labels, mesh, label_placements
                    ).to_local()
                    local_weights = distribute_tensor(
                        global_weights, mesh, label_placements
                    ).to_local()
                    expected_logprobs = distribute_tensor(
                        reference_logprobs.detach(), mesh, label_placements
                    ).to_local()
                    expected_entropy = distribute_tensor(
                        reference_entropy, mesh, label_placements
                    ).to_local()
                    expected_grad = distribute_tensor(
                        reference_logits.grad, mesh, logits_placements
                    ).to_local()

                    logits_type = {tp_group: spmd.S(1)}
                    labels_type = {tp_group: spmd.I}
                    if "dp" in axis_names:
                        dp_group = mesh.get_group("dp")
                        logits_type[dp_group] = spmd.S(0)
                        labels_type[dp_group] = spmd.S(0)
                    spmd.assert_type(local_logits, logits_type)
                    spmd.assert_type(local_labels, labels_type)

                    with set_current_spmd_mesh(mesh):
                        with typecheck(strict_mode="strict"):
                            logprobs, entropy = compute_logprobs(
                                local_logits,
                                local_labels,
                                vocab_parallel_group=tp_group,
                                return_entropy=True,
                                global_vocab_size=vocab_size,
                            )

                    self.assertIs(spmd.get_axis_local_type(logprobs, tp_group), spmd.I)
                    self.assertIs(spmd.get_axis_local_type(entropy, tp_group), spmd.I)
                    self.assertFalse(entropy.requires_grad)
                    torch.testing.assert_close(logprobs, expected_logprobs)
                    torch.testing.assert_close(entropy, expected_entropy)

                    (logprobs * local_weights).sum().backward()
                    torch.testing.assert_close(local_logits.grad, expected_grad)


class _FakeDecoder(nn.Module):
    """Minimal Decoder-like model for testing ChunkedLossWrapper."""

    def __init__(self, dim: int, vocab_size: int):
        super().__init__()
        self.output = nn.Linear(dim, vocab_size, bias=False)
        # Make it look like a Decoder to ChunkedLossWrapper
        self.layers = nn.ModuleDict()
        self.tok_embeddings = None
        self.norm = None

    def forward(self, tokens, skip_lm_head=False):
        if skip_lm_head:
            return tokens  # return hidden states directly
        return self.output(tokens)


class _IdentityAttention:
    inner_attention = None


class _IdentityDecoderBlock(nn.Module):
    def __init__(self):
        super().__init__()
        self.attention = _IdentityAttention()

    def forward(
        self,
        hidden,
        attention_metadata,
        positions,
        *,
        padding_mask=None,
        aux_loss_denominator=None,
    ):
        del attention_metadata, positions, padding_mask, aux_loss_denominator
        return hidden


class _AddMTPBlock(nn.Module):
    def __init__(self):
        super().__init__()
        self.attention = _IdentityAttention()

    def forward(
        self,
        mtp_input_embed,
        prev_embed,
        mtp_input_valid_mask,
        attention_metadata,
        positions,
        *,
        padding_mask=None,
        aux_loss_denominator=None,
    ):
        del attention_metadata, positions, padding_mask, aux_loss_denominator
        return mtp_input_embed + prev_embed * mtp_input_valid_mask.unsqueeze(-1)


class _FakeMTPDecoder(MTPDecoder):
    def __init__(self, *, skip_lm_head: bool, num_mtp_layers: int = 1):
        nn.Module.__init__(self)
        self._skip_lm_head = skip_lm_head
        self.tok_embeddings = nn.Embedding(16, 4)
        self.layers = nn.ModuleDict({"0": _IdentityDecoderBlock()})
        self.norm = nn.Identity()
        self.lm_head = nn.Linear(4, 16, bias=False)
        self.mtp_layers = nn.ModuleList(_AddMTPBlock() for _ in range(num_mtp_layers))

    def _get_attention_metadata(self, positions, **kwargs):
        del positions, kwargs
        return {}


class _WeightedTwoOutputLoss(BaseLoss):
    """Two-output objective used to exercise generic chunked-loss plumbing."""

    @dataclass(kw_only=True, slots=True)
    class Config(BaseLoss.Config):
        pass

    auxiliary_weight = 0.25

    def __init__(self, config: Config):
        del config
        self.fn = cross_entropy_loss

    def __call__(
        self,
        pred: torch.Tensor | tuple[torch.Tensor, ...],
        labels: torch.Tensor | tuple[torch.Tensor, ...],
        global_loss_token_counts: torch.Tensor | None = None,
        **loss_inputs,
    ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
        del loss_inputs
        if not isinstance(pred, tuple) or not isinstance(labels, tuple):
            raise ValueError("Expected prediction and labels tuples.")
        loss = self.fn(pred[0], labels[0]) + self.auxiliary_weight * self.fn(
            pred[1], labels[1]
        )
        if global_loss_token_counts is not None:
            loss = loss / global_loss_token_counts
        return loss, {}


class _RecordingFSDPLinear(nn.Linear):
    def __init__(self, *args, events: list[str], **kwargs):
        super().__init__(*args, **kwargs)
        self.events = events

    def set_reshard_after_forward(self, enabled):
        self.events.append(f"set_reshard_after_forward({enabled})")

    def set_reshard_after_backward(self, enabled):
        self.events.append(f"set_reshard_after_backward({enabled})")

    def set_requires_gradient_sync(self, enabled, *, recurse):
        self.events.append(f"gradient_sync({enabled})")

    def unshard(self):
        self.events.append("unshard")

    def reshard(self):
        self.events.append("reshard")

    def forward(self, input):
        self.events.append("forward")
        return super().forward(input)


class TestChunkedLossWrapper(unittest.TestCase):
    def _make_model_and_loss(self, dim=32, vocab_size=64, num_chunks=4):
        """Create a fake Decoder and ChunkedLossWrapper for testing."""
        model = _FakeDecoder(dim, vocab_size)
        chunked_loss = ChunkedLossWrapper(
            ChunkedLossWrapper.Config(num_chunks=num_chunks)
        )
        # Bypass isinstance(model, Decoder) check for unit testing
        chunked_loss.lm_head = model.output
        return model, chunked_loss

    def _torch_chunk_loss_reference(
        self,
        lm_head: nn.Module,
        hidden_states: torch.Tensor,
        labels: torch.Tensor,
        num_chunks: int,
        global_loss_token_counts: torch.Tensor | None = None,
    ):
        total_loss = hidden_states.new_zeros((), dtype=torch.float32)
        for h_chunk, label_chunk in zip(
            torch.chunk(hidden_states, num_chunks, dim=0),
            torch.chunk(labels, num_chunks, dim=0),
        ):
            chunk_loss = cross_entropy_loss(
                lm_head(h_chunk.contiguous()),
                label_chunk.contiguous(),
            )
            if global_loss_token_counts is not None:
                chunk_loss = chunk_loss / global_loss_token_counts
            total_loss = total_loss + chunk_loss.detach()
        return total_loss

    def test_chunked_loss_matches_torch_chunk_reference_for_supported_shapes(self):
        torch.manual_seed(42)
        T, D, V, num_chunks = 24, 5, 17, 3
        _model, chunked_loss = self._make_model_and_loss(D, V, num_chunks)
        hidden_states = torch.randn(T, D)
        labels = torch.randint(0, V, (T,))

        expected_loss = self._torch_chunk_loss_reference(
            chunked_loss.lm_head,
            hidden_states,
            labels,
            num_chunks,
        )
        loss, _ = chunked_loss(hidden_states, labels)

        torch.testing.assert_close(loss, expected_loss)

    def test_chunked_loss_backward_matches_torch_chunk_reference(self):
        torch.manual_seed(42)
        T, D, V, num_chunks = 24, 5, 17, 3
        model_ref, _ = self._make_model_and_loss(D, V, num_chunks)
        model_chunked, chunked_loss = self._make_model_and_loss(D, V, num_chunks)
        model_chunked.output.load_state_dict(model_ref.output.state_dict())

        hidden = torch.randn(T, D)
        labels = torch.randint(0, V, (T,))
        global_loss_token_counts = float((labels != IGNORE_INDEX).sum().item())

        def torch_chunk_loss(hidden_states):
            total = hidden_states.new_zeros((), dtype=torch.float32)
            for h_chunk, label_chunk in zip(
                torch.chunk(hidden_states, num_chunks, dim=0),
                torch.chunk(labels, num_chunks, dim=0),
            ):
                total = total + cross_entropy_loss(
                    model_ref.output(h_chunk.contiguous()),
                    label_chunk.contiguous(),
                )
            return total / global_loss_token_counts

        ref_hidden = hidden.detach().clone().requires_grad_(True)
        chunk_hidden = hidden.detach().clone().requires_grad_(True)

        ref_loss = torch_chunk_loss(ref_hidden)
        chunk_loss, _ = chunked_loss(chunk_hidden, labels, global_loss_token_counts)

        ref_loss.backward()
        chunk_loss.backward()

        torch.testing.assert_close(chunk_loss, ref_loss)
        torch.testing.assert_close(chunk_hidden.grad, ref_hidden.grad)
        torch.testing.assert_close(
            model_chunked.output.weight.grad, model_ref.output.weight.grad
        )

    def test_weighted_multi_output_matches_full_objective(self):
        torch.manual_seed(42)
        T, D, V, num_chunks = 24, 5, 17, 3
        auxiliary_weight = _WeightedTwoOutputLoss.auxiliary_weight
        model_ref, _ = self._make_model_and_loss(D, V, num_chunks)
        model_chunked, _ = self._make_model_and_loss(D, V, num_chunks)
        model_chunked.output.load_state_dict(model_ref.output.state_dict())
        chunked_loss = ChunkedLossWrapper(
            ChunkedLossWrapper.Config(
                num_chunks=num_chunks,
                loss_fn=_WeightedTwoOutputLoss.Config(),
            )
        )
        chunked_loss.set_lm_head(model_chunked.output)

        hidden = torch.randn(2, T, D)
        labels = torch.randint(0, V, (T,))
        global_loss_token_counts = (labels != IGNORE_INDEX).sum()
        ref_hidden = tuple(
            item.detach().clone().requires_grad_(True) for item in hidden
        )
        chunked_hidden = tuple(
            item.detach().clone().requires_grad_(True) for item in hidden
        )

        ref_loss = (
            cross_entropy_loss(model_ref.output(ref_hidden[0]), labels)
            + auxiliary_weight
            * cross_entropy_loss(model_ref.output(ref_hidden[1]), labels)
        ) / global_loss_token_counts
        chunked_value, _ = chunked_loss(
            chunked_hidden,
            (labels, labels),
            global_loss_token_counts,
        )
        ref_loss.backward()
        chunked_value.backward()

        torch.testing.assert_close(chunked_value, ref_loss)
        for actual, expected in zip(chunked_hidden, ref_hidden, strict=True):
            torch.testing.assert_close(actual.grad, expected.grad)
        torch.testing.assert_close(
            model_chunked.output.weight.grad,
            model_ref.output.weight.grad,
        )

    def test_mtp_decoder_returns_predictions(self):
        tokens = torch.tensor([1, 2, 3, 4])
        positions = torch.tensor([0, 1, 0, 1])
        mtp_tokens, valid_mask = roll_mtp_sequence(
            tokens,
            shift=1,
            positions=positions,
            fill_value=0,
            return_valid_mask=True,
        )
        model_inputs = (tokens, mtp_tokens)

        hidden_outputs = _FakeMTPDecoder(skip_lm_head=True)(
            model_inputs,
            positions,
            mtp_input_valid_masks=(valid_mask,),
        )
        self.assertIsInstance(hidden_outputs, tuple)
        self.assertEqual(len(hidden_outputs), 2)
        self.assertEqual(hidden_outputs[0].shape, (4, 4))
        self.assertEqual(hidden_outputs[1].shape, (4, 4))

        full_outputs = _FakeMTPDecoder(skip_lm_head=False)(
            model_inputs,
            positions,
            mtp_input_valid_masks=(valid_mask,),
        )
        self.assertEqual(full_outputs[0].shape, (4, 16))
        self.assertEqual(full_outputs[1].shape, (4, 16))

    def test_chunked_mtp_matches_full_objective(self):
        torch.manual_seed(42)
        seq_len, dim, vocab_size, num_chunks = 8, 5, 17, 2
        model_ref, _ = self._make_model_and_loss(dim, vocab_size, num_chunks)
        model_chunked, _ = self._make_model_and_loss(dim, vocab_size, num_chunks)
        model_chunked.output.load_state_dict(model_ref.output.state_dict())
        loss_config = MTPLoss.Config(
            mtp_scale=0.3,
            global_vocab_size=vocab_size,
        )
        full_loss = MTPLoss(loss_config)
        chunked_loss = ChunkedLossWrapper(
            ChunkedLossWrapper.Config(
                num_chunks=num_chunks,
                loss_fn=loss_config,
            )
        )
        chunked_loss.set_lm_head(model_chunked.output)

        labels = torch.randint(0, vocab_size, (seq_len,))
        positions = torch.tensor([0, 1, 2, 3, 0, 1, 2, 3])
        loss_labels = (
            labels,
            roll_mtp_sequence(
                labels,
                shift=1,
                positions=positions,
                fill_value=IGNORE_INDEX,
            ),
            roll_mtp_sequence(
                labels,
                shift=2,
                positions=positions,
                fill_value=IGNORE_INDEX,
            ),
        )
        hidden = tuple(torch.randn(seq_len, dim) for _ in range(3))
        reference_hidden = tuple(
            value.detach().clone().requires_grad_(True) for value in hidden
        )
        chunked_hidden = tuple(
            value.detach().clone().requires_grad_(True) for value in hidden
        )
        global_loss_token_counts = torch.stack(
            [(depth_labels != IGNORE_INDEX).sum() for depth_labels in loss_labels]
        )

        reference_value, _ = full_loss(
            tuple(model_ref.output(value) for value in reference_hidden),
            loss_labels,
            global_loss_token_counts,
        )
        chunked_value, _ = chunked_loss(
            chunked_hidden,
            loss_labels,
            global_loss_token_counts,
        )
        reference_value.backward()
        chunked_value.backward()

        torch.testing.assert_close(chunked_value, reference_value)
        for actual, expected in zip(chunked_hidden, reference_hidden, strict=True):
            torch.testing.assert_close(actual.grad, expected.grad)
        torch.testing.assert_close(
            model_chunked.output.weight.grad,
            model_ref.output.weight.grad,
        )

    def test_multi_output_fsdp_lifecycle_spans_all_terms(self):
        events: list[str] = []
        chunked_loss = ChunkedLossWrapper(
            ChunkedLossWrapper.Config(
                num_chunks=2,
                loss_fn=_WeightedTwoOutputLoss.Config(),
            )
        )
        chunked_loss.set_lm_head(_RecordingFSDPLinear(4, 8, bias=False, events=events))
        predictions = (
            torch.randn(4, 4, requires_grad=True),
            torch.randn(4, 4, requires_grad=True),
        )
        labels = (
            torch.randint(0, 8, (4,)),
            torch.randint(0, 8, (4,)),
        )

        with patch(
            "torch.distributed._composable.fsdp.FSDPModule", _RecordingFSDPLinear
        ):
            chunked_loss(predictions, labels)

        self.assertEqual(events.count("unshard"), 1)
        self.assertEqual(events.count("forward"), 4)
        self.assertEqual(events.count("gradient_sync(True)"), 1)
        self.assertEqual(events.count("reshard"), 1)

    def test_fsdp_unshards_once_before_chunk_forwards(self):
        events: list[str] = []

        class FakeFSDPLinear(nn.Linear):
            def set_reshard_after_forward(self, enabled):
                events.append(f"set_reshard_after_forward({enabled})")

            def set_reshard_after_backward(self, enabled):
                events.append(f"set_reshard_after_backward({enabled})")

            def set_requires_gradient_sync(self, enabled, *, recurse):
                events.append(f"set_requires_gradient_sync({enabled})")

            def unshard(self):
                events.append("unshard")

            def reshard(self):
                events.append("reshard")

            def forward(self, input):
                events.append("forward")
                return super().forward(input)

        chunked_loss = ChunkedLossWrapper(ChunkedLossWrapper.Config(num_chunks=2))
        chunked_loss.lm_head = FakeFSDPLinear(4, 8, bias=False)
        hidden_states = torch.randn(4, 4)
        labels = torch.randint(0, 8, (4,))

        with patch("torch.distributed._composable.fsdp.FSDPModule", FakeFSDPLinear):
            chunked_loss(hidden_states, labels)

        self.assertEqual(events.count("unshard"), 1)
        self.assertEqual(events.count("forward"), 2)
        self.assertLess(events.index("unshard"), events.index("forward"))
        self.assertEqual(events[-1], "reshard")

    def test_numerical_equivalence(self):
        """ChunkedLossWrapper must produce the same loss and gradients as the standard path."""
        torch.manual_seed(42)
        T, D, V = 16, 32, 64
        num_chunks = 4

        model_std, _ = self._make_model_and_loss(D, V, num_chunks)
        model_chunked, chunked_loss = self._make_model_and_loss(D, V, num_chunks)

        # Share the same lm_head weights
        model_chunked.output.load_state_dict(model_std.output.state_dict())

        hidden_states = torch.randn(T, D)
        labels = torch.randint(0, V, (T,))
        labels[1] = IGNORE_INDEX
        labels[11] = IGNORE_INDEX
        global_loss_token_counts = float((labels != IGNORE_INDEX).sum().item())

        # Standard path: lm_head + ce_loss + backward
        hidden_std = hidden_states.detach().requires_grad_(True)
        logits_std = model_std.output(hidden_std)
        loss_std = cross_entropy_loss(logits_std, labels)
        scaled_loss_std = loss_std / global_loss_token_counts
        scaled_loss_std.backward()
        grad_std = hidden_std.grad.clone()
        lm_head_grad_std = model_std.output.weight.grad.clone()

        # Chunked path
        hidden_chunked = hidden_states.detach().requires_grad_(True)

        loss_chunked, _ = chunked_loss(hidden_chunked, labels, global_loss_token_counts)
        loss_chunked.backward()
        grad_chunked = hidden_chunked.grad.clone()
        lm_head_grad_chunked = model_chunked.output.weight.grad.clone()

        # Verify loss values match
        torch.testing.assert_close(
            loss_chunked,
            scaled_loss_std,
            atol=1e-5,
            rtol=1e-5,
            msg="Chunked and standard loss values should match",
        )

        # Verify hidden state gradients match
        torch.testing.assert_close(
            grad_chunked.float(),
            grad_std.float(),
            atol=1e-5,
            rtol=1e-5,
            msg="Chunked and standard hidden state gradients should match",
        )

        # Verify lm_head weight gradients match
        torch.testing.assert_close(
            lm_head_grad_chunked.float(),
            lm_head_grad_std.float(),
            atol=1e-5,
            rtol=1e-5,
            msg="Chunked and standard lm_head gradients should match",
        )

    def test_different_chunk_counts(self):
        """Loss should be the same regardless of num_chunks."""
        torch.manual_seed(42)
        T, D, V = 32, 32, 64
        labels = torch.randint(0, V, (T,))
        global_loss_token_counts = float((labels != IGNORE_INDEX).sum().item())
        hidden_states = torch.randn(T, D)

        losses = []
        ref_state_dict = None
        for num_chunks in [1, 2, 4, 8]:
            model, chunked_loss = self._make_model_and_loss(D, V, num_chunks)
            # Use same lm_head weights
            if ref_state_dict is None:
                ref_state_dict = model.output.state_dict()
            else:
                model.output.load_state_dict(ref_state_dict)

            h = hidden_states.detach().requires_grad_(True)

            loss, _ = chunked_loss(h, labels, global_loss_token_counts)
            loss.backward()
            losses.append(loss.item())

        for i in range(1, len(losses)):
            self.assertAlmostEqual(
                losses[0],
                losses[i],
                places=5,
                msg=f"Loss with {2**i} chunks should match loss with 1 chunk",
            )

    def test_symbolic_seq_len_traces_chunking(self):
        from torch._dynamo.decorators import mark_unbacked
        from torch.fx.experimental.proxy_tensor import make_fx

        torch.manual_seed(42)
        T, D, V, num_chunks = 32, 8, 32, 4
        _model, chunked_loss = self._make_model_and_loss(D, V, num_chunks)
        hidden_states = torch.randn(T, D)
        labels = torch.randint(0, V, (T,))
        for tensor in (hidden_states, labels):
            mark_unbacked(
                tensor,
                0,
                hint_override=T,
                min=num_chunks,
                max=T,
                specialize_on=[lambda extent, hint=T: extent == hint],
            )

        traced = make_fx(
            chunked_loss,
            tracing_mode="symbolic",
            _allow_non_fake_inputs=True,
        )(hidden_states, labels)
        expected_loss = self._torch_chunk_loss_reference(
            chunked_loss.lm_head,
            hidden_states,
            labels,
            num_chunks,
        )
        traced_loss, _ = traced(hidden_states, labels)
        torch.testing.assert_close(traced_loss, expected_loss)

        split_nodes = [
            node
            for node in traced.graph.nodes
            if node.target is torch.ops.aten.split_with_sizes.default
        ]
        self.assertEqual(len(split_nodes), 2)

    def test_rejects_non_divisible_sequence_length(self):
        torch.manual_seed(42)
        T, D, V, num_chunks = 10, 8, 32, 4
        _model, chunked_loss = self._make_model_and_loss(D, V, num_chunks)
        hidden_states = torch.randn(T, D)
        labels = torch.randint(0, V, (T,))

        with self.assertRaisesRegex(RuntimeError, "divisible by num_chunks"):
            chunked_loss(hidden_states, labels)

    def test_single_token_chunks_match_torch_chunk_reference(self):
        torch.manual_seed(42)
        T, D, V, num_chunks = 4, 8, 32, 4
        _model, chunked_loss = self._make_model_and_loss(D, V, num_chunks)
        hidden_states = torch.randn(T, D)
        labels = torch.randint(0, V, (T,))
        global_loss_token_counts = float((labels != IGNORE_INDEX).sum().item())

        expected_loss = self._torch_chunk_loss_reference(
            chunked_loss.lm_head,
            hidden_states,
            labels,
            num_chunks,
            global_loss_token_counts,
        )
        loss, _ = chunked_loss(hidden_states, labels, global_loss_token_counts)

        torch.testing.assert_close(loss, expected_loss)


class TestChunkedLossWrapperSPMD(DTensorTestBase):
    @property
    def world_size(self):
        return 2

    @property
    def device_type(self):
        return "cpu"

    def _make_loss(
        self,
        lm_head: nn.Module,
        *,
        num_chunks=4,
    ):
        """Create the ChunkedLossWrapper variant under SPMD typecheck."""
        chunked_loss = ChunkedLossWrapper(
            ChunkedLossWrapper.Config(
                num_chunks=num_chunks,
                loss_fn=CrossEntropyLoss.Config(
                    global_vocab_size=lm_head.out_features,
                ),
            )
        )
        chunked_loss.set_lm_head(lm_head)
        return chunked_loss

    def _make_vocab_parallel_lm_head(
        self,
        dim: int,
        global_vocab_size: int,
        tp_group: dist.ProcessGroup,
    ):
        """Mimic TorchTitan's TP-vocab-parallel lm_head for ChunkedLossWrapper.

        Each rank owns only its local vocab rows. ChunkedLossWrapper consumes those
        sharded logits directly through loss-parallel CE.
        """
        tp_degree = dist.get_world_size(tp_group)
        tp_rank = dist.get_rank(tp_group)
        chunk_size = (global_vocab_size + tp_degree - 1) // tp_degree
        vocab_start = min(global_vocab_size, chunk_size * tp_rank)
        vocab_end = min(global_vocab_size, vocab_start + chunk_size)
        lm_head = nn.Linear(dim, vocab_end - vocab_start, bias=False)
        lm_head.out_features = global_vocab_size
        return lm_head

    @with_comms
    def test_spmd_matches_eager_and_types(self):
        """Check ChunkedLossWrapper numerics and SPMD types with a TP lm_head.

        The reference path runs a full-vocab ``lm_head`` followed by ordinary
        ``cross_entropy_loss``. The SPMD path uses a TP-vocab-sharded
        ``lm_head`` and loss-parallel CE. Strict typechecking must accept the
        local tensor placements, and the final loss plus hidden-state gradients
        must match the eager reference.
        """
        torch.manual_seed(42)
        T, D, V = 16, 32, 64
        num_chunks = 2
        mesh = init_device_mesh(
            self.device_type,
            (1, 1, 2),
            mesh_dim_names=("dp", "cp", "tp"),
        )
        _, _, tp_rank = mesh.get_coordinate()
        _, _, tp_degree = mesh.shape

        hidden_states = torch.randn(T, D, device=self.device_type)
        labels = torch.randint(0, V, (T,), device=self.device_type)
        labels[1] = IGNORE_INDEX
        labels[11] = IGNORE_INDEX

        tp_group = mesh.get_group("tp")
        # create full-weight ref lm_head, sharded lm_head & ChunkedLossWrapper
        lm_head_ref = nn.Linear(D, V, bias=False).to(self.device_type)
        lm_head_spmd = self._make_vocab_parallel_lm_head(D, V, tp_group).to(
            self.device_type
        )
        loss_spmd_fn = self._make_loss(lm_head_spmd, num_chunks=num_chunks)

        # copy over vocab shard
        chunk_size = (V + tp_degree - 1) // tp_degree
        vocab_start = min(V, chunk_size * tp_rank)
        vocab_end = min(V, vocab_start + chunk_size)
        lm_head_spmd.weight.data.copy_(lm_head_ref.weight.data[vocab_start:vocab_end])

        # run reference path for loss, grad
        h_ref = hidden_states.clone().detach().requires_grad_(True)
        loss_ref = cross_entropy_loss(lm_head_ref(h_ref), labels)
        loss_ref.backward()
        h_grad_ref = h_ref.grad.clone()

        # run SPMD path, typecheck
        h_spmd = hidden_states.clone().detach().requires_grad_(True)
        with set_current_spmd_mesh(mesh):
            spmd.assert_type(h_spmd, {tp_group: spmd.R})
            spmd.assert_type(labels, {tp_group: spmd.I})
            spmd.assert_type(lm_head_spmd.weight, {tp_group: spmd.S(0)})
            with typecheck(strict_mode="strict", local=False):
                loss_spmd, _ = loss_spmd_fn(h_spmd, labels)
            # ChunkedLossWrapper returns through an autograd bridge under
            # no_typecheck, so the returned scalar is intentionally untyped.

        # numerics check
        loss_spmd.backward()
        torch.testing.assert_close(loss_spmd, loss_ref)
        h_grad_spmd = h_spmd.grad.clone()
        dist.all_reduce(h_grad_spmd, group=tp_group)
        torch.testing.assert_close(h_grad_spmd, h_grad_ref)


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