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

"""Lightweight CUDA graph wrapper for training steps."""

import logging
import warnings
from collections.abc import Callable, Iterable, Sequence
from dataclasses import dataclass
from typing import Any

import torch
from torch.cuda._graph_annotations import get_kernel_annotations
from torch.nn.attention.flex_attention import BlockMask
from torch.utils import _pytree as pytree

from torchtitan.tools import utils


logger = logging.getLogger(__name__)

NUM_CUDA_GRAPH_WARMUP_STEPS = 2


@dataclass(frozen=True)
class _BlockMaskInputSpec:
    num_leaves: int
    context: tuple[Any, ...]


# TODO(@jinsooihm): Remove this class and use standard pytree flattening after
# attention mask creation moves into model code and BlockMask is no longer an input.
class CUDAGraphInputSpec:
    """Flatten structured inputs while exposing tensors stored in ``BlockMask``."""

    def __init__(self, tree: Any) -> None:
        outer_leaves, self._tree_spec = pytree.tree_flatten(
            tree,
            is_leaf=lambda value: isinstance(value, BlockMask),
        )
        self._leaf_specs: list[_BlockMaskInputSpec | None] = []
        self._num_flat_leaves = 0
        for leaf in outer_leaves:
            if isinstance(leaf, BlockMask):
                block_mask_leaves, context = leaf._flatten()
                self._leaf_specs.append(
                    _BlockMaskInputSpec(len(block_mask_leaves), context)
                )
                self._num_flat_leaves += len(block_mask_leaves)
            else:
                self._leaf_specs.append(None)
                self._num_flat_leaves += 1

    def flatten(self, tree: Any) -> list[Any]:
        outer_leaves, tree_spec = pytree.tree_flatten(
            tree,
            is_leaf=lambda value: isinstance(value, BlockMask),
        )
        if tree_spec != self._tree_spec or len(outer_leaves) != len(self._leaf_specs):
            raise ValueError(
                "CUDA graph input structure must remain constant across "
                "training steps."
            )

        flat_leaves: list[Any] = []
        for leaf, leaf_spec in zip(outer_leaves, self._leaf_specs, strict=True):
            if leaf_spec is None:
                if isinstance(leaf, BlockMask):
                    raise ValueError(
                        "CUDA graph input structure must remain constant "
                        "across training steps."
                    )
                flat_leaves.append(leaf)
                continue

            if not isinstance(leaf, BlockMask):
                raise ValueError(
                    "CUDA graph input structure must remain constant "
                    "across training steps."
                )
            block_mask_leaves, context = leaf._flatten()
            if (
                len(block_mask_leaves) != leaf_spec.num_leaves
                or context != leaf_spec.context
            ):
                raise ValueError(
                    "CUDA graph BlockMask structure must remain constant across "
                    "training steps."
                )
            flat_leaves.extend(block_mask_leaves)

        return flat_leaves

    def unflatten(self, flat_leaves: Sequence[Any]) -> Any:
        if len(flat_leaves) != self._num_flat_leaves:
            raise ValueError(
                f"CUDA graph expected {self._num_flat_leaves} inputs, "
                f"got {len(flat_leaves)}."
            )

        outer_leaves: list[Any] = []
        flat_index = 0
        for leaf_spec in self._leaf_specs:
            if leaf_spec is None:
                outer_leaves.append(flat_leaves[flat_index])
                flat_index += 1
                continue

            block_mask_end = flat_index + leaf_spec.num_leaves
            block_mask_leaves = tuple(flat_leaves[flat_index:block_mask_end])
            outer_leaves.append(
                BlockMask._unflatten(block_mask_leaves, leaf_spec.context)
            )
            flat_index = block_mask_end

        return pytree.tree_unflatten(outer_leaves, self._tree_spec)


