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

"""
Memory policy passes for graph_trainer.

Selective activation checkpointing (SAC) tagging and memory policy dispatch.
Each saved forward activation can independently be tagged as MUST_SAVE,
MUST_RECOMPUTE, or MUST_CPU_OFFLOAD. The ``tag_with_memory_policy_pass`` entry
point selects a tagging strategy with ``compile.memory_policy``.
"""

from __future__ import annotations

import logging

import operator
from collections import defaultdict
from collections.abc import Callable
from typing import TYPE_CHECKING

import torch
from torch._functorch.partitioners import (
    choose_saved_values_set,
    force_save_bw_mutation_src,
    force_save_collectives,
    force_save_effectful_ops,
    get_default_op_list,
    NodeInfo,
)
from torch.utils._ordered_set import OrderedSet
from torch.utils.checkpoint import _is_cacheable_effect, CheckpointPolicy

from torchtitan.distributed.fsdp import get_fsdp_reshard_after_forward_policy
from torchtitan.experiments.graph_trainer.common_utils import (
    _get_layer_id,
    _get_module_fqn,
    _is_backward_node,
    _MODULE_FQN,
    _NOT_IN_LAYERS,
    matches_module_fqn_pattern,
)
from torchtitan.experiments.graph_trainer.cpu_offload import (
    tag_all_offloadable_activations,
)
from torchtitan.experiments.graph_trainer.fsdp_patterns import (
    find_fsdp_unshard_outputs_by_param,
)
from torchtitan.experiments.graph_trainer.log_activation_memory_policy import (
    log_activation_memory_policy,
)
from torchtitan.experiments.graph_trainer.registry import (
    MEMORY_POLICY_REGISTRY,
    register_memory_policy,
)

logger = logging.getLogger(__name__)


if TYPE_CHECKING:
    from torchtitan.experiments.graph_trainer.configs import GraphTrainerCompileConfig


_INF_DISTANCE = int(1e9)


def _get_default_save_ops() -> set:
    """Return the operator save set used by graph-trainer SAC policies.

    Copied from the former operator-level ``SelectiveAC`` policy in
    ``torchtitan/distributed/activation_checkpoint.py``, which now uses
    torch_remat regions instead of an operator save set.
    """
    compute_ops = [
        torch.ops.aten._scaled_dot_product_cudnn_attention.default,
        torch.ops.aten._scaled_dot_product_attention_math.default,
        torch.ops.aten._scaled_dot_product_fused_attention_overrideable.default,
        torch.ops.aten.max.default,
        torch._higher_order_ops.flex_attention,
        torch.ops.aten.linear.default,
        torch.ops.aten.mm.dtype,
        torch.ops.aten.topk.default,
        (torch._higher_order_ops, "inductor_compiled_code"),
        (torch.ops, "torch_attn._varlen_attn.default"),
    ]
    communication_ops = [
        torch.ops._c10d_functional.reduce_scatter_tensor.default,
        torch.ops._c10d_functional.all_to_all_single.default,
        (torch.ops, "deepep.dispatch.default"),
        (torch.ops, "deepep.combine.default"),
        (torch.ops, "hybridep.dispatch.default"),
        (torch.ops, "hybridep.combine.default"),
    ]

    def resolve_ops(op_specs: list) -> set:
        ops = set()
        for spec in op_specs:
            if isinstance(spec, tuple):
                obj, path = spec
                try:
                    for part in path.split("."):
                        obj = getattr(obj, part)
                    ops.add(obj)
                except AttributeError:
                    pass
            else:
                ops.add(spec)
        return ops

    aten_op_types = get_default_op_list()
    save_ops = {
        op.default  # pyrefly: ignore [missing-attribute]
        for op in aten_op_types.compute_intensive_ops
    }
    save_ops.update(resolve_ops(compute_ops))
    save_ops.update(resolve_ops(communication_ops))
    return save_ops


