# 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.

from contextlib import nullcontext
from typing import cast
from unittest.mock import MagicMock, patch

import pytest
import torch
from torchtitan.distributed.cuda_graph import (
    _ForwardBackwardCUDAGraphWrapper,
    _manager,
    CUDAGraphWrapper,
    get_cuda_graph_annotations,
    run_eager_on_cuda_graph_stream,
    wrap_with_cuda_graph,
)


def test_cuda_graph_wrapper_uses_configured_warmup_iterations() -> None:
    with (
        patch.object(_manager, "maybe_initialize"),
        patch.object(_manager, "register"),
    ):
        wrapper = CUDAGraphWrapper(
            lambda value: value,
            (torch.tensor(1),),
            num_warmup_iterations=2,
        )

    assert wrapper._warmup_remaining == 2


def test_cuda_graph_wrapper_rejects_negative_warmup_iterations() -> None:
    with pytest.raises(ValueError, match="must be non-negative"):
        CUDAGraphWrapper(
            lambda value: value,
            (torch.tensor(1),),
            num_warmup_iterations=-1,
        )


def test_run_eager_on_cuda_graph_stream_synchronizes_streams() -> None:
    current_stream = MagicMock()
    graph_stream = MagicMock()
    fn = MagicMock(return_value="output")

    with (
        patch.object(_manager, "maybe_initialize") as maybe_initialize,
        patch.object(_manager, "_stream", graph_stream),
        patch("torch.cuda.current_stream", return_value=current_stream),
        patch("torch.cuda.stream", return_value=nullcontext()) as use_stream,
    ):
        output = run_eager_on_cuda_graph_stream(fn, "arg", keyword="value")

    assert output == "output"
    maybe_initialize.assert_called_once_with()
    graph_stream.wait_stream.assert_called_once_with(current_stream)
    use_stream.assert_called_once_with(graph_stream)
    fn.assert_called_once_with("arg", keyword="value")
    current_stream.wait_stream.assert_called_once_with(graph_stream)


def test_wrap_with_cuda_graph_captures_first_invocation() -> None:
    graph = MagicMock()
    fn = MagicMock(side_effect=lambda value: value)
    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.object(_manager, "maybe_initialize"),
        patch.object(_manager, "register"),
        patch.object(_manager, "_graph_pool", object()),
        patch.object(_manager, "_stream", MagicMock()),
        patch("torch.cuda.current_stream", return_value=MagicMock()),
        patch("torch.cuda.stream", return_value=nullcontext()),
        patch("torch.cuda.CUDAGraph", return_value=graph) as graph_constructor,
        patch("torch.cuda.graph", return_value=nullcontext()),
        patch(
            "torchtitan.distributed.cuda_graph.get_kernel_annotations",
            return_value={},
        ),
    ):
        run = wrap_with_cuda_graph(fn)
        value = torch.tensor(1.0)
        run(value)
        graph_constructor.assert_called_once()
        graph.replay.assert_called_once()
        run(value)
        assert fn.call_count == 1
        assert graph.replay.call_count == 2


def test_tensor_input_indices_control_replay_copies() -> None:
    static_input = torch.tensor(1)
    excluded_input = torch.tensor(2)
    copied_input = torch.tensor(3)

    with (
        patch.object(_manager, "maybe_initialize"),
        patch.object(_manager, "register"),
    ):
        wrapper = CUDAGraphWrapper(
            lambda *args: args,
            (static_input, excluded_input, copied_input),
            static_input_indices=(0,),
            tensor_input_indices=[0, 2],
            num_warmup_iterations=0,
        )

    wrapper._args = (static_input, excluded_input, copied_input)
    graph = cast(torch.cuda.CUDAGraph, MagicMock())
    wrapper._graph = graph
    wrapper._output = "output"

    result = wrapper(torch.tensor(4), torch.tensor(5), torch.tensor(6))

    assert result == "output"
    assert static_input.item() == 1
    assert excluded_input.item() == 2
    assert copied_input.item() == 6
    cast(MagicMock, graph.replay).assert_called_once_with()


def test_cuda_graph_wrapper_collects_annotations() -> None:
    graph = cast(torch.cuda.CUDAGraph, MagicMock())
    annotations = {42: [{"module_fqn": "layers.0"}]}
    graph_pool = object()
    stream = MagicMock()

    with (
        patch.object(_manager, "maybe_initialize"),
        patch.object(_manager, "register"),
        patch.object(_manager, "_graph_pool", graph_pool),
        patch.object(_manager, "_stream", stream),
        patch.object(_manager, "all_annotations", {}),
        patch("torch.cuda.CUDAGraph", return_value=graph),
        patch("torch.cuda.graph", return_value=nullcontext()) as cuda_graph,
        patch(
            "torchtitan.distributed.cuda_graph.get_kernel_annotations",
            return_value=annotations,
        ),
    ):
        wrapper = CUDAGraphWrapper(
            lambda x: x,
            (torch.tensor(1),),
            num_warmup_iterations=0,
        )

        output = wrapper(torch.tensor(2))

        assert output.item() == 2
        assert get_cuda_graph_annotations() == annotations
        cuda_graph.assert_called_once_with(
            graph,
            pool=graph_pool,
            stream=stream,
            enable_annotations=True,
            capture_error_mode="thread_local",
        )


