# 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 contextlib
import subprocess
import sys
import weakref
from functools import partial
from types import SimpleNamespace
from typing import Any, cast
from unittest.mock import MagicMock, patch

import pytest
import torch
from torchtitan.components.data.types import (
    TokenizedTrainingMicrobatch,
    TrainingMicrobatch,
)
from torchtitan.components.optim import Optim
from torchtitan.distributed.cuda_graph import wrap_with_cuda_graph
from torchtitan.experiments.graph_trainer.trainer import GraphTrainingEngine
from torchtitan.observability.metrics import compute_training_performance_metrics
from torchtitan.observability.sdc_replayer import SDCReplayMismatch
from torchtitan.trainer import Trainer
from torchtitan.training_engine import ForwardBackwardResult, TrainingEngine


def test_common_imports_do_not_require_dist_moe() -> None:
    """Ordinary engine and recipe imports keep Dist-MoE optional."""
    script = r"""
import importlib.abc
import sys

class BlockDistMoe(importlib.abc.MetaPathFinder):
    def find_spec(self, fullname, path, target=None):
        if fullname == "dist_moe" or fullname.startswith("dist_moe."):
            raise ModuleNotFoundError("blocked optional import", name=fullname)
        return None

sys.meta_path.insert(0, BlockDistMoe())
import torchtitan.config.transform
import torchtitan.training_engine
import torchtitan_recipes.models.deepseek_v3 as recipes

try:
    recipes.deepseek_v3_671b_dist_moe_bf16(seq_len=128)
except ModuleNotFoundError as error:
    assert error.name == "dist_moe"
    assert "optional dist_moe package" in str(error)
else:
    raise AssertionError("Dist-MoE recipe unexpectedly loaded without its package")
"""
    result = subprocess.run(
        [sys.executable, "-c", script],
        check=False,
        capture_output=True,
        text=True,
    )
    assert result.returncode == 0, result.stderr


def _batch(
    loss_token_counts: torch.Tensor | None = None,
    routing_token_counts: torch.Tensor | None = None,
) -> TokenizedTrainingMicrobatch:
    """One microbatch produced by ``Trainer.microbatch_generator``.

    Built fresh per call because tests may mutate its model-facing dictionary.
    """
    return TokenizedTrainingMicrobatch(
        input=torch.ones(1),
        labels=torch.ones(1, dtype=torch.long),
        positions=torch.zeros(1, dtype=torch.long),
        padding_mask=torch.zeros(1, dtype=torch.bool),
        loss_token_counts=(
            torch.tensor(1) if loss_token_counts is None else loss_token_counts
        ),
        routing_token_counts=(
            torch.tensor([1]) if routing_token_counts is None else routing_token_counts
        ),
    )


class _DictTrainingMicrobatch(TrainingMicrobatch):
    def __init__(
        self,
        input_dict: dict[str, Any],
        loss_kwargs: dict[str, Any] | None = None,
    ) -> None:
        self._input_dict = input_dict
        self._loss_kwargs = loss_kwargs or {}
        self.labels = input_dict["labels"]
        self.loss_token_counts = torch.tensor(self.labels.numel())
        self.routing_token_counts = self.loss_token_counts.unsqueeze(0)
        self.to_input_dict_calls: list[tuple[torch.device | str, bool]] = []
        self.to_loss_kwargs_calls: list[tuple[torch.device | str, bool]] = []

    def as_input_dict(self) -> dict[str, Any]:
        return dict(self._input_dict)

    def to_input_dict(
        self, device: torch.device | str, *, non_blocking: bool = False
    ) -> dict[str, Any]:
        self.to_input_dict_calls.append((device, non_blocking))
        return super().to_input_dict(device, non_blocking=non_blocking)

    def loss_kwargs(self) -> dict[str, Any]:
        return dict(self._loss_kwargs)

    def to_loss_kwargs(
        self, device: torch.device | str, *, non_blocking: bool = False
    ) -> dict[str, Any]:
        self.to_loss_kwargs_calls.append((device, non_blocking))
        return super().to_loss_kwargs(device, non_blocking=non_blocking)


def _dict_microbatch(
    input_dict: dict[str, Any],
    loss_kwargs: dict[str, Any] | None = None,
) -> TrainingMicrobatch:
    return _DictTrainingMicrobatch(input_dict, loss_kwargs)


def _training_loop(trainer: TrainingEngine) -> SimpleNamespace:
    if not hasattr(trainer, "optim_step"):
        trainer.optim_step = lambda: TrainingEngine.optim_step(trainer)
    if not hasattr(trainer.config, "model"):
        trainer.config.model = SimpleNamespace(traverse=lambda _: iter(()))

    return SimpleNamespace(
        engine=trainer,
        config=trainer.config,
        gradient_accumulation_steps=trainer.gradient_accumulation_steps,
        num_pp_microbatches=trainer.num_pp_microbatches,
        metrics_processor=trainer.metrics_processor,
    )


def test_microbatch_generator_preserves_labels() -> None:
    labels = torch.ones(1, dtype=torch.long)
    microbatch = TokenizedTrainingMicrobatch(
        input=torch.ones(1),
        labels=labels,
        positions=torch.zeros(1, dtype=torch.long),
        padding_mask=torch.zeros(1, dtype=torch.bool),
        loss_token_counts=torch.tensor(1),
        routing_token_counts=torch.tensor([1]),
    )
    trainer = cast(
        Trainer,
        SimpleNamespace(
            config=SimpleNamespace(
                training=SimpleNamespace(
                    num_tokens_per_microbatch_per_dp_rank=1,
                )
            ),
            metrics_processor=SimpleNamespace(
                ntokens_since_last_log=0,
                data_loading_times=[],
            ),
        ),
    )

    output = next(Trainer.microbatch_generator(trainer, [microbatch]))

    assert output is microbatch
    assert output.labels is labels
    assert trainer.metrics_processor.ntokens_since_last_log == 1