def _make_default_memory_policy(save_ops: set | None = None) -> Callable:
    """Create a SAC policy function from a set of op targets to save."""
    if save_ops is None:
        save_ops = _get_default_save_ops()

    def policy_fn(node: torch.fx.Node) -> CheckpointPolicy:
        if node.target in save_ops:
            return CheckpointPolicy.MUST_SAVE
        return CheckpointPolicy.PREFER_RECOMPUTE

    return policy_fn


def _make_no_ac_memory_policy() -> Callable:
    """Create a policy that saves every forward activation."""

    def policy_fn(node: torch.fx.Node) -> CheckpointPolicy:
        return CheckpointPolicy.MUST_SAVE

    return policy_fn


def _find_fsdp_unshard_save_nodes(gm: torch.fx.GraphModule) -> set[torch.fx.Node]:
    outputs_by_param = find_fsdp_unshard_outputs_by_param(
        gm.graph.find_nodes(op="placeholder")
    )
    return {output for outputs in outputs_by_param.values() for output in outputs}


def _resolve_op_target(op_name: str) -> object:
    """Resolve ``aten.mm.default``-style names through ``torch.ops``."""
    if op_name.startswith("torch.ops."):
        op_name = op_name.removeprefix("torch.ops.")

    target = torch.ops
    try:
        for component in op_name.split("."):
            target = getattr(target, component)
    except AttributeError as exc:
        raise ValueError(
            f"Unknown op in compile.full_recompute_save_ops: {op_name!r}"
        ) from exc

    if not isinstance(target, (torch._ops.OpOverload, torch._ops.HigherOrderOperator)):
        raise ValueError(
            "Ops in compile.full_recompute_save_ops must name a specific "
            f"overload or higher-order op, got {op_name!r}"
        )
    return target


def _parse_full_recompute_save_ops(
    value: str,
) -> tuple[tuple[str, object], ...]:
    """Parse ``FQN::OP | FQN::OP`` save selectors."""
    if not value.strip():
        return ()

    selectors: list[tuple[str, object]] = []
    for raw_selector in value.split("|"):
        parts = raw_selector.split("::")
        if len(parts) != 2 or not all(part.strip() for part in parts):
            raise ValueError(
                "Invalid compile.full_recompute_save_ops selector "
                f"{raw_selector.strip()!r}; expected 'MODULE_FQN_PATTERN::OP'"
            )
        module_fqn_pattern, op_name = (part.strip() for part in parts)
        selectors.append((module_fqn_pattern, _resolve_op_target(op_name)))
    return tuple(selectors)


def validate_memory_policy_config(
    compile_config: "GraphTrainerCompileConfig",
) -> None:
    """Validate memory-policy options before tracing the training graph."""
    if (
        compile_config.full_recompute_save_ops
        and compile_config.memory_policy != "full"
    ):
        raise ValueError(
            "compile.full_recompute_save_ops requires compile.memory_policy='full'"
        )
    _parse_full_recompute_save_ops(compile_config.full_recompute_save_ops)


def _make_full_memory_policy(save_ops: str = "") -> Callable:
    """Full recompute policy: mark everything as MUST_RECOMPUTE.

    The layer boundary pass in tag_sac_policy will force MUST_SAVE on nodes
    whose output crosses a layer boundary, so only layer outputs are saved.
    This mirrors TorchTitan's former eager full AC, which recomputed the entire
    block -- including attention -- in backward.

    RNG ops (dropout etc.) are the one class that is always saved: the remat
    pass cannot replay their random state, and ``has_recomputable_rng_ops``
    would otherwise raise. Eager full AC recomputes them by forking the RNG
    state (``preserve_rng_state=True``); the graph path lacks that, so it
    saves them instead.

    Higher-order ops (e.g. flex_attention) ARE recomputed: ``node_copy``
    duplicates them together with their ``get_attr`` subgraph references, and
    the subsequent regional_inductor pass compiles the duplicate as well.

    ``save_ops`` can make exact module-FQN-pattern and op pairs exceptions to
    full recompute. Matching nodes are marked MUST_SAVE.
    """

    save_selectors = _parse_full_recompute_save_ops(save_ops)

    def policy_fn(node: torch.fx.Node) -> CheckpointPolicy:
        if torch.Tag.nondeterministic_seeded in getattr(node.target, "tags", set()):
            return CheckpointPolicy.MUST_SAVE
        fqn = _get_module_fqn(node)
        if any(
            node.target == target and matches_module_fqn_pattern(fqn_pattern, fqn)
            for fqn_pattern, target in save_selectors
        ):
            return CheckpointPolicy.MUST_SAVE
        return CheckpointPolicy.MUST_RECOMPUTE

    return policy_fn