def test_structured_wrapper_validates_and_copies_replay_inputs() -> None:
    graph = cast(torch.cuda.CUDAGraph, MagicMock())
    graph_stream = MagicMock()
    current_stream = MagicMock()
    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.object(_manager, "maybe_initialize"),
        patch.object(_manager, "register"),
        patch.object(_manager, "_graph_pool", object()),
        patch.object(_manager, "_stream", graph_stream),
        patch("torch.cuda.current_stream", return_value=current_stream),
        patch("torch.cuda.stream", return_value=nullcontext()),
        patch("torch.cuda.CUDAGraph", return_value=graph),
        patch("torch.cuda.graph", return_value=nullcontext()),
        patch(
            "torchtitan.distributed.cuda_graph.get_kernel_annotations",
            return_value={},
        ),
    ):
        run = wrap_with_cuda_graph(fn)
        torch.testing.assert_close(
            run(
                [{"x": torch.tensor(4.0)}, {"x": torch.tensor(5.0)}],
                scale=torch.tensor(4.0),
            ),
            torch.tensor(20.0),
        )
        torch.testing.assert_close(
            run(
                [{"x": torch.tensor(6.0)}, {"x": torch.tensor(7.0)}],
                scale=torch.tensor(5.0),
            ),
            torch.tensor(20.0),
        )

        with pytest.raises(ValueError, match="structure must remain constant"):
            run([{"x": torch.tensor(1.0)}], scale=torch.tensor(1.0))
        with pytest.raises(ValueError, match="same shape, dtype, and device"):
            run(
                [{"x": torch.ones(2)}, {"x": torch.tensor(1.0)}],
                scale=torch.tensor(1.0),
            )

    assert fn.call_count == 1
    captured_batches = fn.call_args.args[0]
    torch.testing.assert_close(captured_batches[0]["x"], torch.tensor(6.0))
    torch.testing.assert_close(captured_batches[1]["x"], torch.tensor(7.0))
    torch.testing.assert_close(fn.call_args.kwargs["scale"], torch.tensor(5.0))
    assert cast(MagicMock, graph.replay).call_count == 2


def test_cuda_graph_wrapper_restores_capture_allocated_gradients() -> None:
    parameter = torch.nn.Parameter(torch.ones(2))
    frozen_parameter = torch.nn.Parameter(torch.ones(2), requires_grad=False)
    captured_gradient = torch.full_like(parameter, 3.0)
    graph = cast(torch.cuda.CUDAGraph, MagicMock())

    def forward_backward(value):
        parameter.grad = captured_gradient
        return value

    with (
        patch.object(_manager, "maybe_initialize"),
        patch.object(_manager, "register"),
        patch.object(_manager, "_graph_pool", object()),
        patch.object(_manager, "_stream", MagicMock()),
        patch("torch.cuda.CUDAGraph", return_value=graph),
        patch("torch.cuda.graph", return_value=nullcontext()),
        patch(
            "torchtitan.distributed.cuda_graph.get_kernel_annotations",
            return_value={},
        ),
    ):
        wrapper = _ForwardBackwardCUDAGraphWrapper(
            forward_backward,
            (torch.tensor(1.0),),
            num_warmup_iterations=0,
            parameters=(parameter, parameter, frozen_parameter),
        )

        wrapper(torch.tensor(2.0))
        assert parameter.grad is captured_gradient
        assert frozen_parameter.grad is None

        parameter.grad = None
        wrapper(torch.tensor(3.0))
        assert parameter.grad is captured_gradient

        wrapper.teardown()
        assert parameter.grad is None

    assert cast(MagicMock, graph.replay).call_count == 2


@pytest.mark.parametrize("capture_first", [False, True])
def test_cuda_graph_wrapper_requires_cleared_owned_gradients(
    capture_first: bool,
) -> None:
    parameter = torch.nn.Parameter(torch.ones(2))
    graph = cast(torch.cuda.CUDAGraph, MagicMock())

    def forward_backward(value):
        parameter.grad = torch.ones_like(parameter)
        return value

    with (
        patch.object(_manager, "maybe_initialize"),
        patch.object(_manager, "register"),
        patch.object(_manager, "_graph_pool", object()),
        patch.object(_manager, "_stream", MagicMock()),
        patch("torch.cuda.CUDAGraph", return_value=graph),
        patch("torch.cuda.graph", return_value=nullcontext()),
        patch(
            "torchtitan.distributed.cuda_graph.get_kernel_annotations",
            return_value={},
        ),
    ):
        wrapper = _ForwardBackwardCUDAGraphWrapper(
            forward_backward,
            (torch.tensor(1.0),),
            num_warmup_iterations=0,
            parameters=(parameter,),
        )

        if capture_first:
            wrapper(torch.tensor(2.0))
        parameter.grad = torch.ones_like(parameter)
        with pytest.raises(RuntimeError, match="must be None"):
            wrapper(torch.tensor(3.0))