def test_pp_forward_backward_microbatch_group_returns_sentinel_without_last_stage(
    monkeypatch,
) -> None:
    sentinel = torch.full((1,), -1.0)
    trainer = cast(
        TrainingEngine,
        SimpleNamespace(
            pp_has_last_stage=False,
            pp_schedule=SimpleNamespace(step=lambda **kwargs: None),
            parallelism_context=SimpleNamespace(
                activate_spmd=lambda **kwargs: contextlib.nullcontext(),
            ),
            config=SimpleNamespace(
                debug=SimpleNamespace(spmd_typechecking=False),
            ),
            device=torch.device("cpu"),
            _pp_loss_sentinel_on_non_last_stage=sentinel,
        ),
    )
    loss = TrainingEngine._pp_forward_backward_microbatch_group(
        trainer,
        inputs=None,
        labels=None,
        model_kwargs=[{}],
        loss_kwargs={"global_loss_token_counts": torch.tensor(1)},
    )

    assert loss is sentinel


def test_pp_forward_backward_microbatch_group_releases_consumed_loss_graphs(
    monkeypatch,
) -> None:
    activation_refs: list[weakref.ReferenceType[torch.Tensor]] = []
    loss_refs: list[weakref.ReferenceType[torch.Tensor]] = []
    loss_containers: list[list[torch.Tensor]] = []
    gradients: list[torch.Tensor] = []

    def schedule_step(**kwargs) -> None:
        loss_containers.append(kwargs["losses"])
        for value in (1.0, 2.0):
            activation = torch.tensor(value, requires_grad=True)
            loss = activation.square().view(())
            loss.backward()
            assert activation.grad is not None
            gradients.append(activation.grad.detach().clone())
            activation_refs.append(weakref.ref(activation))
            loss_refs.append(weakref.ref(loss))
            kwargs["losses"].append(loss)

    trainer = cast(
        TrainingEngine,
        SimpleNamespace(
            pp_has_last_stage=True,
            pp_schedule=SimpleNamespace(step=schedule_step),
            parallelism_context=SimpleNamespace(
                activate_spmd=lambda **kwargs: contextlib.nullcontext(),
            ),
            config=SimpleNamespace(
                debug=SimpleNamespace(spmd_typechecking=False),
            ),
            device=torch.device("cpu"),
        ),
    )
    reporting_loss = TrainingEngine._pp_forward_backward_microbatch_group(
        trainer,
        inputs=[(torch.ones(1),), (torch.ones(1),)],
        labels=[torch.ones(1), torch.ones(1)],
        model_kwargs=[{}, {}],
        loss_kwargs={"global_loss_token_counts": torch.tensor(2)},
    )

    torch.testing.assert_close(reporting_loss, torch.tensor(5.0))
    torch.testing.assert_close(torch.stack(gradients), torch.tensor([2.0, 4.0]))
    assert not reporting_loss.requires_grad
    assert reporting_loss.grad_fn is None
    assert loss_containers == [[]]
    assert all(reference() is None for reference in loss_refs)
    assert all(reference() is None for reference in activation_refs)


def test_preprocess_microbatch_groups_prepares_structured_pp_inputs(
    monkeypatch,
) -> None:
    spmd_context_active = False

    @contextlib.contextmanager
    def spmd_context(**kwargs):
        nonlocal spmd_context_active
        spmd_context_active = True
        yield
        spmd_context_active = False

    class _FakeModel:
        def preprocess_inputs(self, input_dict, **kwargs):
            assert spmd_context_active
            return (
                input_dict["input"] + 1,
                input_dict["labels"] + 2,
                {"positions": input_dict["positions"] + 3},
            )

    trainer = cast(
        TrainingEngine,
        SimpleNamespace(
            pp_has_first_stage=True,
            pp_has_last_stage=True,
            model_parts=[_FakeModel()],
            parallelism_context=SimpleNamespace(
                pp_enabled=True,
                cp=1,
                dp_replicate_enabled=False,
                activate_spmd=spmd_context,
            ),
            max_num_documents=4,
            preprocess_inputs_kwargs={},
            config=SimpleNamespace(
                parallelism="PARA",
                training=SimpleNamespace(
                    max_context_length=2048,
                    num_tokens_per_microbatch_per_dp_rank=1,
                ),
            ),
            ntokens_seen=0,
            device=torch.device("cpu"),
        ),
    )
    microbatches = [
        _dict_microbatch(
            {
                "input": torch.tensor(1),
                "positions": torch.tensor(10),
                "labels": torch.tensor([3]),
            }
        ),
        _dict_microbatch(
            {
                "input": torch.tensor(2),
                "positions": torch.tensor(20),
                "labels": torch.tensor([4]),
            }
        ),
    ]
    [(arg_mbs, kwarg_mbs, target_mbs)] = TrainingEngine._preprocess_microbatch_groups(
        trainer, [microbatches]
    )

    assert arg_mbs is not None
    assert target_mbs is not None
    torch.testing.assert_close(arg_mbs[0][0], torch.tensor(2))
    torch.testing.assert_close(arg_mbs[1][0], torch.tensor(3))
    torch.testing.assert_close(kwarg_mbs[0]["positions"], torch.tensor(13))
    torch.testing.assert_close(kwarg_mbs[1]["positions"], torch.tensor(23))
    torch.testing.assert_close(target_mbs[0], torch.tensor([5]))
    torch.testing.assert_close(target_mbs[1], torch.tensor([6]))
    assert trainer.ntokens_seen == 2
    for microbatch in microbatches:
        assert isinstance(microbatch, _DictTrainingMicrobatch)
        assert microbatch.to_input_dict_calls == [(trainer.device, True)]
        assert microbatch.to_loss_kwargs_calls == [(trainer.device, True)]


