TensorPlay
latest (dev)
Copy
View Markdown

Latest development documentation · Updated 2026-10-08

Source code for tensorplay.distributed.pipelining.stage

"""Pipeline stage execution and metadata management."""

from abc import ABC
from dataclasses import dataclass
import operator
from typing import Any, Callable

import tensorplay as tp
from .. import config as dist_config
from .. import distributed_core as dist

from ._backward import (
    _autograd_grad_for_inputs,
    stage_backward,
    stage_backward_input,
    stage_backward_weight,
)
from ._utils import (
    _MeshCache,
    PipeliningMetadataError,
    _StageBackwardMeta,
    _StageForwardMeta,
    _StageMeta,
    _DTensorMeta,
    _TensorMeta,
    _derive_grad_metas,
    flatten_args,
    _make_tensor_from_meta,
    InferenceMode,
    extract_tensor_meta,
    extract_tensor_metas,
    to_local_if_dtensor,
    validate_static_arg_grad_correspondence,
    validate_tensors_metadata,
)
from ..tensor import DTensor

__all__ = ["PipelineStage", "build_stage"]


def _normalize_model_output_as_tuple(output: Any) -> tuple[Any, ...]:
    if isinstance(output, list):
        return tuple(output)
    return output if isinstance(output, tuple) else (output,)


@dataclass
class _RecvInfo:
    input_name: str
    source: int | None
    buffer: Any
    tensor_meta: Any
    is_root_arg: bool = False

    def __init__(
        self,
        input_name: str,
        source: int | None,
        buffer: Any,
        tensor_meta: Any,
        is_root_arg: bool = False,
    ) -> None:
        self.input_name = input_name
        self.source = source
        self.buffer = buffer
        self.tensor_meta = tensor_meta
        self.is_root_arg = is_root_arg

    def __repr__(self) -> str:
        if self.is_root_arg:
            return f"_RecvInfo(input={self.input_name}, root_arg=True)"
        meta_type = type(self.tensor_meta).__name__ if self.tensor_meta else "None"
        buffer_shape = self.buffer.size() if self.buffer is not None else "None"
        return f"_RecvInfo(input={self.input_name}, source={self.source}, shape={buffer_shape}, meta={meta_type})"


def _build_p2p_direction_groups(group: Any) -> tuple[Any, Any]:
    if not dist.is_initialized():
        return group, group
    parent = group if group is not None else dist._get_default_group()
    if parent.size() <= 1:
        return group, group
    cache = getattr(_build_p2p_direction_groups, "_cache", None)
    if cache is None:
        cache = _build_p2p_direction_groups._cache = {}
    key = id(parent)
    cached = cache.get(key)
    if cached is not None and cached[0] is parent:
        return cached[1], cached[2]
    split_ranks = [list(range(parent.size()))]
    downstream = dist.split_group(
        parent_pg=parent,
        split_ranks=split_ranks,
        group_desc="pipeline_downstream",
    )
    upstream = dist.split_group(
        parent_pg=parent,
        split_ranks=split_ranks,
        group_desc="pipeline_upstream",
    )
    if downstream is dist.GroupMember.NON_GROUP_MEMBER or upstream is dist.GroupMember.NON_GROUP_MEMBER:
        raise RuntimeError("pipeline direction groups must contain the current rank")
    cache[key] = (parent, downstream, upstream)
    return downstream, upstream