def _make_eager_memory_policy(save_ops: set | None = None) -> Callable:
    """Eager-compatible SAC policy that alternates mm ops between save/recompute.

    Matches TorchTitan's former eager per-op SelectiveAC policy: every second
    mm/linear op is marked PREFER_RECOMPUTE instead of MUST_SAVE. The mm counter
    resets at each layer boundary so every layer sees the same alternation
    pattern.
    """
    if save_ops is None:
        save_ops = _get_default_save_ops()
    mm_ops = {
        torch.ops.aten.mm.default,
        torch.ops.aten.mm.dtype,
        torch.ops.aten.linear.default,
    }
    mm_count = 0
    current_layer = None

    def policy_fn(node: torch.fx.Node) -> CheckpointPolicy:
        nonlocal mm_count, current_layer
        layer_id = _get_layer_id(node)
        if layer_id != _NOT_IN_LAYERS and layer_id != current_layer:
            mm_count = 0
            current_layer = layer_id

        if node.target in mm_ops:
            mm_count += 1
            if node.target in save_ops and mm_count % 2 == 0:
                return CheckpointPolicy.PREFER_RECOMPUTE
        if node.target in save_ops:
            return CheckpointPolicy.MUST_SAVE
        return CheckpointPolicy.PREFER_RECOMPUTE

    return policy_fn