class _CUDAGraphGradientState:
    """Keep capture-created parameter gradients alive across graph replays.

    Why this state can reduce memory use:

    1. ``set_to_none=True`` without this state fails after the
       first capture.
       Autograd allocates a gradient during capture and assigns it to
       ``param.grad``. The next ``zero_grad`` clears that Python reference.
       CUDA graph replay runs the captured GPU work, but it does not repeat the
       ``param.grad`` assignment. The optimizer then sees ``param.grad is None``
       and skips the parameter update.

    2. ``set_to_none=False`` avoids that failure, but eager warmup allocates the
       gradient before capture, outside the CUDA graph memory pool. Capture only
       records writes into that existing gradient buffer. Its lifetime overlaps
       the complete forward-backward graph::

           time           forward             backward          optimizer
           gradient       [===============================================]
           forward temp   [==========]
           backward temp                       [=========]

       The graph allocator cannot reuse the gradient storage for these temporary
       tensors because their lifetimes overlap.

    3. With this state, ``set_to_none=True`` leaves gradients
       absent when capture starts. Autograd allocates them from the CUDA graph
       memory pool during backward, after many forward tensors have expired::

           time           forward             backward          optimizer
           forward temp   [==========]
           gradient                            [=========================]
           shared block   [forward temp]       [gradient                 ]

       During capture, the CUDA graph pool may reuse storage released by forward
       tensors for later gradients. This state then keeps those gradient tensors
       alive and restores each ``param.grad`` reference after replay. This makes
       ``set_to_none=True`` safe.

    Args:
        parameters: Parameters whose capture-created ``.grad`` tensors are
            recorded and restored. Parameters that do not require gradients and
            duplicate parameters are ignored.
    """

    def __init__(self, parameters: Iterable[torch.nn.Parameter]) -> None:
        unique_parameters: list[torch.nn.Parameter] = []
        seen: set[int] = set()
        for parameter in parameters:
            if parameter.requires_grad and id(parameter) not in seen:
                seen.add(id(parameter))
                unique_parameters.append(parameter)
        self._parameters = tuple(unique_parameters)
        self._captured_parameter_gradients: tuple[
            tuple[torch.nn.Parameter, torch.Tensor], ...
        ] | None = None

    def require_cleared(self) -> None:
        if any(parameter.grad is not None for parameter in self._parameters):
            raise RuntimeError(
                "All parameter gradients must be None before CUDA graph capture "
                "and replay."
            )

    def record(self) -> None:
        assert (
            self._captured_parameter_gradients is None
        ), "CUDA graph gradients were already recorded"
        self._captured_parameter_gradients = tuple(
            (parameter, gradient)
            for parameter in self._parameters
            if (gradient := parameter.grad) is not None
        )

    def restore(self) -> None:
        assert (
            self._captured_parameter_gradients is not None
        ), "CUDA graph gradients were not recorded"
        for parameter, gradient in self._captured_parameter_gradients:
            parameter.grad = gradient

    def teardown(self) -> None:
        if self._captured_parameter_gradients is None:
            return
        for parameter, gradient in self._captured_parameter_gradients:
            if parameter.grad is gradient:
                parameter.grad = None
        self._captured_parameter_gradients = None


class _CUDAGraphManager:
    """Singleton that owns a shared graph pool, stream, and annotations."""

    def __init__(self) -> None:
        self._initialized = False
        self._wrappers: list["CUDAGraphWrapper"] = []
        self._graph_pool: Any = None
        self._stream: torch.cuda.Stream | None = None
        self._dummy_graph: torch.cuda.CUDAGraph | None = None
        self.all_annotations: dict[int, list[Any]] = {}

    @property
    def graph_pool(self) -> Any:
        assert self._graph_pool is not None
        return self._graph_pool

    @property
    def stream(self) -> torch.cuda.Stream:
        assert self._stream is not None
        return self._stream

    def maybe_initialize(self) -> None:
        if self._initialized:
            return
        graph_pool = torch.cuda.graph_pool_handle()
        stream = torch.cuda.Stream()
        dummy_graph = torch.cuda.CUDAGraph()
        self._graph_pool = graph_pool
        self._stream = stream
        self._dummy_graph = dummy_graph
        with (
            warnings.catch_warnings(record=True),
            torch.cuda.graph(
                dummy_graph,
                pool=graph_pool,
                stream=stream,
                capture_error_mode="thread_local",
            ),
        ):
            pass
        self._initialized = True

    def register(self, wrapper: "CUDAGraphWrapper") -> None:
        self._wrappers.append(wrapper)

    def teardown(self) -> None:
        if not self._initialized:
            return
        for wrapper in self._wrappers:
            wrapper.teardown()
        self._wrappers.clear()
        self._dummy_graph = None
        self._stream = None
        self._graph_pool = None
        self._initialized = False


_manager = _CUDAGraphManager()


def cuda_graph_teardown() -> None:
    """Destroy all CUDA graphs and release the shared memory pool."""
    _manager.teardown()


def get_cuda_graph_annotations() -> dict[int, list[Any]]:
    """Return all kernel annotations accumulated across CUDA graph captures."""
    return _manager.all_annotations


def run_eager_on_cuda_graph_stream(
    fn: Callable[..., Any], *args: Any, **kwargs: Any
) -> Any:
    """Run a callable eagerly on the stream reserved for CUDA graph capture."""
    _manager.maybe_initialize()
    current_stream = torch.cuda.current_stream()
    _manager.stream.wait_stream(current_stream)
    with torch.cuda.stream(_manager.stream):
        output = fn(*args, **kwargs)
    current_stream.wait_stream(_manager.stream)
    return output