class _PipelineStageBase(ABC):
    def __init__(self, submodule: Any, stage_index: int, num_stages: int, device: Any = None, group: Any = None, dw_builder: Callable[[], Callable[..., None]] | None = None) -> None:
        if stage_index < 0 or stage_index >= num_stages:
            raise ValueError("stage_index is outside the pipeline")
        self.submod = submodule
        self.stage_index = stage_index
        self.num_stages = num_stages
        self.device = device
        self.group = group
        self.dw_builder = dw_builder
        self.p2p_per_direction = bool(dist_config.pipeline_per_direction_p2p)
        if self.p2p_per_direction:
            self._downstream_group, self._upstream_group = _build_p2p_direction_groups(group)
        else:
            self._downstream_group = group
            self._upstream_group = group
        try:
            self.group_rank = int(dist.get_rank(group)) if dist.is_initialized() else stage_index
            self.group_size = int(dist.get_world_size(group)) if dist.is_initialized() else num_stages
        except (RuntimeError, ValueError):
            self.group_rank = stage_index
            self.group_size = num_stages
        if self.group_size > num_stages:
            raise ValueError("pipeline group cannot contain more ranks than stages")
        self.stage_index_to_group_rank = {
            index: index % self.group_size for index in range(num_stages)
        }
        self._has_backward = False
        self.fwd_cache: dict[int, tuple[Any, tuple[Any, ...]]] = {}
        self.bwd_cache: dict[int, Any] = {}
        self.output_chunks: list[Any] = []
        self.args_recv_info: dict[int, tuple[_RecvInfo, ...]] = {}
        self.act_send_info: dict[int, list[Any]] = {}
        self.grad_recv_info: dict[int, tuple[_RecvInfo, ...]] = {}
        self.grad_send_info: list[Any] | None = None
        self.chunks: int | None = None
        self._stage_meta = _StageMeta()
        self._mesh_cache = _MeshCache()
        self._input_chunks: dict[int, tuple[Any, ...]] = {}
        self._forward_inputs: dict[int, tuple[Any, ...]] = {}
        self.backward_state: dict[int, tuple[Any, Any, Any, Any]] = {}
        self.dw_runner: dict[int, Callable[[], Any]] = {}

    @property
    def has_backward(self) -> bool:
        return self._has_backward

    @has_backward.setter
    def has_backward(self, value: bool) -> None:
        self._has_backward = bool(value)

    @property
    def is_first(self) -> bool:
        return self.stage_index == 0

    @property
    def is_last(self) -> bool:
        return self.stage_index == self.num_stages - 1

    def _validate_stage_tensors(self, desc: str, expected: tuple[Any, ...] | None, actual: tuple[Any, ...]) -> None:
        if expected is None:
            raise PipeliningMetadataError(f"{desc}: metadata is unavailable")
        validate_tensors_metadata(desc, expected, actual)

    def _check_chunk_id(self, chunk_id: int) -> None:
        if self.chunks is None or chunk_id < 0 or chunk_id >= self.chunks:
            raise RuntimeError("chunk id is outside the configured range")

    def _create_grad_send_info(self, args_recv_info: tuple[_RecvInfo, ...]) -> list[Any]:
        return [item.source if isinstance(item, _RecvInfo) else None for item in args_recv_info]

    def _prepare_forward_infra(self, num_microbatches: int, args: Any, kwargs: Any, has_backward: bool) -> Any:
        self.chunks = num_microbatches
        self.has_backward = has_backward
        self._stage_meta.forward.input_metas = tuple(meta for meta in (extract_tensor_meta(value) for value in args) if meta is not None)
        self.args_recv_info = {index: tuple(_RecvInfo(str(pos), None, None, extract_tensor_meta(value), True) for pos, value in enumerate(args)) for index in range(num_microbatches)}

    def _prepare_backward_infra(
        self,
        num_microbatches: int,
        loss_fn: Any = None,
        target: Any = None,
        received_grad_meta: Any = None,
        loss_kwargs: Any = None,
    ) -> None:
        del loss_fn, target, loss_kwargs
        self.chunks = num_microbatches
        self.has_backward = True
        self._stage_meta.backward.output_grad_metas = tuple(received_grad_meta or ())
        self.grad_recv_info = {
            index: self._create_grad_recv_info(self.act_send_info)
            for index in range(num_microbatches)
        }
        self.grad_send_info = self._create_grad_send_info(
            self.args_recv_info.get(0, ())
        )

    def _setup_backward_recv_info(self, num_microbatches: int) -> None:
        self.chunks = num_microbatches
        self.grad_recv_info = {
            index: self._create_grad_recv_info(self.act_send_info)
            for index in range(num_microbatches)
        }

    def _create_grad_recv_info(self, act_send_info: Any) -> tuple[_RecvInfo, ...]:
        del act_send_info
        return ()

    def _resolve_peer_global_rank(self, stage_idx: int) -> int:
        peer_group_rank = self.stage_index_to_group_rank[int(stage_idx)]
        if self.group is None:
            return int(peer_group_rank)
        return int(dist.get_global_rank(self.group, peer_group_rank))

    def _get_recv_ops(self, recv_infos: Any, group: Any) -> list[Any]:
        if not dist.is_initialized():
            return []
        process_group = self.group if group is None else group
        operations = []
        for info in recv_infos:
            if not isinstance(info, _RecvInfo) or info.source is None or info.buffer is None:
                continue
            peer_group_rank = self.stage_index_to_group_rank[int(info.source)]
            peer = (
                peer_group_rank
                if process_group is None
                else dist.get_global_rank(process_group, peer_group_rank)
            )
            operations.append(dist.P2POp(dist.irecv, info.buffer, peer, process_group))
        return operations

    def set_local_fwd_input(self, prev_stage_outputs: Any, mb_index: int) -> None:
        values = _normalize_model_output_as_tuple(prev_stage_outputs)
        recv_infos = self.args_recv_info[mb_index]
        if len(recv_infos) != len(values):
            raise PipeliningMetadataError(
                f"Stage {self.stage_index}: local forward input count does not match "
                f"the receive metadata ({len(values)} != {len(recv_infos)})"
            )
        if self.is_first:
            raise AssertionError("local forward input is only valid for a non-first stage")
        for info, value in zip(recv_infos, values, strict=True):
            if info.is_root_arg:
                raise AssertionError("local forward input cannot replace a root argument")
            local_value = to_local_if_dtensor(value)
            if isinstance(local_value, tp.Tensor):
                local_value = local_value.detach()
                if (
                    info.tensor_meta is not None
                    and info.tensor_meta.requires_grad
                    and (local_value.is_floating_point() or local_value.is_complex())
                ):
                    local_value.requires_grad_(True)
            info.buffer = local_value
        self._input_chunks[mb_index] = tuple(info.buffer for info in recv_infos)

    def get_local_bwd_output(self, mb_index: int) -> Any:
        if not self.has_backward:
            raise AssertionError("cannot get a backward output without backward enabled")
        if self.is_first:
            raise AssertionError("the first stage has no local backward output")
        self._check_chunk_id(mb_index)
        return self.bwd_cache.pop(mb_index)

    def set_local_bwd_input(self, next_stage_bwd_outputs: Any, mb_index: int) -> None:
        values = next_stage_bwd_outputs
        if not isinstance(values, tuple):
            raise AssertionError(f"expected a tuple of gradients, got {type(values)}")
        if not self.has_backward:
            raise AssertionError("cannot set a backward input without backward enabled")
        if self.is_last:
            raise AssertionError("the last stage has no local backward input")
        recv_infos = self.grad_recv_info[mb_index]
        if len(recv_infos) != len(values):
            raise PipeliningMetadataError(
                f"Stage {self.stage_index}: local backward input count does not match "
                f"the receive metadata ({len(values)} != {len(recv_infos)})"
            )
        for info, value in zip(recv_infos, values, strict=True):
            if value is None:
                if info.buffer is not None:
                    info.buffer.zero_()
                continue
            if info.is_root_arg:
                raise AssertionError("local backward input cannot target a root argument")
            info.buffer = to_local_if_dtensor(value)

    def get_fwd_recv_ops(self, fwd_chunk_id: int) -> list[Any]:
        self._check_chunk_id(fwd_chunk_id)
        return self._get_recv_ops(
            self.args_recv_info.get(fwd_chunk_id, ()), self._downstream_group
        )

    def get_bwd_recv_ops(self, bwd_chunk_id: int) -> list[Any]:
        self._check_chunk_id(bwd_chunk_id)
        if not self.has_backward or self.is_last:
            return []
        return self._get_recv_ops(
            self.grad_recv_info.get(bwd_chunk_id, ()), self._upstream_group
        )

    def get_fwd_send_ops(self, fwd_chunk_id: int) -> list[Any]:
        self._check_chunk_id(fwd_chunk_id)
        output = self.fwd_cache[fwd_chunk_id][0]
        values = _normalize_model_output_as_tuple(output)
        operations = []
        for index, value in enumerate(values):
            for destination in self.act_send_info.get(index, ()):
                if destination is None:
                    continue
                value = to_local_if_dtensor(value, detach=True)
                if not isinstance(value, tp.Tensor):
                    raise PipeliningMetadataError(
                        f"Stage {self.stage_index}: activation {index} is not a tensor"
                    )
                peer_group_rank = self.stage_index_to_group_rank[int(destination)]
                peer = (
                    peer_group_rank
                    if self._downstream_group is None
                    else dist.get_global_rank(self._downstream_group, peer_group_rank)
                )
                operations.append(
                    dist.P2POp(dist.isend, value, peer, self._downstream_group)
                )
        return operations

    def _get_grad_send_meta(self, input_idx: int) -> Any:
        input_grads = self._stage_meta.input_grads
        if input_grads is not None and input_idx < len(input_grads):
            return input_grads[input_idx]
        inputs = self._stage_meta.inputs
        if inputs is not None and input_idx < len(inputs):
            meta = inputs[input_idx]
            if meta is not None:
                return _derive_grad_metas((meta,))[0]
        raise PipeliningMetadataError(
            f"Stage {self.stage_index}: backward produced a gradient for input "
            f"{input_idx}, but no gradient metadata is available"
        )

    def get_bwd_send_ops(self, bwd_chunk_id: int) -> list[Any]:
        self._check_chunk_id(bwd_chunk_id)
        if not self.has_backward or self.is_first:
            return []
        if self.grad_send_info is None:
            self.grad_send_info = self._create_grad_send_info(
                self.args_recv_info.get(bwd_chunk_id, ())
            )
        gradients = self.bwd_cache.pop(bwd_chunk_id, ())
        operations = []
        for index, (gradient, destination) in enumerate(
            zip(gradients or (), self.grad_send_info, strict=True)
        ):
            if destination is None:
                if gradient is not None:
                    raise PipeliningMetadataError(
                        f"Stage {self.stage_index}: input {index} has a gradient but "
                        "no previous stage receives it"
                    )
                continue
            grad_meta = self._get_grad_send_meta(index)
            if grad_meta is None:
                if gradient is not None:
                    raise PipeliningMetadataError(
                        f"Stage {self.stage_index}: input {index} produced a gradient "
                        "without gradient metadata"
                    )
                continue
            if gradient is None:
                send_tensor = _make_tensor_from_meta(grad_meta, self.device).zero_()
            else:
                send_tensor = to_local_if_dtensor(gradient, detach=True)
                if not isinstance(send_tensor, tp.Tensor):
                    raise PipeliningMetadataError(
                        f"Stage {self.stage_index}: input {index} gradient is not a tensor"
                    )
            peer_group_rank = self.stage_index_to_group_rank[int(destination)]
            peer = (
                peer_group_rank
                if self._upstream_group is None
                else dist.get_global_rank(self._upstream_group, peer_group_rank)
            )
            operations.append(
                dist.P2POp(dist.isend, send_tensor, peer, self._upstream_group)
            )
        return operations

    def clear_runtime_states(self) -> None:
        self.fwd_cache.clear()
        self.bwd_cache.clear()
        self.output_chunks.clear()
        self._input_chunks.clear()
        self._forward_inputs.clear()
        self.backward_state.clear()
        self.dw_runner.clear()
        for recv_infos in self.args_recv_info.values():
            for info in recv_infos:
                if not info.is_root_arg and isinstance(info.buffer, tp.Tensor):
                    info.buffer.grad = None

    def _map_tensor_from_recv_info(self, recv_infos: Any) -> tuple[Any, ...]:
        values = []
        for item in recv_infos:
            if item.is_root_arg:
                raise PipeliningMetadataError("root arguments are not received tensors")
            values.append(item.buffer)
        return tuple(values)

    def _retrieve_recv_activations(self, fwd_chunk_id: int) -> tuple[Any, ...]:
        recv_infos = self.args_recv_info.get(fwd_chunk_id, ())
        values = []
        for index, info in enumerate(recv_infos):
            if info.is_root_arg:
                raise PipeliningMetadataError(
                    f"Stage {self.stage_index}: root input cannot be received"
                )
            if info.buffer is None or info.tensor_meta is None:
                raise PipeliningMetadataError(
                    f"Stage {self.stage_index}: activation {index} has no receive buffer or metadata"
                )
            effective_requires_grad = bool(
                info.tensor_meta.requires_grad
                and self.has_backward
                and tp.is_grad_enabled()
            )
            if isinstance(info.tensor_meta, _DTensorMeta):
                local = info.buffer
                if not isinstance(local, tp.Tensor):
                    raise PipeliningMetadataError(
                        f"Stage {self.stage_index}: DTensor activation buffer is not a tensor"
                    )
                local = local.detach()
                if effective_requires_grad and (
                    local.is_floating_point() or local.is_complex()
                ):
                    local.requires_grad_(True)
                mesh = self._mesh_cache.get_mesh(info.tensor_meta.mesh_cache_key)
                values.append(
                    DTensor.from_local(
                        local,
                        device_mesh=mesh,
                        placements=info.tensor_meta.placements,
                        shape=info.tensor_meta.global_shape,
                        stride=info.tensor_meta.global_stride,
                        run_check=False,
                    )
                )
            else:
                value = info.buffer
                if not isinstance(value, tp.Tensor):
                    raise PipeliningMetadataError(
                        f"Stage {self.stage_index}: activation {index} is not a tensor"
                    )
                value.requires_grad_(
                    effective_requires_grad
                    and (value.is_floating_point() or value.is_complex())
                )
                values.append(value)
        return tuple(values)

    def _retrieve_recv_grads(self, bwd_chunk_id: int) -> tuple[Any, ...]:
        recv_infos = self.grad_recv_info.get(bwd_chunk_id, ())
        values = []
        for index, info in enumerate(recv_infos):
            if info.is_root_arg:
                raise PipeliningMetadataError(
                    f"Stage {self.stage_index}: root input cannot receive a gradient"
                )
            if info.buffer is None:
                if info.tensor_meta is not None:
                    raise PipeliningMetadataError(
                        f"Stage {self.stage_index}: gradient {index} has metadata but no buffer"
                    )
                values.append(None)
                continue
            if info.tensor_meta is None:
                raise PipeliningMetadataError(
                    f"Stage {self.stage_index}: gradient {index} has a buffer but no metadata"
                )
            if isinstance(info.tensor_meta, _DTensorMeta):
                mesh = self._mesh_cache.get_mesh(info.tensor_meta.mesh_cache_key)
                values.append(
                    DTensor.from_local(
                        info.buffer,
                        device_mesh=mesh,
                        placements=info.tensor_meta.placements,
                        shape=info.tensor_meta.global_shape,
                        stride=info.tensor_meta.global_stride,
                        run_check=False,
                    )
                )
            else:
                values.append(info.buffer)
        return tuple(values)

    def forward_maybe_with_nosync(self, *args: Any, **kwargs: Any) -> Any:
        from ...nn.parallel.distributed import DistributedDataParallel

        if isinstance(self.submod, DistributedDataParallel):
            with self.submod.no_sync():
                return self.submod(*args, **kwargs)
        return self.submod(*args, **kwargs)

    def scale_grads(self, grad_scale_factor: float) -> None:
        for param in self.submod.parameters():
            if getattr(param, "grad", None) is not None:
                param.grad.div_(grad_scale_factor)

    def backward_maybe_with_nosync(self, backward_type: Any, bwd_kwargs: dict[str, Any], last_backward: bool = False) -> Any:
        del last_backward

        fsdp_flags = (
            ("set_is_last_backward", False),
            ("set_reshard_after_backward", False),
            ("set_requires_gradient_sync", False),
        )
        for method_name, value in fsdp_flags:
            method = getattr(self.submod, method_name, None)
            if callable(method):
                method(value)
        if backward_type == "full":
            return stage_backward(
                bwd_kwargs["stage_output"],
                bwd_kwargs["output_grads"],
                bwd_kwargs["input_values"],
            ), None
        if backward_type == "input":
            return stage_backward_input(
                bwd_kwargs["stage_output"],
                bwd_kwargs["output_grads"],
                bwd_kwargs["input_values"],
                self.submod.parameters(),
            )
        if backward_type == "weight":
            return stage_backward_weight(
                self.submod.parameters(),
                bwd_kwargs["param_groups"] or [],
            ), None
        raise RuntimeError(f"unknown backward type {backward_type!r}")

    def forward_one_chunk(self, fwd_chunk_id: int, args: tuple[Any, ...], kwargs: dict[str, Any], save_forward_output: bool = True) -> Any:
        self._check_chunk_id(fwd_chunk_id)
        composite_args = args if self.is_first else self._retrieve_recv_activations(fwd_chunk_id)
        output = self.forward_maybe_with_nosync(*composite_args, **kwargs)
        self._forward_inputs[fwd_chunk_id] = tuple(
            value
            for value in flatten_args(composite_args)
            if isinstance(value, tp.Tensor) or value is not None
        ) + tuple(
            value
            for value in flatten_args(kwargs)
            if isinstance(value, tp.Tensor) or value is not None
        )
        output_tuple = _normalize_model_output_as_tuple(output)
        self.fwd_cache[fwd_chunk_id] = (output, output_tuple)
        if save_forward_output:
            while len(self.output_chunks) <= fwd_chunk_id:
                self.output_chunks.append(None)
            self.output_chunks[fwd_chunk_id] = output
        self._stage_meta.forward.output_metas = tuple(meta for meta in (extract_tensor_meta(value) for value in output_tuple) if meta is not None)
        return output

    def backward_one_chunk(self, bwd_chunk_id: int, loss: Any = None, full_backward: bool = True, last_backward: bool = False) -> Any:
        if not self.has_backward:
            return None
        self._check_chunk_id(bwd_chunk_id)
        output, output_values = self.fwd_cache.pop(bwd_chunk_id)
        if self.is_last:
            stage_output = output if loss is None else loss
            output_grads = None
        else:
            stage_output = output_values
            output_grads = self._retrieve_recv_grads(bwd_chunk_id)
        input_values = self._forward_inputs.pop(bwd_chunk_id, ())
        bwd_kwargs = {
            "stage_output": stage_output,
            "output_grads": output_grads,
            "input_values": input_values,
        }
        grads_input: tuple[Any, ...] = ()
        if self.dw_builder is not None:
            grads_input, _ = self.backward_maybe_with_nosync(
                "full", bwd_kwargs, last_backward=last_backward
            )
            if full_backward:
                self.dw_builder()()
            else:
                self.dw_runner[bwd_chunk_id] = self.dw_builder()
        elif full_backward:
            grads_input, _ = self.backward_maybe_with_nosync(
                "full", bwd_kwargs, last_backward=last_backward
            )
        else:
            param_groups = None
            if not self.is_first:
                grads_input, param_groups = self.backward_maybe_with_nosync(
                    "input", bwd_kwargs, last_backward=last_backward
                )
            self.backward_state[bwd_chunk_id] = (
                input_values,
                param_groups,
                stage_output,
                output_grads,
            )
            self.dw_runner[bwd_chunk_id] = lambda: None
        num_forward_inputs = len(self._stage_meta.inputs or ())
        self.bwd_cache[bwd_chunk_id] = tuple(grads_input[:num_forward_inputs])
        return self.bwd_cache[bwd_chunk_id]

    def backward_weight_one_chunk(self, bwd_chunk_id: int, last_backward: bool = False) -> Any:
        if not self.has_backward:
            return None
        runner = self.dw_runner.pop(bwd_chunk_id, None)
        if runner is None:
            raise AssertionError(
                f"backward weight requested for chunk {bwd_chunk_id} without input backward"
            )
        if self.dw_builder is not None:
            return runner()
        input_values, param_groups, stage_output, output_grads = self.backward_state.pop(
            bwd_chunk_id
        )
        if self.is_first:
            self.backward_maybe_with_nosync(
                "full",
                {
                    "stage_output": stage_output,
                    "output_grads": output_grads,
                    "input_values": input_values,
                },
                last_backward=last_backward,
            )
        else:
            self.backward_maybe_with_nosync(
                "weight",
                {"param_groups": param_groups},
                last_backward=last_backward,
            )
        return None

    def _get_init_p2p_neighbors_ops(self) -> list[Any]:
        operations: list[Any] = []
        next_stage_peer_rank = self.stage_index_to_group_rank.get(
            self.stage_index + 1
        )
        previous_stage_peer_rank = self.stage_index_to_group_rank.get(
            self.stage_index - 1
        )
        downstream_recv_tensor = tp.zeros(
            1, device=self.device, dtype=tp.float32
        )
        upstream_recv_tensor = tp.zeros(
            1, device=self.device, dtype=tp.float32
        )
        send_tensor = tp.tensor(
            self.stage_index, device=self.device, dtype=tp.float32
        )
        if not self.is_first:
            operations.append(
                dist.P2POp(
                    dist.irecv,
                    downstream_recv_tensor,
                    group_peer=previous_stage_peer_rank,
                    group=self._downstream_group,
                )
            )
        if not self.is_last:
            operations.append(
                dist.P2POp(
                    dist.isend,
                    send_tensor,
                    group_peer=next_stage_peer_rank,
                    group=self._downstream_group,
                )
            )
        if not self.is_first:
            operations.append(
                dist.P2POp(
                    dist.isend,
                    send_tensor,
                    group_peer=previous_stage_peer_rank,
                    group=self._upstream_group,
                )
            )
        if not self.is_last:
            operations.append(
                dist.P2POp(
                    dist.irecv,
                    upstream_recv_tensor,
                    group_peer=next_stage_peer_rank,
                    group=self._upstream_group,
                )
            )
        return operations

    def perform_reduce_grad(self, grad_scale_factor: float) -> None:
        state_getter = getattr(self.submod, "_get_fsdp_state", None)
        if not callable(state_getter):
            state_getter = getattr(self.submod, "_get_replicate_state", None)
        if callable(state_getter):
            for method_name, value in (
                ("set_is_last_backward", True),
                ("set_reshard_after_backward", True),
                ("set_requires_gradient_sync", True),
            ):
                method = getattr(self.submod, method_name, None)
                if callable(method):
                    method(value)
            state = state_getter()
            state_context = getattr(state, "_state_ctx", None)
            states = (
                getattr(state_context, "all_states", None)
                or getattr(state_context, "states", None)
                or [state]
            )
            for state_item in states:
                groups_getter = getattr(state_item, "_all_param_groups", None)
                if callable(groups_getter):
                    for param_group in groups_getter():
                        param_group.post_backward()
            callback = getattr(state, "_root_post_backward_final_callback", None)
            if callable(callback):
                callback()
        self.scale_grads(grad_scale_factor)