def test_preprocess_microbatch_groups_rejects_pp_loss_kwargs(monkeypatch) -> None:
    trainer = cast(
        TrainingEngine,
        SimpleNamespace(
            parallelism_context=SimpleNamespace(
                pp_enabled=True,
                cp=1,
                activate_spmd=lambda **kwargs: contextlib.nullcontext(),
            ),
            pp_has_first_stage=True,
            pp_has_last_stage=True,
            model_parts=[
                SimpleNamespace(
                    preprocess_inputs=lambda input_dict, **kwargs: (
                        input_dict["input"],
                        input_dict["labels"],
                        {},
                    )
                )
            ],
            max_num_documents=None,
            preprocess_inputs_kwargs={},
            config=SimpleNamespace(
                parallelism="PARA",
                training=SimpleNamespace(
                    max_context_length=1,
                    num_tokens_per_microbatch_per_dp_rank=1,
                ),
            ),
            device=torch.device("cpu"),
            ntokens_seen=0,
        ),
    )
    microbatch = _dict_microbatch(
        {"input": torch.tensor([1]), "labels": torch.tensor([1])},
        {"advantages": torch.tensor([0.1])},
    )
    with pytest.raises(ValueError, match="pipeline parallelism"):
        TrainingEngine._preprocess_microbatch_groups(
            trainer,
            [[microbatch]],
        )


@pytest.mark.parametrize(
    ("per_group_graph", "set_to_none"),
    [(False, True), (True, False)],
)
def test_forward_backward_runs_whole_accumulation(
    monkeypatch,
    per_group_graph: bool,
    set_to_none: bool,
) -> None:
    captured: dict[str, Any] = {}

    class _FakeModel:
        def modules(self):
            return iter(())

        def preprocess_inputs(self, input_dict, **kw):
            captured["preprocess_kwargs"] = kw
            return (
                "INPUTS",
                torch.ones(7),
                {"positions": 1, "aux_loss_denominators": None},
            )

    losses = iter((torch.tensor(1.0), torch.tensor(2.0)))

    def forward_backward_body(*, inputs, labels, model_kwargs, loss_kwargs):
        captured.setdefault("fwd_bwd_args", []).append((inputs, labels, model_kwargs))
        torch.testing.assert_close(
            loss_kwargs["global_loss_token_counts"], torch.tensor(2)
        )
        torch.testing.assert_close(loss_kwargs["advantages"], torch.tensor([0.1]))
        assert loss_kwargs["reduction"] == "sum"
        engine.loss_metrics = {"loss/mean": next(losses)}
        return engine.loss_metrics["loss/mean"]

    engine = object.__new__(TrainingEngine)
    engine.model_parts = [_FakeModel()]
    engine.max_num_documents = 4
    engine.parallelism_context = SimpleNamespace(
        pp_enabled=False,
        cp=1,
        dp_replicate_enabled=False,
        activate_spmd=contextlib.nullcontext,
    )
    engine.config = SimpleNamespace(
        parallelism=SimpleNamespace(
            fsdp_defer_gradient_reduction=False,
            fsdp_reshard_after_forward="default",
        ),
        training=SimpleNamespace(
            disable_cuda_graphs=True,
            max_context_length=2048,
            num_tokens_per_microbatch_per_dp_rank=7,
        ),
    )
    engine.preprocess_inputs_kwargs = {"processor": "VALUE"}
    engine.ntokens_seen = 100
    engine.num_completed_steps = 0
    engine.device = torch.device("cpu")
    engine.garbage_collector = SimpleNamespace(run=MagicMock())
    engine.optim = SimpleNamespace(zero_grad=MagicMock())
    engine._cuda_graph_per_accumulation_group_enabled = per_group_graph
    engine.sdc_replayer = None
    engine._non_pp_forward_backward_microbatch = forward_backward_body
    engine._run_forward_backward = partial(
        TrainingEngine._forward_backward_body,
        engine,
        defer_fsdp_gradient_reduction=False,
    )
    microbatches = [
        _dict_microbatch(
            {"input": index, "labels": torch.zeros(1)},
            {"advantages": torch.tensor([0.1]), "reduction": "sum"},
        )
        for index in range(2)
    ]
    result = TrainingEngine.forward_backward(
        engine,
        microbatch_groups=[[microbatch] for microbatch in microbatches],
        global_loss_token_counts=2,
        global_routing_token_counts=torch.tensor([2]),
    )

    torch.testing.assert_close(result.loss, torch.tensor(3.0))
    assert [metrics["loss/mean"].item() for metrics in result.loss_metrics] == [
        1.0,
        2.0,
    ]
    assert engine.num_accumulation_steps == 2
    assert engine.ntokens_seen == 114
    assert engine.loss is result.loss
    engine.garbage_collector.run.assert_called_once_with(1)
    engine.optim.zero_grad.assert_called_once_with(set_to_none=set_to_none)
    for microbatch in microbatches:
        assert isinstance(microbatch, _DictTrainingMicrobatch)
        assert microbatch.to_input_dict_calls == [(engine.device, True)]
        assert microbatch.to_loss_kwargs_calls == [(engine.device, True)]
    for inputs, labels, model_kwargs in captured["fwd_bwd_args"]:
        assert inputs == "INPUTS"
        assert model_kwargs["positions"] == 1
        torch.testing.assert_close(
            model_kwargs["aux_loss_denominators"], torch.tensor([2])
        )
        assert labels.numel() == 7
    assert captured["preprocess_kwargs"] == {
        "parallelism_context": engine.parallelism_context,
        "parallelism": engine.config.parallelism,
        "max_num_documents": 4,
        "max_context_length": 2048,
        "processor": "VALUE",
    }


def test_cuda_graph_wrapper_returns_graph_owned_output():
    class PassthroughCUDAGraphWrapper:
        def __init__(
            self,
            fn,
            example_inputs,
            *,
            num_warmup_iterations,
        ):
            self.fn = fn
            assert num_warmup_iterations == 0

        def __call__(self, *args):
            return self.fn(*args)

    graph_loss = torch.tensor(0.0)
    fwd_bwd = MagicMock(return_value=graph_loss)

    with (
        patch("torchtitan.distributed.cuda_graph.utils.device_type", "cuda"),
        patch("torch.cuda.is_available", return_value=True),
        patch.object(torch.version, "hip", None),
        patch(
            "torchtitan.distributed.cuda_graph.CUDAGraphWrapper",
            PassthroughCUDAGraphWrapper,
        ),
    ):
        runner = wrap_with_cuda_graph(fwd_bwd)
        for value in (1.0, 2.0, 3.0):
            graph_loss.fill_(value)
            loss = runner(
                torch.ones(1),
                torch.ones(1),
                torch.tensor(1),
                {"position": torch.ones(1)},
            )
            # Sanity check that the wrapper returns the same graph-owned object.
            assert loss is graph_loss

    assert fwd_bwd.call_count == 3
    _, _, global_loss_token_counts, extra_kwargs = fwd_bwd.call_args.args
    torch.testing.assert_close(global_loss_token_counts, torch.tensor(1))
    assert global_loss_token_counts.dtype == torch.int64
    torch.testing.assert_close(extra_kwargs["position"], torch.ones(1))


