# 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.
"""GraphPP runtime action handlers.

GraphPP uses upstream runtime PP schedules for ordering, communication,
microbatch splitting, and stage metadata initialization. This module only maps
schedule actions onto bound stage graph executors.
"""

from __future__ import annotations

import logging
from enum import Enum
from typing import Any, cast, TYPE_CHECKING

import torch
from torch.distributed.pipelining import PipelineStageInfo
from torch.distributed.pipelining.schedules import (
    _Action,
    _PipelineContext,
    _PipelineScheduleRuntime,
    _wait_batch_p2p,
    BACKWARD_INPUT,
    BACKWARD_WEIGHT,
    FORWARD,
    OVERLAP_F_B,
    REDUCE_GRAD,
    RESHARD,
    UNSHARD,
    WAIT_REDUCE_GRAD,
)
from torch.distributed.pipelining.stage import _normalize_model_output_as_tuple

from torchtitan.experiments.graph_trainer.common_utils import accumulate_param_grads_
from torchtitan.experiments.graph_trainer.graph_pp.stage import (
    FSDPBoundaryJointStageGraphs,
    GraphPipelineStage,
    JointStageGraphs,
    NoGradAccumJointStageGraphs,
    OverlapStageGraphs,
    SplitStageGraphs,
    StageGraphs,
    StageGraphsProvider,
)
from torchtitan.experiments.graph_trainer.graph_pp.utils import (
    flatten_graph_values,
    overlap_fw_bw_sub_actions,
)

if TYPE_CHECKING:
    from torchtitan.models.common.dist_moe.runtime import _DistMoeForwardContext


logger = logging.getLogger(__name__)


__all__ = [
    "BACKWARD",
    "BACKWARD_WEIGHT_WITH_REDUCE_GRAD",
    "BACKWARD_WITH_REDUCE_GRAD",
    "FORWARD_BACKWARD",
    "FORWARD_BACKWARD_FIRST_WITH_UNSHARD",
    "FORWARD_BACKWARD_LAST_WITH_REDUCE_GRAD",
    "FORWARD_BACKWARD_NOGRADACCUM",
    "FULL_FORWARD_BACKWARD",
    "GraphRuntime",
    "register_graph_schedule",
]


class _GraphComputationType(Enum):
    # Backward graph with gradient reduction extracted into REDUCE_GRAD.
    BACKWARD = "BACKWARD"
    # Backward graph that includes gradient reduction.
    BACKWARD_WITH_REDUCE_GRAD = "BACKWARD_WITH_REDUCE_GRAD"
    # Weight-backward graph that includes gradient reduction.
    BACKWARD_WEIGHT_WITH_REDUCE_GRAD = "BACKWARD_WEIGHT_WITH_REDUCE_GRAD"
    # Joint graph without FSDP unshard and reduce_grad.
    FORWARD_BACKWARD = "FORWARD_BACKWARD"
    # First joint graph without accumulator inputs; FSDP boundaries match the
    # repeated graph, and its gradients become the accumulators.
    FORWARD_BACKWARD_NOGRADACCUM = "FORWARD_BACKWARD_NOGRADACCUM"
    # First joint graph that unshards parameters and returns them for reuse.
    FORWARD_BACKWARD_FIRST_WITH_UNSHARD = "FORWARD_BACKWARD_FIRST_WITH_UNSHARD"
    # Last joint graph that reduces accumulated gradients.
    FORWARD_BACKWARD_LAST_WITH_REDUCE_GRAD = "FORWARD_BACKWARD_LAST_WITH_REDUCE_GRAD"
    # Joint graph with FSDP unshard and reduce_grad inside.
    FULL_FORWARD_BACKWARD = "FULL_FORWARD_BACKWARD"


BACKWARD = _GraphComputationType.BACKWARD
BACKWARD_WITH_REDUCE_GRAD = _GraphComputationType.BACKWARD_WITH_REDUCE_GRAD
BACKWARD_WEIGHT_WITH_REDUCE_GRAD = (
    _GraphComputationType.BACKWARD_WEIGHT_WITH_REDUCE_GRAD
)
FORWARD_BACKWARD = _GraphComputationType.FORWARD_BACKWARD
FORWARD_BACKWARD_NOGRADACCUM = _GraphComputationType.FORWARD_BACKWARD_NOGRADACCUM
FORWARD_BACKWARD_FIRST_WITH_UNSHARD = (
    _GraphComputationType.FORWARD_BACKWARD_FIRST_WITH_UNSHARD
)
FORWARD_BACKWARD_LAST_WITH_REDUCE_GRAD = (
    _GraphComputationType.FORWARD_BACKWARD_LAST_WITH_REDUCE_GRAD
)
FULL_FORWARD_BACKWARD = _GraphComputationType.FULL_FORWARD_BACKWARD


def _scale_grad_values_(grads: list[Any], grad_scale_factor: int) -> None:
    if grad_scale_factor == 1:
        return
    for grad in grads:
        if isinstance(grad, torch.Tensor):
            grad.div_(grad_scale_factor)


def _scale_graph_pp_sharded_grads(
    stage: GraphPipelineStage,
    schedule: _PipelineScheduleRuntime,
) -> None:
    if stage._graph_pp_grads_scaled:
        return
    grad_scale_factor = schedule._n_microbatches if schedule.scale_grads else 1
    _scale_grad_values_(stage.state.sharded_param_grads, grad_scale_factor)
    stage._graph_pp_grads_scaled = True