def tag_sac_policy(
    gm: torch.fx.GraphModule,
    example_inputs: tuple | None = None,
    *,
    policy_fn: Callable[[torch.fx.Node], CheckpointPolicy] | None = None,
    force_save_nodes: set[torch.fx.Node] | None = None,
) -> torch.fx.GraphModule:
    """Apply selective activation checkpointing on the joint graph.

    Annotates forward ``call_function`` nodes with a ``CheckpointPolicy``
    determined by ``policy_fn``. After tagging, a boundary pass forces
    ``MUST_SAVE`` on recomputable nodes whose output leaves its producing
    layer, since recomputing them would require rerunning that layer. This
    includes the final layer output consumed by the model epilogue.

    ``getitem`` / ``wait_tensor`` nodes inherit the parent's tag.

    The model must have been annotated with ``annotate_module_fqns`` before
    tracing so that nodes carry ``module_fqn`` metadata.

    Args:
        gm: The joint forward-backward graph module.
        policy_fn: Callable that takes a node and returns a CheckpointPolicy.
            Defaults to ``_make_default_memory_policy()`` if None.
        force_save_nodes: Nodes that must be saved independent of ``policy_fn``.
            Used for graph-structure constraints such as FSDP unshards with
            ``reshard_after_forward=False``.

    Returns:
        The annotated graph module
    """
    if policy_fn is None:
        policy_fn = _make_default_memory_policy()
    if force_save_nodes is None:
        force_save_nodes = set()

    # Pass 1: Tag each forward node with a recompute policy.
    for node in gm.graph.nodes:
        if node.op != "call_function":
            continue

        # Skip backward nodes — they must not carry recompute tags,
        # otherwise the remat pass would try to duplicate backward ops.
        if _is_backward_node(node):
            continue

        # Skip the post-layer epilogue (lm_head + loss). Chunked-loss
        # regions split backward into multiple disjoint regions, and the
        # remat pass only supports one region with must_recompute deps.
        fqn = node.meta.get("custom", {}).get(_MODULE_FQN, "")
        if fqn.startswith(("lm_head", "loss")):
            continue

        if _is_cacheable_effect(node.target):
            node.meta["recompute"] = CheckpointPolicy.MUST_SAVE
            continue

        if node in force_save_nodes:
            node.meta["recompute"] = CheckpointPolicy.MUST_SAVE
            continue

        if node.target in (
            operator.getitem,
            torch.ops._c10d_functional.wait_tensor.default,
        ):
            # Propagate from parent: getitem extracts tuple elements,
            # wait_tensor is tied to its async collective — both must
            # share the parent's save/recompute decision.
            parent = node.args[0]
            if isinstance(parent, torch.fx.Node) and "recompute" in parent.meta:
                node.meta["recompute"] = parent.meta["recompute"]
            continue

        # Always save sym-int nodes (shape reads like sym_size/sym_stride, and
        # tensor->int scalar conversions) rather than recompute them: recomputing
        # a shape read pins the parent tensor alive in backward just to reread its
        # size. We key off meta["val"] being a SymInt -- mirroring AOT Autograd's
        # partitioner, which saves SymInts (cheap scalars) but never SymFloats.
        if "val" in node.meta and isinstance(node.meta["val"], torch.SymInt):
            node.meta["recompute"] = CheckpointPolicy.MUST_SAVE
            continue

        # NOTE: The eager SAC policy (activation_checkpoint.py) alternates
        # mm ops between MUST_SAVE and PREFER_RECOMPUTE. We omit that here
        # because the alternating heuristic is arbitrary.
        node.meta["recompute"] = policy_fn(node)

    # Pass 2: Save recomputable outputs consumed outside their producing
    # layer. Recreating one would require rerunning that producer layer. A
    # consumer without a layer FQN is also outside the producer layer; this
    # covers the final layer output consumed by the model epilogue.
    def _is_recomputable(n: torch.fx.Node) -> bool:
        return n.meta.get("recompute") in (
            CheckpointPolicy.PREFER_RECOMPUTE,
            CheckpointPolicy.MUST_RECOMPUTE,
        )

    boundary_saves = 0
    for node in gm.graph.nodes:
        if _is_backward_node(node) or not _is_recomputable(node):
            continue
        node_layer_id = _get_layer_id(node)
        for user in node.users:
            if (
                not _is_backward_node(user)
                and _is_recomputable(user)
                and _get_layer_id(user) != node_layer_id
            ):
                node.meta["recompute"] = CheckpointPolicy.MUST_SAVE
                boundary_saves += 1
                break

    gm.recompile()

    # Per-layer summary from the FINAL policy (after boundary forcing). Counts
    # every forward node carrying a recompute decision — primary policy-tagged
    # ops plus getitem/wait_tensor (which inherit a parent's tag) plus any node
    # the boundary pass forced to MUST_SAVE — so the per-layer MUST_SAVE counts
    # account for boundary saves wherever they land (including on getitem /
    # wait_tensor). The recompute count covers both PREFER_RECOMPUTE and
    # MUST_RECOMPUTE, so the column is labelled generically as RECOMPUTE.
    layer_stats: dict[int, dict[str, int]] = defaultdict(
        lambda: {"save": 0, "recompute": 0}
    )
    for node in gm.graph.nodes:
        if "recompute" not in node.meta:
            continue
        key = (
            "save"
            if node.meta["recompute"] == CheckpointPolicy.MUST_SAVE
            else "recompute"
        )
        layer_stats[_get_layer_id(node)][key] += 1

    logger.info("Applied selective activation checkpointing (SAC) graph pass.")
    if boundary_saves:
        logger.info(f"  Forced {boundary_saves} nodes to MUST_SAVE at layer boundaries")
    for layer_id in sorted(layer_stats):
        stats = layer_stats[layer_id]
        label = "non-layer" if layer_id == _NOT_IN_LAYERS else str(layer_id)
        logger.info(
            f"  Layer {label}: "
            f"{stats['save']} MUST_SAVE, "
            f"{stats['recompute']} RECOMPUTE"
        )
    return gm