def test_cuda_graph_wrapper_preserves_structured_args_and_kwargs():
    class PassthroughCUDAGraphWrapper:
        def __init__(
            self,
            fn,
            example_inputs,
            *,
            num_warmup_iterations,
        ):
            self.fn = fn
            assert num_warmup_iterations == 0

        def __call__(self, *args):
            return self.fn(*args)

    fn = MagicMock(side_effect=lambda batches, *, scale: batches[1]["x"] * scale)
    with (
        patch("torchtitan.distributed.cuda_graph.utils.device_type", "cuda"),
        patch("torch.cuda.is_available", return_value=True),
        patch.object(torch.version, "hip", None),
        patch(
            "torchtitan.distributed.cuda_graph.CUDAGraphWrapper",
            PassthroughCUDAGraphWrapper,
        ),
    ):
        run = wrap_with_cuda_graph(fn)
        output = run(
            [{"x": torch.tensor(1.0)}, {"x": torch.tensor(2.0)}],
            scale=torch.tensor(3.0),
        )

    torch.testing.assert_close(output, torch.tensor(6.0))
    fn.assert_called_once()
    batches = fn.call_args.args[0]
    torch.testing.assert_close(batches[0]["x"], torch.tensor(1.0))
    torch.testing.assert_close(batches[1]["x"], torch.tensor(2.0))
    torch.testing.assert_close(fn.call_args.kwargs["scale"], torch.tensor(3.0))


def test_training_engine_configures_gradient_accumulation_cuda_graph() -> None:
    model = torch.nn.Linear(2, 2)
    eager_forward_backward = MagicMock(
        return_value=ForwardBackwardResult(torch.tensor(1.0), [])
    )
    cuda_graph_forward_backward = MagicMock(
        return_value=ForwardBackwardResult(torch.tensor(2.0), [])
    )
    engine = cast(
        TrainingEngine,
        SimpleNamespace(
            config=SimpleNamespace(
                dist_moe=None,
                sdc_replayer=None,
                debug=SimpleNamespace(spmd_typechecking=False),
                training=SimpleNamespace(
                    cuda_graph_per_accumulation_group=False,
                    disable_cuda_graphs=False,
                ),
                parallelism=SimpleNamespace(
                    enable_sequence_parallel=False,
                    fsdp_defer_gradient_reduction=True,
                    fsdp_reshard_after_forward="never",
                ),
            ),
            parallelism_context=SimpleNamespace(pp_enabled=False),
            model_parts=[model],
            _forward_backward_body=eager_forward_backward,
        ),
    )

    with (
        patch(
            "torchtitan.training_engine.wrap_fwd_bwd_with_cuda_graph",
            return_value=cuda_graph_forward_backward,
        ) as wrap,
        patch("torchtitan.training_engine.cuda_graphs_supported", return_value=True),
    ):
        TrainingEngine._initialize_forward_backward(engine)
        torch.testing.assert_close(
            engine._run_forward_backward([(), ()], torch.tensor(0)).loss,
            torch.tensor(2.0),
        )

    wrap.assert_called_once()
    assert wrap.call_args.kwargs["num_warmup_iterations"] == 2
    assert tuple(wrap.call_args.kwargs["parameters"]) == tuple(model.parameters())
    cuda_graph_forward_backward.assert_called_once()


def test_training_engine_replays_one_cuda_graph_per_accumulation_group() -> None:
    graph_loss = torch.tensor(0.0)
    graph_metric = torch.tensor(0.0)

    def run_group(
        microbatch_groups: list[tuple[int]],
        global_loss_token_counts: torch.Tensor,
    ) -> ForwardBackwardResult:
        torch.testing.assert_close(global_loss_token_counts, torch.tensor(3))
        assert len(microbatch_groups) == 1
        graph_loss.fill_(microbatch_groups[0][0])
        graph_metric.fill_(microbatch_groups[0][0])
        return ForwardBackwardResult(graph_loss, [{"group": graph_metric}])

    group_runner = MagicMock(side_effect=run_group)
    engine = cast(
        TrainingEngine,
        SimpleNamespace(
            config=SimpleNamespace(
                dist_moe=None,
                sdc_replayer=None,
                debug=SimpleNamespace(spmd_typechecking=False),
                training=SimpleNamespace(
                    cuda_graph_per_accumulation_group=True,
                    disable_cuda_graphs=False,
                ),
                parallelism=SimpleNamespace(
                    fsdp_defer_gradient_reduction=False,
                ),
            ),
            parallelism_context=SimpleNamespace(pp_enabled=False),
            _forward_backward_body=MagicMock(),
        ),
    )

    with (
        patch(
            "torchtitan.training_engine.wrap_with_cuda_graph",
            return_value=group_runner,
        ) as wrap_group,
        patch("torchtitan.training_engine.wrap_fwd_bwd_with_cuda_graph") as wrap_step,
        patch("torchtitan.training_engine.cuda_graphs_supported", return_value=True),
    ):
        TrainingEngine._initialize_forward_backward(engine)
        result = engine._run_forward_backward(
            [(1,), (2,), (3,)],
            torch.tensor(3),
        )

    torch.testing.assert_close(result.loss, torch.tensor(6.0))
    assert [metrics["group"].item() for metrics in result.loss_metrics] == [
        1.0,
        2.0,
        3.0,
    ]
    assert engine._cuda_graph_per_accumulation_group_enabled
    wrap_group.assert_called_once()
    wrap_step.assert_not_called()
    assert [call.args[0] for call in group_runner.call_args_list] == [
        [(1,)],
        [(2,)],
        [(3,)],
    ]