def _accumulate_flat_grad_values_(
    accumulated: list[Any],
    grads: list[Any],
    *,
    label: str,
    runtime_validate: bool,
) -> None:
    """Accumulate raw flat graph-gradient values before boundary rewrapping."""
    if len(grads) != len(accumulated):
        raise ValueError(
            f"GraphPP {label} grad count mismatch: "
            f"expected {len(accumulated)}, got {len(grads)}"
        )
    for index, grad in enumerate(grads):
        if grad is None:
            continue
        if not isinstance(grad, torch.Tensor):
            if accumulated[index] is None:
                accumulated[index] = grad
            elif (
                runtime_validate
                and accumulated[index] is not grad
                and accumulated[index] != grad
            ):
                raise ValueError(
                    "GraphPP flat gradient metadata changed across "
                    f"microbatches at index {index}: "
                    f"{accumulated[index]!r} != {grad!r}"
                )
            continue
        if accumulated[index] is None:
            accumulated[index] = grad
        else:
            accumulated[index] += grad


def _ensure_unsharded_param_values(
    stage: GraphPipelineStage,
    graphs: StageGraphs,
) -> None:
    if stage.state.unsharded_param_values:
        return
    stage.state.unsharded_param_values = graphs.unshard_params(
        stage.state.sharded_param_values,
        runtime_validate=stage._runtime_validate,
    )


def _accumulate_stage_unsharded_grads(
    stage: GraphPipelineStage,
    grads: list[Any],
) -> None:
    _accumulate_flat_grad_values_(
        stage.state.unsharded_param_grads,
        grads,
        label="unsharded",
        runtime_validate=stage._runtime_validate,
    )


def _grad_reduction_runs_in_backward(action: _Action) -> bool:
    if action.computation_type in (
        BACKWARD_WITH_REDUCE_GRAD,
        BACKWARD_WEIGHT_WITH_REDUCE_GRAD,
    ):
        return True
    if action.computation_type in (BACKWARD, BACKWARD_WEIGHT):
        return False
    raise ValueError(f"Unsupported backward action: {action}")


def _prepare_fwd_user_args(
    stage: GraphPipelineStage,
    mb_index: int,
    ctx: _PipelineContext,
) -> tuple[tuple[Any, ...], dict[str, Any], Any]:
    """Package runtime forward inputs for a stage graph call.

    Upstream PP owns microbatch splitting and stage metadata initialization.
    This helper only translates those upstream runtime containers into the
    GraphPP forward calling convention.
    """
    arg_mbs = ctx.arg_mbs
    kwarg_mbs = ctx.kwarg_mbs
    kwargs = {} if kwarg_mbs is None else kwarg_mbs[mb_index]
    if stage.is_first:
        args = () if arg_mbs is None else arg_mbs[mb_index]
    else:
        args = _normalize_model_output_as_tuple(
            stage._retrieve_recv_activations(mb_index)
        )
    target = ctx.target_mbs[mb_index] if stage.is_last and ctx.target_mbs else None
    return tuple(args), kwargs, target


def _stage_map_and_stage_from_action(
    schedule: _PipelineScheduleRuntime,
    action: _Action,
) -> tuple[dict[int, GraphPipelineStage], GraphPipelineStage]:
    """Return local GraphPP stages for an upstream schedule action.

    This is PP schedule mechanics, not runtime state: upstream schedules address
    stages by global stage index, while GraphPP handlers need both the selected
    local stage and the local-stage map for same-rank neighbor propagation.
    """
    stage_index_to_stage = {
        stage.stage_index: cast(GraphPipelineStage, stage) for stage in schedule._stages
    }
    return (
        stage_index_to_stage,
        stage_index_to_stage[action.stage_index],
    )


def _prepare_fwd_common(
    schedule: _PipelineScheduleRuntime,
    action: _Action,
) -> tuple[dict[int, GraphPipelineStage], GraphPipelineStage, int, bool]:
    """Prepare PP schedule state before one GraphPP forward graph call."""
    # 1. Resolve the local stage and rank-local topology for this action.
    stage_index_to_stage, stage = _stage_map_and_stage_from_action(schedule, action)
    mb_index = action.microbatch_index
    if mb_index is None:
        raise ValueError(f"GraphPP FORWARD action must have microbatch index: {action}")
    is_next_stage_on_this_rank = stage.stage_index + 1 in stage_index_to_stage
    is_prev_stage_on_this_rank = stage.stage_index - 1 in stage_index_to_stage

    # 2. Remote previous-stage activations must arrive before the graph reads
    #    the upstream PP receive buffer. Same-rank V-schedule neighbors bypass
    #    P2P and use PipelineStage local input caches.
    if not stage.is_first and not is_prev_stage_on_this_rank:
        fwd_recv_ops = schedule.fwd_recv_ops
        if (stage.stage_index, mb_index) not in fwd_recv_ops:
            raise ValueError(
                "GraphPP missing forward recv op for "
                f"stage {stage.stage_index}, microbatch {mb_index}."
            )
        _wait_batch_p2p(fwd_recv_ops.pop((stage.stage_index, mb_index)))
    return (
        stage_index_to_stage,
        stage,
        mb_index,
        is_next_stage_on_this_rank,
    )