class _PipelineStage(_PipelineStageBase):
    def __init__(self, stage_module: Any, stage_index: int, pipe_info: Any, device: Any = None, group: Any = None) -> None:
        super().__init__(stage_module, stage_index, pipe_info.num_stages, device, group)
        self.pipe_info = pipe_info
        graph_owner = getattr(pipe_info, "graph", None)
        self.graph = getattr(graph_owner, "graph", graph_owner)
        submod_nodes = [
            node
            for node in getattr(self.graph, "nodes", ())
            if getattr(node, "op", None) == "call_module"
        ]
        if len(submod_nodes) != self.num_stages:
            raise PipeliningMetadataError(
                f"Number of submodules in pipe graph {len(submod_nodes)} does not match "
                f"number of stages {self.num_stages}"
            )
        self.node = submod_nodes[stage_index]
        self.name = self.node.name
        self.submod_to_stage_index = {
            getattr(node, "name", ""): index
            for index, node in enumerate(submod_nodes)
        }
        self._move_submod_to_device()

    def _move_submod_to_device(self) -> None:
        parameters = getattr(self.submod, "parameters", None)
        if callable(parameters) and any(
            bool(getattr(parameter, "is_meta", False))
            for parameter in parameters()
        ):
            return
        if self.device is not None and hasattr(self.submod, "to"):
            self.submod.to(self.device)

    def get_stage_index_of_submod(self, submod_name: str) -> int:
        try:
            return self.submod_to_stage_index[submod_name]
        except KeyError as exc:
            raise PipeliningMetadataError(
                f"stage {submod_name!r} is not present"
            ) from exc

    def _tensor_from_meta(self, meta: Any, value: Any = None) -> Any:
        if isinstance(value, tp.Tensor):
            result = value.detach().clone()
        elif meta is not None and hasattr(meta, "to_tensor"):
            result = meta.to_tensor(self.device)
        elif meta is not None and hasattr(meta, "shape"):
            result = tp.empty(tuple(meta.shape), dtype=meta.dtype, device=self.device)
        else:
            result = value
        if isinstance(result, tp.Tensor) and self.has_backward:
            if result.is_floating_point() or result.is_complex():
                result.requires_grad_(True)
        return result

    def _create_act_recv_info(self) -> tuple[_RecvInfo, ...]:
        if self.node is None:
            raise PipeliningMetadataError(
                f"Stage {self.stage_index}: graph stage node is unavailable"
            )
        stage_graph = getattr(self.submod, "graph", None)
        placeholders = [
            node
            for node in getattr(stage_graph, "nodes", ())
            if getattr(node, "op", None) == "placeholder"
        ]
        outer_args = tuple(getattr(self.node, "args", ()))
        result: list[_RecvInfo] = []
        if len(placeholders) != len(outer_args):
            raise PipeliningMetadataError(
                f"Stage {self.stage_index}: graph placeholder and dependency counts differ"
            )
        for placeholder, arg_node in zip(placeholders, outer_args, strict=True):
            meta_value = getattr(placeholder, "meta", {}).get("val")
            if meta_value is None:
                raise PipeliningMetadataError(
                    f"Stage {self.stage_index}: placeholder metadata is unavailable"
                )
            if isinstance(meta_value, DTensor):
                raise PipeliningMetadataError(
                    f"Stage {self.stage_index}: distributed tensor metadata is unsupported for graph stages"
                )
            if getattr(arg_node, "op", None) == "placeholder":
                result.append(
                    _RecvInfo(
                        f"root_input_{getattr(placeholder, 'name', 'input')}",
                        None,
                        None,
                        _TensorMeta.from_tensor(meta_value),
                        True,
                    )
                )
                continue
            while getattr(arg_node, "target", None) is operator.getitem:
                arg_node = arg_node.args[0]
            if getattr(arg_node, "op", None) != "call_module":
                raise PipeliningMetadataError(
                    f"Stage {self.stage_index}: expected a stage dependency"
                )
            source = self.get_stage_index_of_submod(getattr(arg_node, "name", ""))
            meta = _TensorMeta(
                shape=tuple(meta_value.shape),
                stride=tuple(meta_value.stride()),
                dtype=meta_value.dtype,
                requires_grad=bool(
                    self.has_backward
                    and (
                        meta_value.is_floating_point()
                        or meta_value.is_complex()
                    )
                ),
            )
            result.append(
                _RecvInfo(
                    getattr(arg_node, "name", getattr(placeholder, "name", "input")),
                    source,
                    _make_tensor_from_meta(meta, self.device),
                    meta,
                )
            )
        return tuple(result)

    def _prepare_forward_infra(self, num_microbatches: int, args: Any, kwargs: Any, has_backward: bool) -> Any:
        del kwargs
        self.chunks = int(num_microbatches)
        self.has_backward = bool(has_backward)
        for index in range(self.chunks):
            self.args_recv_info[index] = self._create_act_recv_info()
        recv_infos = self.args_recv_info[0]
        if self.is_first:
            if not isinstance(args, tuple):
                raise AssertionError("first stage requires real tensor args")
            self._stage_meta.inputs = tuple(
                info.tensor_meta for info in recv_infos[: len(args)]
            )
        else:
            self._stage_meta.inputs = tuple(
                info.tensor_meta for info in recv_infos if not info.is_root_arg
            )
        self.act_send_info = self._create_act_send_info()

    def _prepare_backward_infra(
        self,
        num_microbatches: int,
        loss_fn: Any = None,
        target: Any = None,
        received_grad_meta: Any = None,
        loss_kwargs: Any = None,
    ) -> None:
        del loss_fn, target, received_grad_meta, loss_kwargs
        if self._stage_meta.inputs is None:
            raise PipeliningMetadataError(
                f"Stage {self.stage_index}: inputs metadata required for backward inference."
            )
        self._stage_meta.input_grads = _derive_grad_metas(self._stage_meta.inputs)
        self._setup_backward_recv_info(num_microbatches)
        return None

    def find_dst_rank(self, user: Any) -> int:
        if getattr(user, "op", None) != "call_module":
            return None
        return self.get_stage_index_of_submod(getattr(user, "name", ""))

    def _create_act_send_info(self) -> dict[int, list[int]]:
        if self.node is None:
            return {0: [self.stage_index + 1]} if not self.is_last else {0: []}
        result: dict[int, list[int]] = {}
        for user in getattr(self.node, "users", ()):
            if getattr(user, "target", None) is operator.getitem:
                output_index = int(user.args[1])
                destinations = result.setdefault(output_index, [])
                for child in getattr(user, "users", ()):
                    destination = self.find_dst_rank(child)
                    if destination is not None and destination not in destinations:
                        destinations.append(destination)
            else:
                destination = self.find_dst_rank(user)
                if destination is not None:
                    destinations = result.setdefault(0, [])
                    if destination not in destinations:
                        destinations.append(destination)
        output_node = self._get_output_node()
        if output_node is not None:
            values = output_node.args[0] if getattr(output_node, "args", ()) else ()

            def flatten_graph_values(value: Any) -> list[Any]:
                if isinstance(value, (tuple, list)):
                    result_values: list[Any] = []
                    for item in value:
                        result_values.extend(flatten_graph_values(item))
                    return result_values
                if isinstance(value, dict):
                    result_values = []
                    for item in value.values():
                        result_values.extend(flatten_graph_values(item))
                    return result_values
                return [value]

            output_metas: list[_TensorMeta] = []
            for index, value in enumerate(flatten_graph_values(values)):
                example_value = getattr(value, "meta", {}).get("val")
                if example_value is None:
                    raise PipeliningMetadataError(
                        f"Stage {self.stage_index}: output metadata is unavailable at index {index}"
                    )
                if isinstance(example_value, DTensor):
                    raise PipeliningMetadataError(
                        f"Stage {self.stage_index}: distributed tensor metadata is unsupported for graph stages"
                    )
                if not isinstance(example_value, tp.Tensor):
                    raise PipeliningMetadataError(
                        f"Stage {self.stage_index}: output {index} is not a tensor"
                    )
                output_metas.append(
                    _TensorMeta(
                        shape=tuple(example_value.shape),
                        stride=tuple(example_value.stride()),
                        dtype=example_value.dtype,
                        requires_grad=bool(
                            self.has_backward
                            and (
                                example_value.is_floating_point()
                                or example_value.is_complex()
                            )
                        ),
                    )
                )
            self._stage_meta.outputs = tuple(output_metas)
        return result

    def _create_grad_recv_info(self, act_send_info: Any) -> tuple[_RecvInfo, ...]:
        if self._stage_meta.outputs is None:
            raise PipeliningMetadataError(
                f"Stage {self.stage_index}: outputs metadata required for grad recv info."
            )

        outputs_meta = self._stage_meta.outputs
        output_grads_metas: list[Any] = []
        grad_recv_infos: list[_RecvInfo] = []
        for out_idx, out_meta in enumerate(outputs_meta):
            dst_list = act_send_info.get(out_idx, [])
            grad_src = dst_list[0] if dst_list else self.stage_index + 1
            if not dst_list or not out_meta.requires_grad:
                output_grads_metas.append(None)
                grad_recv_infos.append(
                    _RecvInfo(
                        f"recv_grad_for_{self.stage_index}_none_{out_idx}",
                        grad_src,
                        None,
                        None,
                    )
                )
                continue
            grad_meta = _TensorMeta(
                shape=out_meta.shape,
                stride=out_meta.stride,
                dtype=out_meta.dtype,
                requires_grad=False,
            )
            output_grads_metas.append(grad_meta)
            if len(dst_list) != 1:
                raise PipeliningMetadataError(
                    "Backward of skip connections not supported yet"
                )
            grad_recv_infos.append(
                _RecvInfo(
                    f"recv_grad_for_{self.stage_index}_from_{grad_src}",
                    grad_src,
                    _make_tensor_from_meta(grad_meta, self.device),
                    grad_meta,
                )
            )
        self._stage_meta.output_grads = tuple(output_grads_metas)
        if self._stage_meta.inputs is not None:
            self._stage_meta.input_grads = _derive_grad_metas(self._stage_meta.inputs)
        return tuple(grad_recv_infos)

    def backward_one_chunk(self, bwd_chunk_id: int, loss: Any = None, full_backward: bool = True, last_backward: bool = False) -> Any:
        self._check_chunk_id(bwd_chunk_id)
        return super().backward_one_chunk(
            bwd_chunk_id,
            loss=loss,
            full_backward=full_backward,
            last_backward=last_backward,
        )

    def _get_output_node(self) -> Any:
        for graph in (getattr(self.submod, "graph", None), self.graph):
            output_node = next(
                (
                    node
                    for node in getattr(graph, "nodes", ())
                    if getattr(node, "op", None) == "output"
                ),
                None,
            )
            if output_node is not None:
                return output_node
        return None