def test_training_engine_skips_gradient_accumulation_graph_when_unsupported() -> None:
    eager_forward_backward = MagicMock(
        return_value=ForwardBackwardResult(torch.tensor(1.0), [])
    )
    engine = cast(
        TrainingEngine,
        SimpleNamespace(
            config=SimpleNamespace(
                dist_moe=None,
                sdc_replayer=None,
                debug=SimpleNamespace(spmd_typechecking=False),
                training=SimpleNamespace(
                    cuda_graph_per_accumulation_group=False,
                    disable_cuda_graphs=False,
                ),
                parallelism=SimpleNamespace(
                    enable_sequence_parallel=False,
                    fsdp_defer_gradient_reduction=False,
                ),
            ),
            parallelism_context=SimpleNamespace(pp_enabled=False),
            _forward_backward_body=eager_forward_backward,
        ),
    )

    with (
        patch("torchtitan.training_engine.wrap_fwd_bwd_with_cuda_graph") as wrap,
        patch("torchtitan.training_engine.cuda_graphs_supported", return_value=False),
    ):
        TrainingEngine._initialize_forward_backward(engine)
        torch.testing.assert_close(
            engine._run_forward_backward([], torch.tensor(0)).loss,
            torch.tensor(1.0),
        )

    eager_forward_backward.assert_called_once()
    wrap.assert_not_called()


def test_graph_training_engine_rejects_optimizer_cuda_graph() -> None:
    config = SimpleNamespace(optim=SimpleNamespace(enable_cuda_graph=True))

    with (
        patch.object(TrainingEngine, "__init__") as init,
        pytest.raises(ValueError, match="not supported with GraphTrainer"),
    ):
        GraphTrainingEngine(
            config,
            model_config=MagicMock(),
            max_num_documents=None,
            output_dir="",
        )

    init.assert_not_called()


def test_graph_training_engine_rejects_local_compile_regions() -> None:
    config = SimpleNamespace(optim=SimpleNamespace(enable_cuda_graph=False))

    with (
        patch.object(TrainingEngine, "__init__") as init,
        pytest.raises(ValueError, match="local_compile_regions"),
    ):
        GraphTrainingEngine(
            config,
            model_config=SimpleNamespace(local_compile_regions=["loss"]),
            max_num_documents=None,
            output_dir="",
        )

    init.assert_not_called()


def test_optim_update_clips_before_parameter_update() -> None:
    events = []
    optimizers = MagicMock()
    optimizers.step.side_effect = lambda: events.append("step")
    optim = cast(
        Optim,
        SimpleNamespace(
            config=SimpleNamespace(max_norm=1.0),
            parallelism_context=SimpleNamespace(
                pp_enabled=False,
                ep_enabled=False,
                get_optional_mesh=lambda name: None,
            ),
            parameters=[],
            optimizers=optimizers,
        ),
    )

    def clip_grad_norm(*args, **kwargs):
        events.append("clip")
        return torch.tensor(2.0)

    with patch(
        "torchtitan.components.optim.optim.dist_utils.clip_grad_norm_",
        side_effect=clip_grad_norm,
    ):
        grad_norm = Optim._update(optim, torch.tensor(1.0))

    torch.testing.assert_close(grad_norm, torch.tensor(2.0))
    assert events == ["clip", "step"]


def test_initialize_optim_builds_component() -> None:
    optim = MagicMock()
    optim_config = SimpleNamespace(build=MagicMock(return_value=optim))
    engine = cast(
        TrainingEngine,
        SimpleNamespace(
            config=SimpleNamespace(
                optim=optim_config,
                training=SimpleNamespace(steps=10),
            ),
            model_parts=[MagicMock()],
            pp_has_last_stage=True,
            model_cls=SimpleNamespace(_register_optimizer_hooks=MagicMock()),
            parallelism_context=MagicMock(),
        ),
    )

    TrainingEngine._initialize_optim(engine)

    optim_config.build.assert_called_once_with(
        model_parts=engine.model_parts,
        parallelism_context=engine.parallelism_context,
        training_steps=10,
        pp_has_last_stage=True,
    )
    engine.model_cls._register_optimizer_hooks.assert_called_once_with(
        optim.optimizers,
        engine.model_parts,
        engine.parallelism_context,
    )


def test_optim_step_waits_for_checkpoint() -> None:
    events = []
    loss = torch.tensor(1.0)
    optim = SimpleNamespace(
        step=lambda actual_loss, *, current_step: (
            events.append(f"optimization_{current_step}"),
            torch.testing.assert_close(actual_loss, loss),
            torch.tensor(2.0),
        )[2]
    )
    engine = cast(
        TrainingEngine,
        SimpleNamespace(
            num_completed_steps=2,
            loss=loss,
            checkpointer=SimpleNamespace(
                maybe_wait_for_staging=lambda: events.append("checkpoint")
            ),
            optim=optim,
        ),
    )

    grad_norm = TrainingEngine.optim_step(engine)

    torch.testing.assert_close(grad_norm, torch.tensor(2.0))
    assert events == ["checkpoint", "optimization_3"]
    assert engine.num_completed_steps == 3