def _post_fwd_common(
    stage: GraphPipelineStage,
    mb_index: int,
    output: Any,
    saved_values_for_backward: tuple[Any, ...],
    schedule: _PipelineScheduleRuntime,
    stage_index_to_stage: dict[int, GraphPipelineStage],
    is_next_stage_on_this_rank: bool,
) -> None:
    """Store forward graph outputs and propagate same-rank activations."""
    # 1. Record upstream transport ownership and explicit graph backward values.
    #    Last-stage losses are kept in upstream's internal loss list so
    #    schedule._update_losses() remains the only public loss updater.
    output_tuple = _normalize_model_output_as_tuple(output)
    if stage.is_last:
        stage.output_chunks.append(output)
        schedule._internal_losses.append(output)
    stage._record_graph_forward(
        mb_index,
        output_tuple,
        saved_values_for_backward,
    )

    # 2. Adjacent same-rank stages avoid SEND/RECV actions, so hand the output
    #    directly to the next stage's upstream local forward-input cache.
    if is_next_stage_on_this_rank:
        stage_index_to_stage[stage.stage_index + 1].set_local_fwd_input(
            output, mb_index
        )


def _prepare_backward_values(
    stage: GraphPipelineStage,
    mb_index: int,
) -> tuple[tuple[Any, ...], tuple[Any, ...]]:
    """Package cached forward values and received output grads for backward."""
    saved_values_for_backward = stage._take_graph_backward_values(mb_index)
    if stage.is_last:
        output_grads_from_next = ()
    else:
        output_grads_from_next = _normalize_model_output_as_tuple(
            stage._retrieve_recv_grads(mb_index)
        )
    return saved_values_for_backward, output_grads_from_next


def _prepare_backward_common(
    schedule: _PipelineScheduleRuntime,
    action: _Action,
) -> tuple[dict[int, GraphPipelineStage], GraphPipelineStage, int, bool]:
    """Prepare PP schedule state before one GraphPP backward graph call."""
    # 1. Resolve the local stage and rank-local topology for this action.
    stage_index_to_stage, stage = _stage_map_and_stage_from_action(schedule, action)
    mb_index = action.microbatch_index
    if mb_index is None:
        raise ValueError(
            f"GraphPP backward action must have microbatch index: {action}"
        )
    is_next_stage_on_this_rank = stage.stage_index + 1 in stage_index_to_stage
    is_prev_stage_on_this_rank = stage.stage_index - 1 in stage_index_to_stage

    # 2. Remote next-stage output grads must arrive before backward reads the
    #    upstream PP grad receive buffer. Same-rank neighbors use local caches.
    if not stage.is_last and not is_next_stage_on_this_rank:
        bwd_recv_ops = schedule.bwd_recv_ops
        if (stage.stage_index, mb_index) not in bwd_recv_ops:
            raise ValueError(
                "GraphPP missing backward recv op for "
                f"stage {stage.stage_index}, microbatch {mb_index}."
            )
        _wait_batch_p2p(bwd_recv_ops.pop((stage.stage_index, mb_index)))

    # 3. Match upstream runtime semantics: full-backward and backward-input
    #    count as the microbatch backward action for reduce-grad placement.
    schedule.backward_counter[stage.stage_index] += 1
    return (
        stage_index_to_stage,
        stage,
        mb_index,
        is_prev_stage_on_this_rank,
    )


def _post_backward_common(
    stage: GraphPipelineStage,
    mb_index: int,
    input_grads: list[Any],
    stage_index_to_stage: dict[int, GraphPipelineStage],
    is_prev_stage_on_this_rank: bool,
) -> None:
    """Store input gradients and propagate same-rank backward inputs."""
    # 1. Cache input grads in upstream's stage cache so SEND_B actions, or
    #    same-rank previous stages, can retrieve them by microbatch.
    stage.bwd_cache[mb_index] = tuple(input_grads)

    # 2. Adjacent same-rank stages avoid SEND/RECV actions, so hand the grads
    #    directly to the previous stage's upstream local backward-input cache.
    if is_prev_stage_on_this_rank:
        stage_index_to_stage[stage.stage_index - 1].set_local_bwd_input(
            stage.get_local_bwd_output(mb_index),
            mb_index,
        )