class CUDAGraphWrapper:
    """Wrap a callable with CUDA graph capture and replay.

    Args:
        fn: The callable (forward+backward step) to wrap.
        example_inputs: Inputs that define the fixed input structure and tensor
            metadata for capture and replay.
        static_input_indices: Indices of inputs whose tensor addresses
            are stable across calls (e.g. model weights/buffers).
        should_check_address: Whether to verify static input tensor addresses
            before each replay. This should only be enabled for debugging.
        tensor_input_indices: Indices of inputs that should be copied before
            replay. When omitted, these are inferred from ``example_inputs``.
        num_warmup_iterations: Number of eager invocations before capture.
    Raises:
        ValueError: If ``num_warmup_iterations`` is negative.
    """

    def __init__(
        self,
        fn: Callable,
        example_inputs: Sequence[Any],
        static_input_indices: Sequence[int] | None = None,
        should_check_address: bool = False,
        tensor_input_indices: Sequence[int] | None = None,
        *,
        num_warmup_iterations: int,
    ):
        if num_warmup_iterations < 0:
            raise ValueError("num_warmup_iterations must be non-negative")
        self._fn = fn
        self._num_inputs = len(example_inputs)
        self._static_input_indices = set(static_input_indices or ())
        invalid_static_indices = {
            i for i in self._static_input_indices if i < 0 or i >= self._num_inputs
        }
        if invalid_static_indices:
            raise ValueError(
                "CUDA graph static input indices are out of range: "
                f"{sorted(invalid_static_indices)}"
            )

        if tensor_input_indices is not None:
            self._input_indices_to_copy = [
                i for i in tensor_input_indices if i not in self._static_input_indices
            ]
        else:
            self._input_indices_to_copy = [
                i
                for i, inp in enumerate(example_inputs)
                if isinstance(inp, torch.Tensor) and i not in self._static_input_indices
            ]
        self._tensor_metadata = {
            i: (inp.shape, inp.dtype, inp.device)
            for i, inp in enumerate(example_inputs)
            if isinstance(inp, torch.Tensor)
        }
        self._non_tensor_inputs = {
            i: inp
            for i, inp in enumerate(example_inputs)
            if not isinstance(inp, torch.Tensor)
        }
        self._graph: torch.cuda.CUDAGraph | None = None
        self._warmup_remaining = num_warmup_iterations
        self._args: tuple | None = None
        self._output: Any = None
        self._should_check_address = should_check_address
        self._static_input_addresses: dict[int, int] = {}

        _manager.maybe_initialize()
        _manager.register(self)

    def _record_static_input_addresses(self, args: tuple[Any, ...]) -> None:
        for i in self._static_input_indices:
            arg = args[i]
            if isinstance(arg, torch.Tensor):
                self._static_input_addresses[i] = arg.data_ptr()

    def _check_static_input_addresses(self, args: tuple[Any, ...]) -> None:
        for i, expected in self._static_input_addresses.items():
            arg = args[i]
            assert isinstance(
                arg, torch.Tensor
            ), f"Static input at index {i} changed from a tensor to {type(arg)}"
            actual = arg.data_ptr()
            assert expected == actual, (
                "Expected the same static tensor address at index "
                f"{i}, but found {expected} != {actual}"
            )

    def _validate_inputs(self, args: tuple[Any, ...]) -> None:
        if len(args) != self._num_inputs:
            raise ValueError(
                f"CUDA graph expected {self._num_inputs} inputs, got {len(args)}"
            )

        for i, expected_metadata in self._tensor_metadata.items():
            arg = args[i]
            if not isinstance(arg, torch.Tensor):
                raise ValueError(
                    f"CUDA graph input {i} changed from a tensor to {type(arg)}"
                )
            actual_metadata = (arg.shape, arg.dtype, arg.device)
            if actual_metadata != expected_metadata:
                raise ValueError(
                    "CUDA graph tensor inputs must keep the same shape, dtype, "
                    f"and device, but input {i} changed from "
                    f"{expected_metadata} to {actual_metadata}"
                )

        for i, expected in self._non_tensor_inputs.items():
            actual = args[i]
            if type(actual) is not type(expected) or actual != expected:
                raise ValueError(
                    "CUDA graph non-tensor inputs must remain constant, but input "
                    f"{i} changed from {expected!r} to {actual!r}"
                )

    def __call__(self, *args):
        self._validate_inputs(args)

        if self._warmup_remaining > 0:
            self._warmup_remaining -= 1
            return run_eager_on_cuda_graph_stream(self._fn, *args)

        if self._graph is None:
            self._args = args
            self._record_static_input_addresses(args)
            self._graph = torch.cuda.CUDAGraph()
            with torch.cuda.graph(
                self._graph,
                pool=_manager.graph_pool,
                stream=_manager.stream,
                enable_annotations=True,
                capture_error_mode="thread_local",
            ):
                self._output = self._fn(*args)
            _manager.all_annotations.update(get_kernel_annotations())
            logger.info("Recorded CUDA graph")

        if self._should_check_address:
            self._check_static_input_addresses(args)

        assert self._args is not None
        assert self._graph is not None
        for i in self._input_indices_to_copy:
            self._args[i].copy_(args[i])
        self._graph.replay()
        return self._output

    def teardown(self) -> None:
        self._graph = None
        self._args = None
        self._output = None
        self._static_input_addresses.clear()
        self._non_tensor_inputs.clear()