@register_memory_policy("none")
def _no_ac_memory_policy_pass(
    gm: torch.fx.GraphModule,
    *,
    config: "GraphTrainer.Config",
) -> torch.fx.GraphModule:
    """Save every forward activation without rematerialization."""
    tag_sac_policy(gm, policy_fn=_make_no_ac_memory_policy())
    return gm


@register_memory_policy("default")
def _default_memory_policy_pass(
    gm: torch.fx.GraphModule,
    *,
    config: "GraphTrainer.Config",
) -> torch.fx.GraphModule:
    """SAC policy that saves compute-intensive ops and required FSDP unshards."""
    fsdp_reshard_after_forward = get_fsdp_reshard_after_forward_policy(
        config.parallelism.fsdp_reshard_after_forward,
        pp_enabled=config.parallelism.pipeline_parallel_degree > 1,
    )
    force_save_nodes = (
        _find_fsdp_unshard_save_nodes(gm) if not fsdp_reshard_after_forward else None
    )
    tag_sac_policy(
        gm,
        policy_fn=_make_default_memory_policy(),
        force_save_nodes=force_save_nodes,
    )
    return gm


@register_memory_policy("full")
def _full_memory_policy_pass(
    gm: torch.fx.GraphModule,
    *,
    config: "GraphTrainer.Config",
) -> torch.fx.GraphModule:
    """Full recompute except for user-selected module operations."""
    tag_sac_policy(
        gm,
        policy_fn=_make_full_memory_policy(config.compile.full_recompute_save_ops),
    )
    return gm


@register_memory_policy("eager")
def _eager_memory_policy_pass(
    gm: torch.fx.GraphModule,
    *,
    config: "GraphTrainer.Config",
) -> torch.fx.GraphModule:
    """SAC policy that alternates mm ops between save/recompute."""
    tag_sac_policy(gm, policy_fn=_make_eager_memory_policy())
    return gm


def _is_backward_side(node: torch.fx.Node, backward_side: set[torch.fx.Node]) -> bool:
    return _is_backward_node(node) or any(
        inp in backward_side for inp in node.all_input_nodes
    )


def _backward_side_nodes(
    gm: torch.fx.GraphModule,
) -> OrderedSet[torch.fx.Node]:
    backward_side = OrderedSet()
    for node in gm.graph.nodes:
        if node.op != "output" and _is_backward_side(node, backward_side):
            backward_side.add(node)
    return backward_side


def _node_info_for_graph_trainer(
    gm: torch.fx.GraphModule,
    backward_side: OrderedSet[torch.fx.Node],
) -> NodeInfo | None:
    nodes = list(gm.graph.nodes)
    required_bw_nodes = OrderedSet(
        node for node in nodes if node in backward_side and node.op != "output"
    )
    if not required_bw_nodes:
        return None

    required_fw_nodes = OrderedSet(
        node for node in nodes if node not in required_bw_nodes and node.op != "output"
    )
    fw_order = {node: idx for idx, node in enumerate(required_fw_nodes)}
    static_lifetime_input_nodes = OrderedSet(
        node for node in required_fw_nodes if node.op in ("placeholder", "get_attr")
    )

    for node in reversed(nodes):
        if node.op == "output":
            node.dist_from_bw = _INF_DISTANCE
        elif node in required_bw_nodes:
            node.dist_from_bw = 0
        elif node in required_fw_nodes:
            user_distances = [
                getattr(user, "dist_from_bw", _INF_DISTANCE) + 1 for user in node.users
            ]
            node.dist_from_bw = min(user_distances, default=_INF_DISTANCE)
        else:
            node.dist_from_bw = _INF_DISTANCE

    return NodeInfo(
        list(static_lifetime_input_nodes),
        required_fw_nodes,
        required_bw_nodes.copy(),
        required_bw_nodes.copy(),
        OrderedSet(),
        fw_order,
        static_lifetime_input_nodes,
    )