class GraphRuntime:
    """Execute a schedule through bound stage graph handlers.

    ``GraphRuntime`` registers bound action handlers on an upstream
    ``_PipelineScheduleRuntime``. Ownership is split along the runtime
    boundary:

    1. Upstream PP owns microbatch splitting, schedule action ordering, P2P op
       creation, and stage metadata initialization.
    2. ``GraphRuntime`` owns per-step runtime state: graph readiness,
       action dispatch, ``loss_kwargs``, overlap graphs, grad scaling, and
       final parameter-gradient accumulation.
    3. ``SplitStageGraphs`` owns graph calling conventions, metadata packing,
       FX execution, and output unwrapping for one stage.
    4. Graph construction and tracing stay outside the runtime in a
       ``StageGraphsProvider`` implementation.

    The runtime does not inspect GraphTrainer graph metadata. Calling
    conventions live behind ``StageGraphs`` implementations attached to
    each ``GraphPipelineStage``. Module-private helpers in this file own the
    upstream PP schedule mechanics that surround each graph call: recv waits,
    forward/backward caches, and same-rank local propagation.

    Args:
        schedule (_PipelineScheduleRuntime): Runtime pipeline schedule to
            execute.
        graph_provider (StageGraphsProvider | None): Optional provider
            that attaches bound stage graphs before the first runtime action.
            If omitted, every local stage must already have ``stage.graphs``
            populated.
        is_spmd (bool): Whether the schedule is SPMD schedule,
            that does not require PP specific initialization
            (e.g. pipeline communication buffers).
        liveness_schedule: Optional pre-rewrite schedule retained for analyses
            whose action vocabulary differs from GraphPP execution.
    Raises:
        TypeError: If any local schedule stage is not a ``GraphPipelineStage``.
    """

    def __init__(
        self,
        schedule: _PipelineScheduleRuntime,
        *,
        graph_provider: StageGraphsProvider | None = None,
        is_spmd: bool = False,
        liveness_schedule: _PipelineScheduleRuntime | None = None,
    ) -> None:
        self.schedule = schedule
        self._liveness_schedule = liveness_schedule or schedule
        self.graph_provider = graph_provider
        self.is_spmd = is_spmd
        self.overlap_graphs: dict[tuple[int, int], OverlapStageGraphs] = {}
        self.stage_graphs: dict[int, StageGraphs] = {}
        self._dist_moe_forward_context: _DistMoeForwardContext | None = None
        self.loss_kwargs: dict[str, Any] = {}
        self._graph_pp_ready = False
        self._joint_gradient_accumulation_stage_indices = {
            action.stage_index
            for rank_actions in schedule.pipeline_order_with_comms.values()
            for action in rank_actions
            if action.computation_type
            in (
                FORWARD_BACKWARD_NOGRADACCUM,
                FORWARD_BACKWARD_FIRST_WITH_UNSHARD,
            )
        }
        self.schedule._has_backward = True
        for stage in schedule._stages:
            if not isinstance(stage, GraphPipelineStage):
                raise TypeError(
                    "GraphRuntime requires GraphPipelineStage instances, got "
                    f"{type(stage).__name__}"
                )
        assert schedule._stages

    @property
    def num_microbatches(self) -> int:
        """Return the number of microbatches owned by this runtime schedule."""
        return self.schedule._n_microbatches

    @property
    def pipeline_liveness_schedule(self) -> _PipelineScheduleRuntime:
        """Return the pre-rewrite schedule used for activation liveness."""
        return self._liveness_schedule

    def set_dist_moe_forward_context(
        self,
        forward_context: _DistMoeForwardContext | None,
    ) -> None:
        """Install or remove the Dist-MoE slot resolver between schedule steps."""
        if forward_context is not None and self._graph_pp_ready:
            raise RuntimeError(
                "GraphPP Dist-MoE context cannot change during a schedule step"
            )
        self._dist_moe_forward_context = forward_context

    def _resolve_dist_moe_activation_slot(
        self,
        action: _Action,
    ) -> torch.Tensor | None:
        """Resolve the immutable Dist-MoE slot view for a forward action."""
        forward_context = self._dist_moe_forward_context
        if forward_context is None:
            return None
        microbatch_index = action.microbatch_index
        if microbatch_index is None:
            raise ValueError(
                f"GraphPP forward action must have microbatch index: {action}"
            )
        return forward_context.resolve_activation_slot(
            PipelineStageInfo(
                stage_index=action.stage_index,
                microbatch_index=microbatch_index,
            )
        )

    def ensure_ready(self, ctx: _PipelineContext) -> None:
        """Ensure local stage graphs and runtime state are ready for execution.

        Args:
            ctx (_PipelineContext): Pipeline schedule context for the current
                step. The graph provider uses it to derive trace inputs and
                overlap graph pairs.

        Raises:
            ValueError: If a local stage has no bound graph executor after the
                optional graph provider runs.
        """
        if self._graph_pp_ready:
            return
        if self.graph_provider is not None:
            self.overlap_graphs = self.graph_provider.prepare_graphs(
                self.schedule,
                ctx,
                loss_kwargs=self.loss_kwargs,
                dist_moe_forward_context=self._dist_moe_forward_context,
            )
        self.stage_graphs = {}
        for stage in self.schedule._stages:
            graph_stage = cast(GraphPipelineStage, stage)
            if graph_stage.graphs is None:
                raise ValueError(
                    "GraphPP stage graphs must be built before runtime "
                    f"execution. Missing graphs for stage {graph_stage.stage_index}."
                )
            self.stage_graphs[graph_stage.stage_index] = cast(
                StageGraphs, graph_stage.graphs
            )
            self._populate_stage_states(graph_stage)
        self._graph_pp_ready = True

    def _populate_stage_states(self, stage: GraphPipelineStage) -> None:
        sharded_param_values = []
        buffer_values = []
        trainable_params = []
        for _, value in stage.submod.named_parameters(remove_duplicate=False):
            sharded_param_values.extend(flatten_graph_values([value]))
            if value.requires_grad:
                trainable_params.append(value)
        for _, value in stage.submod.named_buffers(remove_duplicate=False):
            buffer_values.extend(flatten_graph_values([value]))
        stage.state.sharded_param_values = sharded_param_values
        stage.state.buffer_values = buffer_values
        stage.state.trainable_params = trainable_params
        stage.state.unsharded_param_values = []
        stage.state.unsharded_param_grads = []
        stage.state.sharded_param_grads = []
        stage._graph_pp_grads_scaled = False

    @staticmethod
    def _initialize_split_grad_accumulators(
        stage: GraphPipelineStage,
        grads: list[Any],
    ) -> None:
        # PP uses runtime-owned slots. SPMD with gradient accumulation bypasses
        # this helper and carries references in state.
        if not stage.state.unsharded_param_grads:
            stage.state.unsharded_param_grads = [None] * len(grads)

    def _ensure_reduced_grads(self, stage: GraphPipelineStage) -> None:
        if stage.state.sharded_param_grads:
            return
        if not any(grad is not None for grad in stage.state.unsharded_param_grads):
            return
        graphs = self.stage_graphs[stage.stage_index]
        stage.state.sharded_param_grads = graphs.reduce_grads(
            stage.state.unsharded_param_grads,
            runtime_validate=stage._runtime_validate,
        )
        _scale_graph_pp_sharded_grads(stage, self.schedule)

    def _accumulate_stage_sharded_grads(self, stage: GraphPipelineStage) -> None:
        self._ensure_reduced_grads(stage)
        if not stage.state.sharded_param_grads:
            return
        graphs = self.stage_graphs[stage.stage_index]
        param_grads = graphs.param_grads_for_accumulation(
            stage.state.sharded_param_grads
        )
        accumulate_param_grads_(stage.state.trainable_params, param_grads)

    def _accumulate_direct_stage_backward_grads(
        self,
        stage: GraphPipelineStage,
        graphs: StageGraphs,
        grads: list[Any],
    ) -> None:
        grad_scale_factor = (
            self.schedule._n_microbatches if self.schedule.scale_grads else 1
        )
        _scale_grad_values_(grads, grad_scale_factor)
        param_grads = graphs.param_grads_for_accumulation(grads)
        accumulate_param_grads_(
            stage.state.trainable_params,
            param_grads,
            clone_grads_to_initialize_param_grad=True,
        )

    def _accumulate_split_stage_backward_grads(
        self,
        stage: GraphPipelineStage,
        graphs: SplitStageGraphs,
        grads: list[Any],
        *,
        grad_reduction_in_backward: bool,
    ) -> None:
        if grad_reduction_in_backward:
            self._accumulate_direct_stage_backward_grads(stage, graphs, grads)
            return

        self._initialize_split_grad_accumulators(stage, grads)
        _accumulate_stage_unsharded_grads(stage, grads)

    def _handle_forward_backward(self, action: _Action, ctx: _PipelineContext) -> None:
        self.ensure_ready(ctx)
        _, stage = _stage_map_and_stage_from_action(self.schedule, action)
        mb_index = action.microbatch_index
        if mb_index is None:
            raise ValueError(
                f"GraphPP {action.computation_type.value} action must have microbatch "
                f"index: {action}"
            )
        args, kwargs, target = _prepare_fwd_user_args(stage, mb_index, ctx)
        graphs = cast(JointStageGraphs, self.stage_graphs[stage.stage_index])
        if not stage.has_backward:
            raise NotImplementedError(
                "GraphPP joint forward/backward does not support forward-only "
                "execution"
            )
        _ensure_unsharded_param_values(stage, graphs)
        initializes_grad_accumulators = (
            action.computation_type == FORWARD_BACKWARD_NOGRADACCUM
        )
        if initializes_grad_accumulators:
            no_grad_accum_graphs = cast(NoGradAccumJointStageGraphs, graphs)
            loss, param_grads = no_grad_accum_graphs.forward_backward_nogradaccum(
                args,
                kwargs,
                target,
                self.loss_kwargs,
                unsharded_param_values=stage.state.unsharded_param_values,
                buffer_values=stage.state.buffer_values,
                runtime_validate=stage._runtime_validate,
            )
        else:
            loss, param_grads = graphs.forward_backward(
                args,
                kwargs,
                target,
                self.loss_kwargs,
                unsharded_param_values=stage.state.unsharded_param_values,
                buffer_values=stage.state.buffer_values,
                grad_accumulators=stage.state.unsharded_param_grads,
                runtime_validate=stage._runtime_validate,
            )
        self.schedule.backward_counter[stage.stage_index] += 1
        if initializes_grad_accumulators:
            stage.state.unsharded_param_grads = param_grads
        elif stage.stage_index not in self._joint_gradient_accumulation_stage_indices:
            self._accumulate_direct_stage_backward_grads(stage, graphs, param_grads)
        stage.output_chunks.append(loss)
        self.schedule._internal_losses.append(loss)

    def _handle_forward_backward_first_with_unshard(
        self, action: _Action, ctx: _PipelineContext
    ) -> None:
        self.ensure_ready(ctx)
        _, stage = _stage_map_and_stage_from_action(self.schedule, action)
        mb_index = action.microbatch_index
        if mb_index is None:
            raise ValueError(
                "GraphRuntime FORWARD_BACKWARD_FIRST_WITH_UNSHARD action must "
                f"have microbatch index: {action}"
            )
        args, kwargs, target = _prepare_fwd_user_args(stage, mb_index, ctx)
        graphs = cast(
            FSDPBoundaryJointStageGraphs,
            self.stage_graphs[stage.stage_index],
        )
        (
            loss,
            param_grads,
            unsharded_param_values,
        ) = graphs.forward_backward_with_unshard(
            args,
            kwargs,
            target,
            self.loss_kwargs,
            sharded_param_values=stage.state.sharded_param_values,
            buffer_values=stage.state.buffer_values,
            runtime_validate=stage._runtime_validate,
        )
        stage.state.unsharded_param_values = unsharded_param_values
        stage.state.unsharded_param_grads = param_grads
        self.schedule.backward_counter[stage.stage_index] += 1
        stage.output_chunks.append(loss)
        self.schedule._internal_losses.append(loss)

    def _handle_forward_backward_last_with_reduce_grad(
        self, action: _Action, ctx: _PipelineContext
    ) -> None:
        self.ensure_ready(ctx)
        _, stage = _stage_map_and_stage_from_action(self.schedule, action)
        mb_index = action.microbatch_index
        if mb_index is None:
            raise ValueError(
                "GraphRuntime FORWARD_BACKWARD_LAST_WITH_REDUCE_GRAD action "
                f"must have microbatch index: {action}"
            )
        args, kwargs, target = _prepare_fwd_user_args(stage, mb_index, ctx)
        graphs = cast(
            FSDPBoundaryJointStageGraphs,
            self.stage_graphs[stage.stage_index],
        )
        _ensure_unsharded_param_values(stage, graphs)
        loss, sharded_param_grads = graphs.forward_backward_with_reduce_grad(
            args,
            kwargs,
            target,
            self.loss_kwargs,
            unsharded_param_values=stage.state.unsharded_param_values,
            buffer_values=stage.state.buffer_values,
            grad_accumulators=stage.state.unsharded_param_grads,
            runtime_validate=stage._runtime_validate,
        )
        stage.state.sharded_param_grads = sharded_param_grads
        _scale_graph_pp_sharded_grads(stage, self.schedule)
        self.schedule.backward_counter[stage.stage_index] += 1
        stage.output_chunks.append(loss)
        self.schedule._internal_losses.append(loss)

    def _handle_forward(self, action: _Action, ctx: _PipelineContext) -> None:
        self.ensure_ready(ctx)
        (
            stage_index_to_stage,
            stage,
            mb_index,
            is_next_stage_on_this_rank,
        ) = _prepare_fwd_common(self.schedule, action)
        args, kwargs, target = _prepare_fwd_user_args(stage, mb_index, ctx)
        graphs = self.stage_graphs[stage.stage_index]
        _ensure_unsharded_param_values(stage, graphs)
        output, saved_values_for_backward = graphs.forward(
            args,
            kwargs,
            target,
            self.loss_kwargs,
            unsharded_param_values=stage.state.unsharded_param_values,
            buffer_values=stage.state.buffer_values,
            activation_slot_id_1=self._resolve_dist_moe_activation_slot(action),
            runtime_validate=stage._runtime_validate,
        )
        _post_fwd_common(
            stage,
            mb_index,
            output,
            saved_values_for_backward,
            self.schedule,
            stage_index_to_stage,
            is_next_stage_on_this_rank,
        )

    def _handle_backward(self, action: _Action, ctx: _PipelineContext) -> None:
        self.ensure_ready(ctx)
        (
            stage_index_to_stage,
            stage,
            mb_index,
            is_prev_stage_on_this_rank,
        ) = _prepare_backward_common(self.schedule, action)
        if not stage.has_backward:
            return
        graphs = cast(SplitStageGraphs, self.stage_graphs[stage.stage_index])
        (
            saved_values_for_backward,
            output_grads_from_next,
        ) = _prepare_backward_values(stage, mb_index)
        input_grads, param_grads = graphs.full_backward(
            saved_values_for_backward,
            output_grads_from_next,
            runtime_validate=stage._runtime_validate,
        )
        self._accumulate_split_stage_backward_grads(
            stage,
            graphs,
            param_grads,
            grad_reduction_in_backward=_grad_reduction_runs_in_backward(action),
        )
        _post_backward_common(
            stage,
            mb_index,
            input_grads,
            stage_index_to_stage,
            is_prev_stage_on_this_rank,
        )

    def _handle_backward_input(self, action: _Action, ctx: _PipelineContext) -> None:
        self.ensure_ready(ctx)
        _, stage = _stage_map_and_stage_from_action(self.schedule, action)
        graphs = cast(SplitStageGraphs, self.stage_graphs[stage.stage_index])
        if not graphs.supports_backward_input_weight_split:
            logger.debug(
                "GraphPP skipping BACKWARD_INPUT for stage %s", stage.stage_index
            )
            return
        (
            stage_index_to_stage,
            stage,
            mb_index,
            is_prev_stage_on_this_rank,
        ) = _prepare_backward_common(self.schedule, action)
        if not stage.has_backward:
            return
        graphs = cast(SplitStageGraphs, self.stage_graphs[stage.stage_index])
        (
            saved_values_for_backward,
            output_grads_from_next,
        ) = _prepare_backward_values(stage, mb_index)
        input_grads, saved_values_for_backward_weight = graphs.backward_input(
            saved_values_for_backward,
            output_grads_from_next,
            runtime_validate=stage._runtime_validate,
        )
        stage.saved_values_for_backward_weight_cache[
            mb_index
        ] = saved_values_for_backward_weight
        _post_backward_common(
            stage,
            mb_index,
            input_grads,
            stage_index_to_stage,
            is_prev_stage_on_this_rank,
        )

    def _handle_backward_weight(self, action: _Action, ctx: _PipelineContext) -> None:
        self.ensure_ready(ctx)
        _, stage = _stage_map_and_stage_from_action(self.schedule, action)
        mb_index = action.microbatch_index
        if mb_index is None:
            raise ValueError(
                f"GraphPP BACKWARD_WEIGHT action must have microbatch index: {action}"
            )
        graphs = cast(SplitStageGraphs, self.stage_graphs[stage.stage_index])
        if not graphs.supports_backward_input_weight_split:
            backward_type = (
                BACKWARD_WITH_REDUCE_GRAD
                if _grad_reduction_runs_in_backward(action)
                else BACKWARD
            )
            new_action = _Action(
                action.stage_index,
                cast(Any, backward_type),
                action.microbatch_index,
                action.sub_actions,
            )
            self._handle_backward(new_action, ctx)
            return
        if not stage.has_backward:
            return
        saved_values_for_backward_weight = (
            stage.saved_values_for_backward_weight_cache.pop(mb_index)
        )
        param_grads = graphs.backward_weight(saved_values_for_backward_weight)
        self._accumulate_split_stage_backward_grads(
            stage,
            graphs,
            param_grads,
            grad_reduction_in_backward=_grad_reduction_runs_in_backward(action),
        )

    def _handle_unshard(self, action: _Action, ctx: _PipelineContext) -> None:
        self.ensure_ready(ctx)
        _, stage = _stage_map_and_stage_from_action(self.schedule, action)
        graphs = self.stage_graphs[stage.stage_index]
        _ensure_unsharded_param_values(stage, graphs)

    def _handle_reshard(self, action: _Action, ctx: _PipelineContext) -> None:
        self.ensure_ready(ctx)
        _, stage = _stage_map_and_stage_from_action(self.schedule, action)
        stage.state.unsharded_param_values = []

    def _handle_reduce_grad(self, action: _Action, ctx: _PipelineContext) -> None:
        self.ensure_ready(ctx)
        _, stage = _stage_map_and_stage_from_action(self.schedule, action)
        self._ensure_reduced_grads(stage)

    def _handle_wait_reduce_grad(
        self,
        action: _Action,
        ctx: _PipelineContext,
    ) -> None:
        """Skip the eager FSDP wait after an explicit reduction graph.

        GraphPP reduction graphs return the reduced tensors directly and do
        not create ``PipelineStage._gradient_reduction_handle``. Tensor
        dependencies carry the collective ordering into later graph work.
        """
        del action, ctx

    def _handle_overlap_fw_bw(self, action: _Action, ctx: _PipelineContext) -> None:
        fw_action, bw_action = overlap_fw_bw_sub_actions(
            action,
            backward_computation_types=(
                BACKWARD,
                BACKWARD_WITH_REDUCE_GRAD,
            ),
        )

        self.ensure_ready(ctx)
        (
            stage_index_to_stage,
            fw_stage,
            fw_mb_index,
            fw_is_next_stage_on_this_rank,
        ) = _prepare_fwd_common(self.schedule, fw_action)
        (
            _,
            bw_stage,
            bw_mb_index,
            bw_is_prev_stage_on_this_rank,
        ) = _prepare_backward_common(self.schedule, bw_action)
        if not bw_stage.has_backward:
            return

        args, kwargs, target = _prepare_fwd_user_args(fw_stage, fw_mb_index, ctx)
        fw_graphs = cast(SplitStageGraphs, self.stage_graphs[fw_stage.stage_index])
        bw_graphs = cast(SplitStageGraphs, self.stage_graphs[bw_stage.stage_index])
        _ensure_unsharded_param_values(fw_stage, fw_graphs)
        pair = (fw_action.stage_index, bw_action.stage_index)
        # The multiplexed graph is runtime-owned state because it is built once
        # from the graph provider and reused across OVERLAP_F_B actions.
        overlap_graph = self.overlap_graphs.get(pair)
        if overlap_graph is None:
            raise ValueError(
                "GraphPP overlap graph must be built before OVERLAP_F_B runtime "
                f"execution for pair {pair}."
            )
        (
            bw_saved_values_for_backward,
            output_grads_from_next,
        ) = _prepare_backward_values(bw_stage, bw_mb_index)
        (
            input_grads,
            param_grads,
            output,
            saved_values_for_backward,
        ) = overlap_graph.forward_backward(
            backward_saved_values_for_backward=bw_saved_values_for_backward,
            output_grads_from_next=output_grads_from_next,
            forward_args=args,
            forward_kwargs=kwargs,
            forward_target=target,
            forward_loss_kwargs=self.loss_kwargs,
            forward_unsharded_param_values=fw_stage.state.unsharded_param_values,
            forward_buffer_values=fw_stage.state.buffer_values,
            forward_activation_slot_id_1=(
                self._resolve_dist_moe_activation_slot(fw_action)
            ),
            runtime_validate=(fw_stage._runtime_validate or bw_stage._runtime_validate),
        )

        self._accumulate_split_stage_backward_grads(
            bw_stage,
            bw_graphs,
            param_grads,
            grad_reduction_in_backward=_grad_reduction_runs_in_backward(bw_action),
        )
        _post_fwd_common(
            fw_stage,
            fw_mb_index,
            output,
            saved_values_for_backward,
            self.schedule,
            stage_index_to_stage,
            fw_is_next_stage_on_this_rank,
        )
        _post_backward_common(
            bw_stage,
            bw_mb_index,
            input_grads,
            stage_index_to_stage,
            bw_is_prev_stage_on_this_rank,
        )

    def _skip_spmd_stage_initialization(self, *, has_backward: bool) -> None:
        """Skip pipeline-only stage initialization for SPMD schedule,
        as it does not require any pipeline communication buffers.
        """
        if not self.is_spmd:
            return
        self.schedule._stages_forward_initialized = True
        self.schedule._stages_backward_initialized = has_backward

    def step(self, *args: Any, **kwargs: Any) -> None:
        """Run one training step through the wrapped pipeline schedule.

        Args:
            *args (Any): Positional arguments forwarded to ``schedule.step``.
            **kwargs (Any): Keyword arguments forwarded to ``schedule.step``.
                GraphPP reads ``loss_kwargs`` from this mapping and forwards
                all kwargs to the upstream schedule unchanged.
        """
        self._graph_pp_ready = False
        self.loss_kwargs = kwargs.get("loss_kwargs") or {}
        step_succeeded = False
        try:
            self._skip_spmd_stage_initialization(has_backward=True)
            self.schedule.step(*args, **kwargs)
            step_succeeded = True
        finally:
            for stage in self.schedule._stages:
                graph_stage = cast(GraphPipelineStage, stage)
                if step_succeeded:
                    self._accumulate_stage_sharded_grads(graph_stage)
                graph_stage.state.clear()
                graph_stage.clear_runtime_states()
            self.loss_kwargs = {}
            self.stage_graphs = {}
            self._graph_pp_ready = False

    def eval(self, *args: Any, **kwargs: Any) -> Any:
        """Run evaluation through the wrapped pipeline schedule.

        Evaluation reuses upstream PP eval semantics, which disable backward
        and delegate to ``schedule.step``. GraphPP only mirrors the per-step
        runtime-state setup and cleanup from training, without accumulating
        parameter gradients.

        Args:
            *args (Any): Positional arguments forwarded to ``schedule.eval``.
            **kwargs (Any): Keyword arguments forwarded to ``schedule.eval``.

        Returns:
            Any: The value returned by ``schedule.eval``.
        """
        self._graph_pp_ready = False
        self.loss_kwargs = kwargs.get("loss_kwargs") or {}
        try:
            self._skip_spmd_stage_initialization(has_backward=False)
            return self.schedule.eval(*args, **kwargs)
        finally:
            for stage in self.schedule._stages:
                graph_stage = cast(GraphPipelineStage, stage)
                graph_stage.state.clear()
                graph_stage.clear_runtime_states()
            self.loss_kwargs = {}
            self.stage_graphs = {}
            self._graph_pp_ready = False