@pytest.mark.parametrize(
    ("loss_token_counts", "routing_token_counts"),
    (
        (torch.tensor(1), torch.tensor([1])),
        (torch.tensor([1, 1]), torch.tensor([1, 1])),
    ),
)
def test_trainer_accumulates_reused_cuda_graph_losses(
    loss_token_counts, routing_token_counts
):
    graph_loss = torch.tensor(0.0)
    loss_values = iter((1.0, 2.0, 3.0, 4.0, 5.0, 6.0))

    def forward_backward(
        *, microbatch_groups, global_loss_token_counts, global_routing_token_counts
    ):
        assert len(microbatch_groups) == 3
        torch.testing.assert_close(global_loss_token_counts, loss_token_counts * 3)
        torch.testing.assert_close(
            global_routing_token_counts, routing_token_counts * 3
        )
        graph_loss.fill_(sum(next(loss_values) for _ in microbatch_groups))
        return ForwardBackwardResult(graph_loss, [])

    metrics_processor = SimpleNamespace(
        should_log=MagicMock(return_value=True),
        step_last_log=0,
        reset=MagicMock(),
        log=MagicMock(),
    )
    trainer = cast(
        TrainingEngine,
        SimpleNamespace(
            config=SimpleNamespace(
                training=SimpleNamespace(
                    cuda_graph_per_accumulation_group=False,
                    disable_cuda_graphs=False,
                ),
            ),
            optim=SimpleNamespace(
                zero_grad=MagicMock(),
                lr_schedulers=SimpleNamespace(
                    get_metrics=MagicMock(return_value={}),
                ),
                step=MagicMock(return_value=torch.tensor(4.0)),
            ),
            parallelism_context=SimpleNamespace(
                dp_enabled=False,
                pp_enabled=False,
                dp_cp_enabled=False,
                ep_enabled=False,
                dp_replicate_enabled=False,
                get_optional_mesh=lambda name: None,
            ),
            gradient_accumulation_steps=3,
            num_pp_microbatches=1,
            device=torch.device("cpu"),
            forward_backward=forward_backward,
            loss=graph_loss,
            sdc_replayer=None,
            model_parts=[],
            model_config=SimpleNamespace(mtp_layers=None),
            checkpointer=SimpleNamespace(maybe_wait_for_staging=MagicMock()),
            metrics_processor=metrics_processor,
            num_completed_steps=0,
            ntokens_seen=3,
        ),
    )
    data_iterator = iter(
        [_batch(loss_token_counts, routing_token_counts) for _ in range(3)]
    )

    Trainer.train_step(_training_loop(trainer), data_iterator)

    metrics_processor.log.assert_called_once_with(
        1,
        6.0,
        6.0,
        4.0,
        extra_metrics={"n_tokens_seen": 3},
    )
    # The first step after loading starts a new metrics window.
    metrics_processor.reset.assert_called_once()
    assert trainer.num_completed_steps == 1

    metrics_processor.should_log.return_value = False
    metrics_processor.log.reset_mock()
    metrics_processor.reset.reset_mock()
    Trainer.train_step(
        _training_loop(trainer),
        data_iterator=iter(
            [_batch(loss_token_counts, routing_token_counts) for _ in range(3)]
        ),
    )

    metrics_processor.log.assert_not_called()
    # A step in the middle of a window does not start a new one.
    metrics_processor.reset.assert_not_called()
    assert trainer.num_completed_steps == 2


def test_engine_replay_checks_whole_accumulation() -> None:
    run_forward_backward = MagicMock(
        return_value=ForwardBackwardResult(torch.tensor(1.0), [])
    )
    replayer = SimpleNamespace(
        run_fwd_bwd=MagicMock(side_effect=lambda fn, **kwargs: fn()),
    )
    engine = object.__new__(TrainingEngine)
    engine.config = SimpleNamespace(
        training=SimpleNamespace(disable_cuda_graphs=True),
        parallelism=SimpleNamespace(
            fsdp_defer_gradient_reduction=False,
            fsdp_reshard_after_forward="default",
        ),
    )
    engine.parallelism_context = SimpleNamespace(dp_enabled=False)
    engine.model_parts = []
    engine.device = torch.device("cpu")
    engine.garbage_collector = SimpleNamespace(run=MagicMock())
    engine.optim = SimpleNamespace(zero_grad=MagicMock())
    engine._cuda_graph_per_accumulation_group_enabled = False
    engine.sdc_replayer = replayer
    engine.num_completed_steps = 0
    engine._preprocess_microbatch_groups = MagicMock(
        side_effect=lambda groups: [
            ("input", group[0].labels, {}, {}) for group in groups
        ]
    )
    engine._run_forward_backward = run_forward_backward

    TrainingEngine.forward_backward(
        engine,
        microbatch_groups=[[_batch()], [_batch()]],
        global_loss_token_counts=2,
        global_routing_token_counts=torch.tensor([2]),
    )

    replayer.run_fwd_bwd.assert_called_once()
    assert replayer.run_fwd_bwd.call_args.kwargs["step"] == 1
    assert callable(replayer.run_fwd_bwd.call_args.kwargs["get_loss"])
    run_forward_backward.assert_called_once()


def test_replay_failure_propagates_from_engine():
    mismatch = SDCReplayMismatch(
        step=1,
        local_step=1,
        replay=1,
        rank=0,
        signature_mismatch="loss",
    )
    engine = object.__new__(TrainingEngine)
    engine.config = SimpleNamespace(
        training=SimpleNamespace(disable_cuda_graphs=True),
        parallelism=SimpleNamespace(
            fsdp_defer_gradient_reduction=False,
            fsdp_reshard_after_forward="default",
        ),
    )
    engine.parallelism_context = SimpleNamespace(dp_enabled=False)
    engine.model_parts = []
    engine.device = torch.device("cpu")
    engine.garbage_collector = SimpleNamespace(run=MagicMock())
    engine.optim = SimpleNamespace(zero_grad=MagicMock())
    engine._cuda_graph_per_accumulation_group_enabled = False
    engine.sdc_replayer = SimpleNamespace(run_fwd_bwd=MagicMock(side_effect=mismatch))
    engine.num_completed_steps = 0
    engine._preprocess_microbatch_groups = MagicMock(
        return_value=[("input", "labels", {}, {})]
    )
    engine._run_forward_backward = MagicMock()

    with pytest.raises(SDCReplayMismatch):
        TrainingEngine.forward_backward(
            engine,
            microbatch_groups=[[_batch()]],
            global_loss_token_counts=torch.tensor(1),
            global_routing_token_counts=torch.tensor([1]),
        )


def test_loading_checkpoint_rearms_replay_schedule():
    replayer = SimpleNamespace(reset_schedule=MagicMock())
    trainer = cast(TrainingEngine, SimpleNamespace(sdc_replayer=replayer))

    TrainingEngine.load_state_dict(trainer, {"step": 12, "ntokens_seen": 34})

    assert trainer.num_completed_steps == 12
    assert trainer.ntokens_seen == 34
    replayer.reset_schedule.assert_called_once_with()

    disabled = cast(TrainingEngine, SimpleNamespace(sdc_replayer=None))
    TrainingEngine.load_state_dict(disabled, {"step": 1, "ntokens_seen": 2})
    assert disabled.num_completed_steps == 1