def tag_min_cut_saved_values(
    gm: torch.fx.GraphModule,
    backward_side: OrderedSet[torch.fx.Node],
    saved_values: set[torch.fx.Node],
) -> None:
    required_fw_nodes = {
        node
        for node in gm.graph.nodes
        if node not in backward_side and node.op != "output"
    }
    saved_boundaries = set(saved_values)
    op_types = get_default_op_list()
    pending = list(saved_boundaries)
    while pending:
        node = pending.pop()
        if node not in saved_boundaries:
            continue
        if node not in required_fw_nodes or node.op != "call_function":
            continue
        if (
            node.target == torch.ops.aten.detach.default or op_types.is_view(node)
        ) and any(inp in required_fw_nodes for inp in node.all_input_nodes):
            saved_boundaries.remove(node)
            for inp in node.all_input_nodes:
                if inp in required_fw_nodes and inp not in saved_boundaries:
                    saved_boundaries.add(inp)
                    pending.append(inp)

    saved_boundaries.update(
        node
        for node in required_fw_nodes
        if node.meta.get("recompute") == CheckpointPolicy.MUST_SAVE
    )
    for node in saved_boundaries:
        if node in required_fw_nodes:
            node.meta["recompute"] = CheckpointPolicy.MUST_SAVE

    seen = set()

    def visit(node: torch.fx.Node) -> None:
        if node in seen or node in saved_boundaries:
            return
        seen.add(node)
        if node in backward_side:
            for inp in node.all_input_nodes:
                visit(inp)
            return
        if node not in required_fw_nodes or node.op in ("placeholder", "get_attr"):
            return
        if node.op == "call_function":
            node.meta["recompute"] = CheckpointPolicy.MUST_RECOMPUTE
            for inp in node.all_input_nodes:
                visit(inp)

    for node in backward_side:
        for inp in node.all_input_nodes:
            visit(inp)


@register_memory_policy("min_cut")
def _min_cut_memory_policy_pass(
    gm: torch.fx.GraphModule,
    *,
    config: "GraphTrainer.Config",
) -> torch.fx.GraphModule:
    """Choose saved activations with the min-cut partitioner."""
    backward_side = _backward_side_nodes(gm)
    node_info = _node_info_for_graph_trainer(gm, backward_side)
    if node_info is None:
        return gm

    force_save_collectives(gm)
    force_save_effectful_ops(gm)
    force_save_bw_mutation_src(gm)
    saved_values = choose_saved_values_set(gm.graph, node_info)
    tag_min_cut_saved_values(gm, backward_side, set(saved_values))
    return gm


@register_memory_policy("sac_and_offload")
def _sac_and_offload_memory_policy_pass(
    gm: torch.fx.GraphModule,
    *,
    config: "GraphTrainer.Config",
) -> torch.fx.GraphModule:
    """SAC + CPU offload: apply default SAC, then offload within budget."""
    _default_memory_policy_pass(gm, config=config)
    tag_all_offloadable_activations(
        gm,
        cpu_budget_gb=config.compile.cpu_offload_budget_gb,
    )
    return gm


def tag_with_memory_policy_pass(
    gm: torch.fx.GraphModule,
    example_inputs: tuple | None = None,
    *,
    config: "GraphTrainer.Config",
) -> torch.fx.GraphModule:
    """Tag forward nodes with MUST_SAVE, PREFER_RECOMPUTE, or MUST_CPU_OFFLOAD.

    The ``config.compile.memory_policy`` selects the tagging strategy:
        none: save every forward activation without rematerialization.
        default: SAC with all compute-intensive ops saved.
        full: full recompute except user-selected module operations.
        eager: SAC alternating mm ops between save/recompute.
        min_cut: choose saved activations with the min-cut partitioner.
        sac_and_offload: SAC + CPU offload within budget.

    Other memory policies combining SAC and CPU offload can be added
    via ``register_memory_policy`` without modifying this function.

    """
    memory_policy = config.compile.memory_policy
    if memory_policy not in MEMORY_POLICY_REGISTRY:
        raise ValueError(
            f"Unknown memory_policy: {memory_policy!r}. "
            f"Available: {list(MEMORY_POLICY_REGISTRY.keys())}"
        )

    gm = MEMORY_POLICY_REGISTRY[memory_policy](gm, config=config)
    log_activation_memory_policy(gm)
    return gm