def register_graph_schedule(
    schedule: _PipelineScheduleRuntime,
    *,
    graph_provider: StageGraphsProvider | None = None,
    is_spmd: bool = False,
    liveness_schedule: _PipelineScheduleRuntime | None = None,
) -> GraphRuntime:
    """Register graph action handlers on a runtime schedule.

    Args:
        schedule (_PipelineScheduleRuntime): Runtime pipeline schedule whose
            compute actions should be handled by GraphPP.
        graph_provider (StageGraphsProvider | None): Optional provider
            that builds or attaches stage graphs before the first runtime
            action in each step.
        is_spmd (bool): Whether this schedule is SPMD and does not require
            any PP only processing (e.g. pipeline comms buffers).
        liveness_schedule: Optional pre-rewrite schedule retained for
            activation-liveness analysis.
    Returns:
        GraphRuntime: Runtime that owns the registered bound action handlers.

    Raises:
        TypeError: If any local schedule stage is not a ``GraphPipelineStage``.
    """
    runtime = GraphRuntime(
        schedule,
        graph_provider=graph_provider,
        is_spmd=is_spmd,
        liveness_schedule=liveness_schedule,
    )
    # Calling convention:
    # Upstream computation types use PyTorch's validated
    # (action, context) -> None handler API.
    for computation_type, handler in (
        (FORWARD, runtime._handle_forward),
        (UNSHARD, runtime._handle_unshard),
        (RESHARD, runtime._handle_reshard),
        (REDUCE_GRAD, runtime._handle_reduce_grad),
        (WAIT_REDUCE_GRAD, runtime._handle_wait_reduce_grad),
        (BACKWARD_INPUT, runtime._handle_backward_input),
        (BACKWARD_WEIGHT, runtime._handle_backward_weight),
        (OVERLAP_F_B, runtime._handle_overlap_fw_bw),
    ):
        schedule.register_custom_function(computation_type, handler)
    # Calling convention:
    # GraphRuntime-only computation types use the same
    # (action, context) -> None handlers but are outside PyTorch's accepted
    # computation types.
    for computation_type, handler in (
        (BACKWARD, runtime._handle_backward),
        (BACKWARD_WITH_REDUCE_GRAD, runtime._handle_backward),
        (BACKWARD_WEIGHT_WITH_REDUCE_GRAD, runtime._handle_backward_weight),
        (FORWARD_BACKWARD, runtime._handle_forward_backward),
        (FORWARD_BACKWARD_NOGRADACCUM, runtime._handle_forward_backward),
        (FULL_FORWARD_BACKWARD, runtime._handle_forward_backward),
        (
            FORWARD_BACKWARD_FIRST_WITH_UNSHARD,
            runtime._handle_forward_backward_first_with_unshard,
        ),
        (
            FORWARD_BACKWARD_LAST_WITH_REDUCE_GRAD,
            runtime._handle_forward_backward_last_with_reduce_grad,
        ),
    ):
        schedule._comp_type_to_function_map[
            computation_type
        ] = handler  # pyrefly: ignore[bad-assignment]
    return runtime