def test_seed_checkpoint_initialize_skips_forward_backward():
    events = []
    model_mem_stats = object()
    engine = cast(
        TrainingEngine,
        SimpleNamespace(
            device_memory_monitor=SimpleNamespace(
                get_peak_stats=lambda: (
                    events.append("model_memory"),
                    model_mem_stats,
                )[1]
            ),
            _initialize_model=MagicMock(
                side_effect=lambda *args, **kwargs: events.append("model")
            ),
            _initialize_optim=MagicMock(
                side_effect=lambda *args, **kwargs: events.append("optim")
            ),
            _initialize_checkpointer=MagicMock(
                side_effect=lambda *args, **kwargs: events.append("checkpointer")
            ),
            _initialize_forward_backward=MagicMock(
                side_effect=lambda **_kwargs: events.append("forward_backward")
            ),
            _dist_moe_runtime=None,
            state_dict_adapter=None,
        ),
    )
    TrainingEngine.initialize(
        engine,
        hf_assets_path="",
        create_seed_checkpoint=True,
    )

    assert events == [
        "model",
        "model_memory",
        "optim",
        "checkpointer",
    ]
    assert engine.model_device_mem_stats is model_mem_stats
    engine._initialize_model.assert_called_once_with(
        hf_assets_path="",
        create_seed_checkpoint=True,
    )
    engine._initialize_optim.assert_called_once_with()
    engine._initialize_checkpointer.assert_called_once_with(
        dataloader=None,
        sd_adapter=engine.state_dict_adapter,
    )
    engine._initialize_forward_backward.assert_not_called()


def test_compute_training_performance_metrics():
    metrics = compute_training_performance_metrics(
        num_tokens=20,
        elapsed_time=2.0,
        non_data_parallel_size=2,
        num_flops_per_token=200,
        gpu_peak_flops=1000,
        has_quantization=False,
    )

    assert metrics == {
        "tokens_per_second": 5.0,
        "tflops": 1e-9,
        "mfu_percent": 100.0,
    }


@pytest.mark.parametrize(
    ("device_type", "cuda_available", "hip_version"),
    [
        ("cpu", False, None),
        ("cuda", False, None),
        ("cuda", True, "6.3"),
        ("xpu", False, None),
    ],
)
def test_cuda_graph_wrapper_is_noop_without_nvidia_cuda(
    device_type: str,
    cuda_available: bool,
    hip_version: str | None,
) -> None:
    fwd_bwd = MagicMock()

    with (
        patch("torchtitan.distributed.cuda_graph.utils.device_type", device_type),
        patch("torch.cuda.is_available", return_value=cuda_available),
        patch.object(torch.version, "hip", hip_version),
        patch("torchtitan.distributed.cuda_graph.logger.warning") as warning,
    ):
        runner = wrap_with_cuda_graph(fwd_bwd)

    assert runner is fwd_bwd
    warning.assert_called_once()


def test_cuda_graph_accumulation_supports_eager_gradient_reduction() -> None:
    engine = cast(
        TrainingEngine,
        SimpleNamespace(
            parallelism_context=SimpleNamespace(pp_enabled=False),
            config=SimpleNamespace(
                dist_moe=None,
                training=SimpleNamespace(
                    cuda_graph_per_accumulation_group=False,
                    disable_cuda_graphs=False,
                ),
                sdc_replayer=None,
                parallelism=SimpleNamespace(
                    fsdp_defer_gradient_reduction=False,
                    fsdp_reshard_after_forward="never",
                ),
            ),
            model_parts=[],
            _forward_backward_body=MagicMock(),
        ),
    )

    with (
        patch(
            "torchtitan.training_engine.wrap_fwd_bwd_with_cuda_graph",
            side_effect=lambda fn, **_: fn,
        ),
        patch("torchtitan.training_engine.cuda_graphs_supported", return_value=True),
    ):
        TrainingEngine._initialize_forward_backward(engine)
        engine._run_forward_backward([(), ()], torch.tensor(2))

    engine._forward_backward_body.assert_called_once_with(
        [(), ()],
        torch.tensor(2),
        defer_fsdp_gradient_reduction=False,
    )


@pytest.mark.parametrize("configured_defer", [False, True])
def test_initialize_forward_backward_uses_eager_fsdp_reduction_config(
    configured_defer: bool,
) -> None:
    forward_backward_body = MagicMock(
        return_value=ForwardBackwardResult(torch.tensor(1.0), [])
    )
    engine = cast(
        TrainingEngine,
        SimpleNamespace(
            config=SimpleNamespace(
                dist_moe=None,
                training=SimpleNamespace(
                    cuda_graph_per_accumulation_group=False,
                    disable_cuda_graphs=True,
                ),
                parallelism=SimpleNamespace(
                    fsdp_defer_gradient_reduction=configured_defer,
                    fsdp_reshard_after_forward="default",
                ),
                sdc_replayer=None,
            ),
            parallelism_context=SimpleNamespace(
                pp_enabled=False,
            ),
            _forward_backward_body=forward_backward_body,
        ),
    )

    TrainingEngine._initialize_forward_backward(engine)
    engine._run_forward_backward(
        [(), ()],
        torch.tensor(2),
    )

    assert (
        forward_backward_body.call_args.kwargs["defer_fsdp_gradient_reduction"]
        is configured_defer
    )


class _RecordingFSDPPart:
    def __init__(self) -> None:
        self.requires_all_reduce_calls: list[bool] = []
        self.is_last_backward_calls: list[bool] = []
        self.reshard_after_backward_calls: list[bool] = []
        self.requires_gradient_sync_calls: list[bool] = []

    def set_requires_all_reduce(self, flag: bool, *, recurse: bool = True) -> None:
        assert recurse is True
        self.requires_all_reduce_calls.append(flag)

    def set_is_last_backward(self, flag: bool) -> None:
        self.is_last_backward_calls.append(flag)

    def set_reshard_after_backward(self, flag: bool) -> None:
        self.reshard_after_backward_calls.append(flag)

    def set_requires_gradient_sync(self, flag: bool, *, recurse: bool = True) -> None:
        assert recurse is True
        self.requires_gradient_sync_calls.append(flag)

    def parameters(self):
        return iter(())

    def preprocess_inputs(self, input_dict, **kwargs):
        return input_dict["input"], input_dict["labels"], {}