[docs]
def build_stage(stage_module: Any, stage_index: int, pipe_info: Any, device: Any = None, group: Any = None) -> _PipelineStage:
    return _PipelineStage(stage_module, stage_index, pipe_info, device, group)




[docs]
class PipelineStage(_PipelineStageBase):
    def __init__(self, submodule: Any, stage_index: int, num_stages: int, device: Any = None, input_args: tuple[Any, ...] | None = None, output_args: Any = None, output_grads: Any = None, input_grads: Any = None, group: Any = None, dw_builder: Callable[[], Callable[..., None]] | None = None, get_mesh: Any = None) -> None:
        super().__init__(submodule, stage_index, num_stages, device, group, dw_builder)
        self._mesh_cache = _MeshCache(get_mesh_cb=get_mesh)
        self._input_example = _normalize_model_output_as_tuple(input_args) if input_args is not None else ()
        self._output_example = _normalize_model_output_as_tuple(output_args) if output_args is not None else None
        input_grad_values = _normalize_model_output_as_tuple(input_grads) if input_grads is not None else None
        output_grad_values = _normalize_model_output_as_tuple(output_grads) if output_grads is not None else None
        self._user_meta = _StageMeta()
        self._user_meta.inputs = extract_tensor_metas(self._input_example) if self._input_example else None
        self._user_meta.outputs = extract_tensor_metas(self._output_example) if self._output_example is not None else None
        self._user_meta.input_grads = extract_tensor_metas(input_grad_values, allow_none=True) if input_grad_values is not None else None
        self._user_meta.output_grads = extract_tensor_metas(output_grad_values, allow_none=True) if output_grad_values is not None else None
        for values in (self._input_example, self._output_example, input_grad_values, output_grad_values):
            if values:
                self._mesh_cache.update_from_tensors(values)
        if self._user_meta.has_dtensors():
            if self._input_example and input_grad_values:
                validate_static_arg_grad_correspondence(
                    self.stage_index,
                    self._input_example,
                    input_grad_values,
                    is_input=True,
                )
            if self._output_example and output_grad_values:
                validate_static_arg_grad_correspondence(
                    self.stage_index,
                    self._output_example,
                    output_grad_values,
                    is_input=False,
                )
        self._inference_mode: InferenceMode | None = None
        self._fwd_outputs_for_bwd_meta: tuple[Any, ...] | None = None
        self._fwd_inputs_for_bwd_meta: tuple[Any, ...] | None = None
        self._fwd_kwargs_tensors_for_bwd_meta: tuple[Any, ...] | None = None
        self._metadata_inference_buffer_backup: list[tuple[Any, Any]] | None = None
        self._inference_mode = None

    def _prepare_forward_infra(self, num_microbatches: int, args: Any, kwargs: Any, has_backward: bool) -> Any:
        self.chunks = int(num_microbatches)
        self.has_backward = bool(has_backward)
        self._inference_mode = (
            InferenceMode.DYNAMIC
            if InferenceMode.needs_dynamic(self._user_meta, has_backward)
            else InferenceMode.STATIC
        )
        source_args = args
        if source_args is None or source_args == ():
            source_args = self._input_example
        fwd_meta_output = None
        if self._inference_mode == InferenceMode.DYNAMIC:
            fwd_meta_output = self._forward_metadata_inference(
                source_args, kwargs, has_backward
            )
        else:
            self._stage_meta.inputs = self._user_meta.inputs
            self._stage_meta.outputs = self._user_meta.outputs
        if self._stage_meta.inputs is None and source_args:
            self._stage_meta.inputs = extract_tensor_metas(tuple(source_args))
        if self._stage_meta.outputs is None and self._output_example is not None:
            self._stage_meta.outputs = extract_tensor_metas(self._output_example)
        self._setup_forward_recv_info(self.chunks, has_backward)
        self._setup_forward_send_info()
        return fwd_meta_output

    def _prepare_backward_infra(
        self,
        num_microbatches: int,
        loss_fn: Any = None,
        target: Any = None,
        received_grad_meta: Any = None,
        loss_kwargs: Any = None,
    ) -> Any:
        self.chunks = int(num_microbatches)
        self.has_backward = True
        if self._inference_mode == InferenceMode.DYNAMIC:
            result = self._backward_metadata_inference(
                loss_fn,
                target,
                received_grad_meta,
                loss_kwargs,
            )
            self._validate_inferred_metadata()
        else:
            result = None
            self._stage_meta.inputs = self._user_meta.inputs
            self._stage_meta.outputs = self._user_meta.outputs
            self._stage_meta.input_grads = self._user_meta.input_grads
            self._stage_meta.output_grads = self._user_meta.output_grads
        if isinstance(received_grad_meta, _StageBackwardMeta):
            self._stage_meta.output_grads = received_grad_meta.input_grad_metas
        if self._stage_meta.output_grads is None:
            if self._stage_meta.outputs is None:
                raise PipeliningMetadataError(
                    f"Stage {self.stage_index}: output metadata is required for backward inference."
                )
            self._stage_meta.output_grads = _derive_grad_metas(self._stage_meta.outputs)
        if self._stage_meta.input_grads is None:
            if self._stage_meta.inputs is None:
                raise PipeliningMetadataError(
                    f"Stage {self.stage_index}: input metadata is required for backward inference."
                )
            self._stage_meta.input_grads = _derive_grad_metas(self._stage_meta.inputs)
        self._setup_backward_recv_info(num_microbatches)
        self.grad_send_info = self._create_grad_send_info(self.args_recv_info.get(0, ()))
        return result

    def get_fwd_recv_ops(self, fwd_chunk_id: int) -> list[Any]:
        self._check_chunk_id(fwd_chunk_id)
        if self.is_first:
            return []
        return self._get_recv_ops(
            self.args_recv_info[fwd_chunk_id], self._downstream_group
        )

    def _recv_meta(self, src_stage: int) -> Any:
        objects = [None]
        dist.recv_object_list(
            objects,
            src=self._resolve_peer_global_rank(src_stage),
            group=self.group,
            device=self.device,
        )
        if len(objects) != 1:
            raise PipeliningMetadataError(
                f"expected one metadata object, got {len(objects)}"
            )
        return objects[0]

    def _send_meta(self, meta: Any, dst_stage: int) -> None:
        dist.send_object_list(
            [meta],
            dst=self._resolve_peer_global_rank(dst_stage),
            group=self.group,
            device=self.device,
        )

    def _is_same_rank(self, other_stage: int) -> bool:
        return self.stage_index_to_group_rank[int(other_stage)] == self.group_rank

    def _warmup_forward_vote(self, has_backward: bool, received_acc: Any = None) -> Any:
        my_vote = 0 if InferenceMode.needs_dynamic(self._user_meta, has_backward) else 1
        vote = tp.tensor([my_vote], dtype=tp.int32, device=self.device)
        if self.is_first:
            accumulated = vote
        elif self._is_same_rank(self.stage_index - 1):
            if received_acc is None:
                raise AssertionError("forward vote is missing the accumulated value")
            accumulated = received_acc * vote
        else:
            accumulated = tp.zeros(1, dtype=tp.int32, device=self.device)
            dist.recv(
                accumulated,
                src=self._resolve_peer_global_rank(self.stage_index - 1),
                group=self.group,
            )
            accumulated = accumulated * vote
        if not self.is_last and not self._is_same_rank(self.stage_index + 1):
            dist.send(
                accumulated,
                dst=self._resolve_peer_global_rank(self.stage_index + 1),
                group=self.group,
            )
        return accumulated

    def _warmup_backward_result(self, received_result: Any = None) -> Any:
        if self.is_last or self._is_same_rank(self.stage_index + 1):
            if received_result is None:
                raise AssertionError("backward vote is missing the accumulated value")
            result = received_result
        else:
            result = tp.zeros(1, dtype=tp.int32, device=self.device)
            dist.recv(
                result,
                src=self._resolve_peer_global_rank(self.stage_index + 1),
                group=self.group,
            )
        if not self.is_first and not self._is_same_rank(self.stage_index - 1):
            dist.send(
                result,
                dst=self._resolve_peer_global_rank(self.stage_index - 1),
                group=self.group,
            )
        return result

    def _compute_outputs(self, *args: Any, module: Any = None, **kwargs: Any) -> Any:
        return (self.submod if module is None else module)(*args, **kwargs)

    def _compute_input_grads(
        self,
        outputs: Any,
        all_fwd_inputs: Any,
        grad_outputs: Any = None,
    ) -> tuple[Any, ...]:
        return _autograd_grad_for_inputs(
            tuple(outputs),
            tuple(all_fwd_inputs),
            None if grad_outputs is None else tuple(grad_outputs),
            allow_unused=True,
        )

    def backward_one_chunk(self, bwd_chunk_id: int, loss: Any = None, full_backward: bool = True, last_backward: bool = False) -> Any:
        return super().backward_one_chunk(
            bwd_chunk_id,
            loss=loss,
            full_backward=full_backward,
            last_backward=last_backward,
        )

    def _to_tensor(self, arg: Any) -> Any:
        if isinstance(arg, DTensor):
            local = arg.to_local().detach()
            if getattr(arg, "requires_grad", False) and (
                local.is_floating_point() or local.is_complex()
            ):
                local.requires_grad_(True)
            return DTensor.from_local(
                local,
                device_mesh=arg.device_mesh,
                placements=arg.placements,
                shape=arg.shape,
                stride=arg.stride(),
            )
        if isinstance(arg, tp.Tensor):
            result = arg.detach()
            if arg.requires_grad:
                result.requires_grad_(True)
            return result
        if isinstance(arg, _DTensorMeta):
            mesh = self._mesh_cache.get_mesh(arg.mesh_cache_key)
            local = _make_tensor_from_meta(arg, self.device)
            if arg.requires_grad and (
                local.is_floating_point() or local.is_complex()
            ):
                local.requires_grad_(True)
            return DTensor.from_local(
                local,
                device_mesh=mesh,
                placements=arg.placements,
                shape=arg.global_shape,
                stride=arg.global_stride,
            )
        if isinstance(arg, _TensorMeta):
            result = arg.to_tensor(self.device)
            if arg.requires_grad and (
                result.is_floating_point() or result.is_complex()
            ):
                result.requires_grad_(True)
            return result
        raise PipeliningMetadataError(
            f"unsupported metadata value {type(arg).__name__}"
        )

    def _ones_from_metadata(self, meta: Any) -> Any:
        local = tp.ones(meta.shape, dtype=meta.dtype, device=self.device)
        if isinstance(meta, _DTensorMeta):
            mesh = self._mesh_cache.get_mesh(meta.mesh_cache_key)
            return DTensor.from_local(
                local,
                device_mesh=mesh,
                placements=meta.placements,
                shape=meta.global_shape,
                stride=meta.global_stride,
            )
        return local

    def _pre_metadata_inference_backup(self) -> None:
        if self._inference_mode != InferenceMode.DYNAMIC:
            return
        if self._metadata_inference_buffer_backup is not None:
            raise RuntimeError("metadata inference backup is already active")
        named_buffers = getattr(self.submod, "named_buffers", None)
        if callable(named_buffers):
            self._metadata_inference_buffer_backup = [
                (buffer, buffer.detach().clone())
                for _, buffer in named_buffers(remove_duplicate=False)
            ]

    def _forward_metadata_inference(self, args: Any, kwargs: Any, has_backward: bool) -> Any:
        kwargs = kwargs or {}
        if self.is_first:
            if args is None or isinstance(args, _StageForwardMeta):
                raise PipeliningMetadataError(
                    f"Stage {self.stage_index}: first stage requires tensor inputs"
                )
            values = tuple(args)
            self._stage_meta.inputs = extract_tensor_metas(values)
            inference_args = tuple(self._to_tensor(value) for value in values)
        elif self._is_same_rank(self.stage_index - 1) or isinstance(args, _StageForwardMeta):
            if not isinstance(args, _StageForwardMeta):
                raise PipeliningMetadataError(
                    f"Stage {self.stage_index}: forward metadata from the previous stage is required"
                )
            input_metas = args.forward_metas
            self._stage_meta.inputs = tuple(input_metas)
            inference_args = tuple(self._to_tensor(meta) for meta in input_metas)
        else:
            recv_meta = self._recv_meta(self.stage_index - 1)
            if not isinstance(recv_meta, _StageForwardMeta):
                raise PipeliningMetadataError(
                    f"Stage {self.stage_index}: invalid forward metadata received from the previous stage"
                )
            input_metas = recv_meta.forward_metas
            self._stage_meta.inputs = tuple(input_metas)
            inference_args = tuple(self._to_tensor(meta) for meta in input_metas)
        inference_kwargs = {
            key: self._to_tensor(value) if isinstance(value, tp.Tensor) else value
            for key, value in kwargs.items()
        }
        with (tp.enable_grad() if has_backward else tp.no_grad()):
            output = self._compute_outputs(
                *inference_args,
                module=self.submod,
                **inference_kwargs,
            )
        output_values = _normalize_model_output_as_tuple(output)
        self._stage_meta.outputs = tuple(
            meta
            for meta in (extract_tensor_meta(value) for value in output_values)
            if meta is not None
        )
        self._fwd_outputs_for_bwd_meta = output_values
        self._fwd_inputs_for_bwd_meta = inference_args
        self._fwd_kwargs_tensors_for_bwd_meta = tuple(
            value
            for value in flatten_args(inference_kwargs)
            if isinstance(value, tp.Tensor) or isinstance(value, DTensor)
        )
        fwd_meta = _StageForwardMeta(forward_metas=self._stage_meta.outputs)
        if self.is_last or self._is_same_rank(self.stage_index + 1):
            return fwd_meta
        self._send_meta(fwd_meta, self.stage_index + 1)
        return None

    def _backward_metadata_inference(self, loss_fn: Any, target: Any, received_grad_meta: Any, loss_kwargs: Any) -> Any:
        fwd_outputs = self._fwd_outputs_for_bwd_meta
        fwd_inputs = self._fwd_inputs_for_bwd_meta
        if fwd_outputs is None or fwd_inputs is None:
            raise PipeliningMetadataError(
                f"Stage {self.stage_index}: forward metadata inference must run first"
            )
        all_inputs = list(fwd_inputs) + list(self._fwd_kwargs_tensors_for_bwd_meta or ())
        if self.is_last:
            if loss_fn is None or target is None:
                raise PipeliningMetadataError(
                    f"Stage {self.stage_index}: loss_fn and target are required for backward inference"
                )
            output_value = fwd_outputs[0] if len(fwd_outputs) == 1 else fwd_outputs
            loss = loss_fn(output_value, self._to_tensor(target), **(loss_kwargs or {}))
            input_grads = self._compute_input_grads((loss,), all_inputs)
            self._stage_meta.output_grads = None
        else:
            if self._is_same_rank(self.stage_index + 1) or (
                not dist.is_initialized() and received_grad_meta is not None
            ):
                if not isinstance(received_grad_meta, _StageBackwardMeta):
                    raise PipeliningMetadataError(
                        f"Stage {self.stage_index}: backward metadata from the next stage is required"
                    )
                output_grad_metas = received_grad_meta.backward_metas
            else:
                recv_meta = self._recv_meta(self.stage_index + 1)
                if not isinstance(recv_meta, _StageBackwardMeta):
                    raise PipeliningMetadataError(
                        f"Stage {self.stage_index}: invalid backward metadata received from the next stage"
                    )
                output_grad_metas = recv_meta.backward_metas
            self._stage_meta.output_grads = output_grad_metas
            if len(fwd_outputs) != len(output_grad_metas):
                raise PipeliningMetadataError(
                    f"Stage {self.stage_index}: output and gradient metadata counts differ"
                )
            filtered_outputs = []
            filtered_grad_outputs = []
            for index, (output, grad_meta) in enumerate(
                zip(fwd_outputs, output_grad_metas, strict=True)
            ):
                if not isinstance(output, (tp.Tensor, DTensor)):
                    raise PipeliningMetadataError(
                        f"Stage {self.stage_index}: output {index} is not a tensor"
                    )
                if not output.requires_grad and getattr(output, "grad_fn", None) is None:
                    if grad_meta is not None:
                        raise PipeliningMetadataError(
                            f"Stage {self.stage_index}: output {index} has gradient metadata but does not require gradients"
                        )
                    continue
                filtered_outputs.append(output)
                filtered_grad_outputs.append(
                    self._ones_from_metadata(grad_meta) if grad_meta is not None else None
                )
            if filtered_outputs:
                input_grads = self._compute_input_grads(
                    filtered_outputs,
                    all_inputs,
                    filtered_grad_outputs,
                )
            else:
                input_grads = tuple(None for _ in all_inputs)
        input_metas = self._stage_meta.inputs or ()
        if len(input_grads) < len(input_metas):
            raise PipeliningMetadataError(
                f"Stage {self.stage_index}: backward returned too few input gradients"
            )
        self._stage_meta.input_grads = tuple(
            extract_tensor_meta(gradient)
            if isinstance(gradient, (tp.Tensor, DTensor))
            else (
                _derive_grad_metas((meta,))[0]
                if meta is not None and meta.requires_grad
                else None
            )
            for meta, gradient in zip(input_metas, input_grads)
        )
        bwd_meta = _StageBackwardMeta(backward_metas=self._stage_meta.input_grads)
        if self.is_first or self._is_same_rank(self.stage_index - 1):
            return bwd_meta
        self._send_meta(bwd_meta, self.stage_index - 1)
        return None

    def _post_metadata_inference_cleanup(self) -> None:
        if self._metadata_inference_buffer_backup is not None:
            with tp.no_grad():
                for buffer, saved in self._metadata_inference_buffer_backup:
                    buffer.copy_(saved)
            self._metadata_inference_buffer_backup = None
        self._fwd_outputs_for_bwd_meta = None
        self._fwd_inputs_for_bwd_meta = None
        self._fwd_kwargs_tensors_for_bwd_meta = None
        self.clear_runtime_states()

    def _validate_inferred_metadata(self) -> None:
        if not self._stage_meta.outputs:
            raise PipeliningMetadataError("stage output metadata is empty")
        for user_meta, inferred_meta, label in (
            (self._user_meta.inputs, self._stage_meta.inputs, "input"),
            (self._user_meta.outputs, self._stage_meta.outputs, "output"),
            (self._user_meta.input_grads, self._stage_meta.input_grads, "input_grad"),
            (self._user_meta.output_grads, self._stage_meta.output_grads, "output_grad"),
        ):
            if user_meta is not None and inferred_meta is not None:
                validate_tensors_metadata(
                    f"Stage {self.stage_index} {label}",
                    user_meta,
                    inferred_meta,
                    raise_on_mismatch=False,
                    warn_on_mismatch=True,
                )

    def _setup_forward_recv_info(self, num_microbatches: int, has_backward: bool) -> None:
        del has_backward
        if self._stage_meta.inputs is None:
            raise PipeliningMetadataError(
                f"Stage {self.stage_index}: inputs metadata is required for receive setup."
            )
        self.args_recv_info = {}
        for chunk_id in range(num_microbatches):
            if self.is_first:
                infos = tuple(
                    _RecvInfo(
                        f"root_input_{index}",
                        None,
                        None,
                        meta,
                        True,
                    )
                    for index, meta in enumerate(self._stage_meta.inputs)
                )
            else:
                infos = tuple(
                    _RecvInfo(
                        f"recv_for_{self.stage_index}_from_{self.stage_index - 1}",
                        self.stage_index - 1,
                        self._to_tensor(meta),
                        meta,
                        False,
                    )
                    for meta in self._stage_meta.inputs
                )
            self.args_recv_info[chunk_id] = infos

    def _setup_forward_send_info(self) -> None:
        if self._stage_meta.outputs is None:
            raise PipeliningMetadataError(
                f"Stage {self.stage_index}: outputs metadata is required for send setup."
            )
        self.act_send_info = {
            index: [self.stage_index + 1] if not self.is_last else []
            for index in range(len(self._stage_meta.outputs))
        }

    def _create_grad_recv_info(
        self,
        act_send_info: dict,
    ) -> tuple[_RecvInfo, ...]:
        grad_recv_infos: list[_RecvInfo] = []
        if not self.is_last:
            if self._stage_meta.output_grads is None:
                raise PipeliningMetadataError(
                    f"Stage {self.stage_index}: output_grads metadata is required for creating grad recv info."
                )
            output_grads = self._stage_meta.output_grads
            for index, destinations in act_send_info.items():
                if destinations is None or not destinations:
                    raise PipeliningMetadataError(
                        f"Stage {self.stage_index}: output {index} is not sent to any stage."
                    )
                source = destinations[0]
                grad_meta = output_grads[index]
                grad_recv_infos.append(
                    _RecvInfo(
                        f"recv_grad_for_{self.stage_index}_from_{source}",
                        source,
                        _make_tensor_from_meta(grad_meta, self.device)
                        if grad_meta is not None
                        else None,
                        grad_meta,
                    )
                )
        return tuple(grad_recv_infos)
Ask DeepWiki