class _ForwardBackwardCUDAGraphWrapper(CUDAGraphWrapper):
    """Preserve capture-created parameter gradients across graph replays."""

    def __init__(
        self,
        fn: Callable,
        example_inputs: Sequence[Any],
        *,
        parameters: Iterable[torch.nn.Parameter],
        num_warmup_iterations: int,
    ) -> None:
        self._gradient_state = _CUDAGraphGradientState(parameters)
        super().__init__(
            fn,
            example_inputs,
            num_warmup_iterations=num_warmup_iterations,
        )

    def __call__(self, *args: Any) -> Any:
        if self._warmup_remaining > 0:
            return super().__call__(*args)

        self._gradient_state.require_cleared()
        should_record_gradients = self._graph is None
        output = super().__call__(*args)
        if should_record_gradients:
            self._gradient_state.record()
        self._gradient_state.restore()
        return output

    def teardown(self) -> None:
        self._gradient_state.teardown()
        super().teardown()


def cuda_graphs_supported() -> bool:
    """Whether CUDA graph capture can actually run on this build/device.

    ROCm is excluded: capture is unsupported there, so anything gated on CUDA
    graphs must agree that they are inert.
    """
    return (
        utils.device_type == "cuda"
        and torch.cuda.is_available()
        and torch.version.hip is None
    )


# TODO: Unify PP and non-PP callable signatures to restore strict input typing.
def _wrap_with_cuda_graph(
    fn: Callable[..., Any],
    *,
    num_warmup_iterations: int,
    gradient_parameters: tuple[torch.nn.Parameter, ...] | None,
) -> Callable[..., Any]:
    """Implement structured CUDA graph wrapping."""

    if not cuda_graphs_supported():
        logger.warning(
            "CUDA graph capture is only supported on NVIDIA CUDA; "
            "using eager execution."
        )
        return fn

    # Every wrapper is registered to the manager in this module and persists
    # until cuda_graph_teardown is called.
    graph_wrapper: CUDAGraphWrapper | None = None
    input_spec: CUDAGraphInputSpec | None = None

    def run(*args: Any, **kwargs: Any) -> Any:
        nonlocal graph_wrapper, input_spec

        if graph_wrapper is None:
            input_spec = CUDAGraphInputSpec((args, kwargs))

            def flat_fn(*flat_inputs: Any) -> torch.Tensor:
                assert input_spec is not None
                step_args, step_kwargs = input_spec.unflatten(flat_inputs)
                return fn(*step_args, **step_kwargs)

            flat_inputs = input_spec.flatten((args, kwargs))
            if gradient_parameters is None:
                graph_wrapper = CUDAGraphWrapper(
                    flat_fn,
                    flat_inputs,
                    num_warmup_iterations=num_warmup_iterations,
                )
            else:
                graph_wrapper = _ForwardBackwardCUDAGraphWrapper(
                    flat_fn,
                    flat_inputs,
                    parameters=gradient_parameters,
                    num_warmup_iterations=num_warmup_iterations,
                )
        else:
            assert input_spec is not None
            flat_inputs = input_spec.flatten((args, kwargs))

        return graph_wrapper(*flat_inputs)

    return run


def wrap_with_cuda_graph(
    fn: Callable[..., Any],
    *,
    num_warmup_iterations: int = 0,
) -> Callable[..., Any]:
    """Decorate a structured callable with CUDA graph capture and replay.

    The positional and keyword inputs must keep the same pytree structure and
    tensor metadata across calls. After capture, tensor outputs alias
    graph-owned storage that is overwritten by the next replay.

    Args:
        fn: Callable to capture.
        num_warmup_iterations: Number of eager invocations before capture.
    """

    return _wrap_with_cuda_graph(
        fn,
        num_warmup_iterations=num_warmup_iterations,
        gradient_parameters=None,
    )


def wrap_fwd_bwd_with_cuda_graph(
    fn: Callable[..., Any],
    *,
    parameters: Iterable[torch.nn.Parameter],
    num_warmup_iterations: int = 0,
) -> Callable[..., Any]:
    """Wrap forward-backward and preserve gradients allocated during capture.

    Args:
        fn: Forward-backward callable to capture.
        parameters: Parameters whose capture-created gradients must remain alive.
        num_warmup_iterations: Number of eager invocations before capture.
    """

    return _wrap_with_cuda_graph(
        fn,
        num_warmup_iterations=num_warmup_iterations,
        gradient_parameters=tuple(parameters),
    )