def _run_forward_backward_recording_all_reduce(
    *,
    dp_replicate_enabled: bool,
    gradient_accumulation_steps: int,
    defer_fsdp_gradient_reduction: bool = False,
) -> list[bool]:
    part = _RecordingFSDPPart()
    engine = cast(
        TrainingEngine,
        SimpleNamespace(
            parallelism_context=SimpleNamespace(
                pp_enabled=False,
                dp_replicate_enabled=dp_replicate_enabled,
            ),
            model_parts=[part],
            _non_pp_forward_backward_microbatch=MagicMock(
                return_value=torch.tensor(1.0)
            ),
        ),
    )
    TrainingEngine._forward_backward_body(
        engine,
        [("input", "labels", {}, {})] * gradient_accumulation_steps,
        torch.tensor(gradient_accumulation_steps),
        defer_fsdp_gradient_reduction=defer_fsdp_gradient_reduction,
    )
    return part.requires_all_reduce_calls


@pytest.mark.parametrize("defer_fsdp_gradient_reduction", [False, True])
def test_hsdp_skips_replicate_all_reduce_until_last_accum_group(
    defer_fsdp_gradient_reduction: bool,
) -> None:
    flags = _run_forward_backward_recording_all_reduce(
        dp_replicate_enabled=True,
        gradient_accumulation_steps=3,
        defer_fsdp_gradient_reduction=defer_fsdp_gradient_reduction,
    )
    assert flags == [False, False, True]


def test_hsdp_keeps_all_reduce_on_single_accum_group() -> None:
    flags = _run_forward_backward_recording_all_reduce(
        dp_replicate_enabled=True,
        gradient_accumulation_steps=1,
    )
    assert flags == [True]


@pytest.mark.parametrize("defer_fsdp_gradient_reduction", [False, True])
def test_pp_hsdp_skips_replicate_all_reduce_until_last_accum_group(
    defer_fsdp_gradient_reduction: bool,
) -> None:
    model_parts = [_RecordingFSDPPart(), _RecordingFSDPPart()]
    engine = cast(
        TrainingEngine,
        SimpleNamespace(
            parallelism_context=SimpleNamespace(
                pp_enabled=True,
                dp_replicate_enabled=True,
            ),
            model_parts=model_parts,
            _pp_forward_backward_microbatch_group=MagicMock(
                side_effect=(torch.tensor(1.0), torch.tensor(2.0))
            ),
        ),
    )

    TrainingEngine._forward_backward_body(
        engine,
        [(None, [{}], None)] * 2,
        torch.tensor(2),
        defer_fsdp_gradient_reduction=defer_fsdp_gradient_reduction,
    )

    for model_part in model_parts:
        assert model_part.requires_all_reduce_calls == [False, True]


def test_pure_fsdp_does_not_toggle_requires_all_reduce():
    flags = _run_forward_backward_recording_all_reduce(
        dp_replicate_enabled=False,
        gradient_accumulation_steps=3,
    )
    assert flags == []


@pytest.mark.parametrize("defer_fsdp_gradient_reduction", [False, True])
def test_fsdp_gradient_accumulation_reduction_policy(
    defer_fsdp_gradient_reduction: bool,
) -> None:
    fsdp_root = _RecordingFSDPPart()
    engine = cast(
        TrainingEngine,
        SimpleNamespace(
            parallelism_context=SimpleNamespace(
                pp_enabled=False,
                dp_replicate_enabled=False,
            ),
            model_parts=[fsdp_root],
            _non_pp_forward_backward_microbatch=MagicMock(
                side_effect=(torch.tensor(1.0), torch.tensor(2.0))
            ),
        ),
    )

    result = TrainingEngine._forward_backward_body(
        engine,
        [("input", "labels", {}, {})] * 2,
        torch.tensor(2),
        defer_fsdp_gradient_reduction=defer_fsdp_gradient_reduction,
    )

    torch.testing.assert_close(result.loss, torch.tensor(3.0))
    assert result.loss_metrics == [{}, {}]
    if defer_fsdp_gradient_reduction:
        assert fsdp_root.is_last_backward_calls == [False, True]
        assert fsdp_root.reshard_after_backward_calls == [False, True]
        assert fsdp_root.requires_gradient_sync_calls == [False, True]
    else:
        assert fsdp_root.is_last_backward_calls == []
        assert fsdp_root.reshard_after_backward_calls == []
        assert fsdp_root.requires_gradient_sync_calls == []


@pytest.mark.parametrize(
    ("defer_fsdp_gradient_reduction", "expected_finalize_gradients"),
    [(False, [True, True]), (True, [False, True])],
)
def test_pp_gradient_accumulation_finalization_policy(
    defer_fsdp_gradient_reduction: bool,
    expected_finalize_gradients: list[bool],
) -> None:
    pp_forward_backward = MagicMock(side_effect=(torch.tensor(1.0), torch.tensor(2.0)))
    engine = cast(
        TrainingEngine,
        SimpleNamespace(
            parallelism_context=SimpleNamespace(
                pp_enabled=True,
                dp_replicate_enabled=False,
            ),
            model_parts=[],
            _pp_forward_backward_microbatch_group=pp_forward_backward,
        ),
    )

    result = TrainingEngine._forward_backward_body(
        engine,
        [(None, [{}], None)] * 2,
        torch.tensor(2),
        defer_fsdp_gradient_reduction=defer_fsdp_gradient_reduction,
    )

    torch.testing.assert_close(result.loss, torch.tensor(3.0))
    assert [
        call.kwargs["finalize_gradients"] for call in pp_forward_backward.call_args_list
    ] == expected_finalize_gradients
