# Source code for tensorplay.distributed.pipelining.schedules Source: https://www.tensorplay.cn/docs/_modules/tensorplay/distributed/pipelining/schedules.html ``` """Microbatch pipeline schedules.""" import csv import logging import re from collections import Counter, defaultdict from dataclasses import dataclass from enum import Enum from typing import Any, Callable, Literal, Protocol, cast import tensorplay as tp from .. import distributed_core as dist from ..fsdp._fully_shard import FSDPModule, UnshardHandle from ._utils import InferenceMode, generate_rank_to_stage_mapping, generate_stage_to_rank_mapping from .microbatch import TensorChunkSpec, merge_chunks, split_args_kwargs_into_chunks, _split_tensor logger = logging.getLogger(__name__) __all__ = [ "get_schedule_class", "PipelineScheduleSingle", "PipelineScheduleMulti", "Schedule1F1B", "ScheduleGPipe", "ScheduleInterleaved1F1B", "ScheduleLoopedBFS", "ScheduleInterleavedZeroBubble", "ScheduleZBVZeroBubble", "ScheduleDualPipeV", ] class _ComputationType(str, Enum): FORWARD = "F" BACKWARD_INPUT = "I" BACKWARD_WEIGHT = "W" UNSHARD = "UNSHARD" RESHARD = "RESHARD" SEND_F = "SEND_F" RECV_F = "RECV_F" SEND_B = "SEND_B" RECV_B = "RECV_B" FULL_BACKWARD = "B" OVERLAP_F_B = "OVERLAP_F_B" REDUCE_GRAD = "REDUCE_GRAD" @staticmethod def from_str(action: str) -> "_ComputationType": return _ComputationType(action) class _Action(tuple): __slots__ = () def __new__(cls, stage_index: int, computation_type: _ComputationType, microbatch_index: int | None = None, sub_actions: tuple["_Action", ...] | None = None): return tuple.__new__(cls, (stage_index, computation_type, microbatch_index, sub_actions)) @property def stage_index(self) -> int: return self[0] @property def computation_type(self) -> _ComputationType: return self[1] @property def microbatch_index(self) -> int | None: return self[2] @property def sub_actions(self) -> tuple["_Action", ...] | None: return self[3] @property def is_compute_op(self) -> bool: return self.computation_type in {_ComputationType.FORWARD, _ComputationType.BACKWARD_INPUT, _ComputationType.BACKWARD_WEIGHT, _ComputationType.FULL_BACKWARD, _ComputationType.OVERLAP_F_B} def __repr__(self) -> str: if self.sub_actions is not None: return f"({';'.join(map(repr, self.sub_actions))}){self.computation_type.value}" return f"{self.stage_index}{self.computation_type.value}{'' if self.microbatch_index is None else self.microbatch_index}" def __str__(self) -> str: return self.__repr__() @staticmethod def from_str(action_string: str) -> "_Action | None": action_string = action_string.strip() if not action_string: return None if action_string.startswith("(") and ")" in action_string: end = action_string.index(")") sub = tuple(item for item in (_Action.from_str(part) for part in action_string[1:end].split(";")) if item is not None) return _Action(-1, _ComputationType.from_str(action_string[end + 1:]), None, sub) match = re.fullmatch(r"(\d+)(F|I|B|W|UNSHARD|RESHARD|REDUCE_GRAD|SEND_F|RECV_F|SEND_B|RECV_B)(\d*)", action_string) if match is None: raise ValueError(f"invalid pipeline action: {action_string}") stage, kind, microbatch = match.groups() return _Action(int(stage), _ComputationType(kind), int(microbatch) if microbatch else None) def _get_profiler_function_name(action: _Action) -> str: return f"TP:{action}" def _format_pipeline_order(pipeline_order: dict[int, list[_Action | None]], error_step_number: int | None = None) -> str: steps = max((len(actions) for actions in pipeline_order.values()), default=0) rows = ["step " + " ".join(f"rank {rank}" for rank in sorted(pipeline_order))] for index in range(steps): values = [str(pipeline_order.get(rank, [None] * steps)[index] or "") for rank in sorted(pipeline_order)] suffix = " " if index == error_step_number else "" rows.append(f"{index}: " + " ".join(values) + suffix) return "\n".join(rows) class _PipelineSchedule: def __init__(self, n_microbatches: int, loss_fn: Any = None, args_chunk_spec: Any = None, kwargs_chunk_spec: Any = None, output_merge_spec: Any = None, scale_grads: bool = True) -> None: if n_microbatches <= 0: raise ValueError("n_microbatches must be positive") self._n_microbatches = int(n_microbatches) self._loss_fn = loss_fn self._args_chunk_spec = args_chunk_spec self._kwargs_chunk_spec = kwargs_chunk_spec self._output_merge_spec = output_merge_spec self._scale_grads = scale_grads self._has_backward = loss_fn is not None self._stages: list[Any] = [] def _maybe_compute_loss(self, stage: Any, output: Any, target_mbs: Any, mb_index: int, loss_kwargs: dict[str, Any] | None) -> Any: if not getattr(stage, "is_last", False) or self._loss_fn is None or target_mbs is None: return None loss = self._loss_fn(output, target_mbs[mb_index], **(loss_kwargs or {})) self._losses.append(loss) return loss def _maybe_get_loss(self, stage: Any, mb_index: int) -> Any: if not getattr(stage, "is_last", False): return None return self._losses[mb_index] if 0 <= mb_index < len(self._losses) else None def _update_losses(self, stages: Any, losses: list[Any] | None) -> None: stage_list = stages if isinstance(stages, (list, tuple)) else [stages] if losses is not None and any(getattr(stage, "is_last", False) for stage in stage_list): if len(self._losses) != self._n_microbatches: raise RuntimeError( f"expected {self._n_microbatches} losses, got {len(self._losses)}" ) losses.clear() losses.extend(self._losses) self._losses.clear() def _warmup_p2p(self, stages: Any, has_backward: bool, p2p_done: Any) -> None: del p2p_done pipeline_stages = [ stage for stage in stages if hasattr(stage, "_user_meta") ] if len(pipeline_stages) != len(stages): if not dist.is_initialized(): return operations = [ operation for stage in stages for operation in stage._get_init_p2p_neighbors_ops() ] _wait_batch_p2p(_batch_p2p(operations)) return if not dist.is_initialized(): for stage in pipeline_stages: stage._inference_mode = ( InferenceMode.DYNAMIC if InferenceMode.needs_dynamic(stage._user_meta, has_backward) else InferenceMode.STATIC ) return has_cross_rank = any( (not stage.is_first and not stage._is_same_rank(stage.stage_index - 1)) or (not stage.is_last and not stage._is_same_rank(stage.stage_index + 1)) for stage in pipeline_stages ) if has_cross_rank and any( dist.get_backend(stage.group) == "fake" for stage in pipeline_stages ): for stage in pipeline_stages: if InferenceMode.needs_dynamic(stage._user_meta, has_backward): raise RuntimeError( f"Stage {stage.stage_index} requires dynamic metadata with a fake process group" ) stage._inference_mode = InferenceMode.STATIC return accumulated = None for stage in pipeline_stages: accumulated = stage._warmup_forward_vote( has_backward, received_acc=accumulated, ) result = accumulated for stage in reversed(pipeline_stages): result = stage._warmup_backward_result(received_result=result) stage._inference_mode = ( InferenceMode.STATIC if int(result.item()) == 1 else InferenceMode.DYNAMIC ) def _initialize_pp_stages( self, stages: list[Any], args: Any, kwargs: Any, target: Any, fwd_initialized: Any, bwd_initialized: Any, loss_kwargs: Any, ) -> tuple[bool, bool]: if fwd_initialized and self._has_backward != bwd_initialized: fwd_initialized = False bwd_initialized = False if not fwd_initialized: self._warmup_p2p(stages, self._has_backward, fwd_initialized) for stage in stages: stage.has_backward = self._has_backward backup = getattr(stage, "_pre_metadata_inference_backup", None) if callable(backup): backup() try: next_stage_args = None for stage in stages: stage_args = args if stage.is_first else next_stage_args next_stage_args = stage._prepare_forward_infra( self._n_microbatches, stage_args, kwargs, self._has_backward, ) fwd_initialized = True if self._has_backward and not bwd_initialized: previous_grad_meta = None for stage in reversed(stages): previous_grad_meta = stage._prepare_backward_infra( self._n_microbatches, loss_fn=self._loss_fn, target=target, received_grad_meta=previous_grad_meta, loss_kwargs=loss_kwargs, ) bwd_initialized = True finally: for stage in stages: cleanup = getattr(stage, "_post_metadata_inference_cleanup", None) if callable(cleanup): cleanup() elif self._has_backward and not bwd_initialized: previous_grad_meta = None for stage in reversed(stages): previous_grad_meta = stage._prepare_backward_infra( self._n_microbatches, loss_fn=self._loss_fn, target=target, received_grad_meta=previous_grad_meta, loss_kwargs=loss_kwargs, ) bwd_initialized = True return fwd_initialized, bwd_initialized def _step_microbatches(self, arg_mbs: list[tuple[Any, ...]], kwarg_mbs: list[dict[str, Any]], target_mbs: list[Any] | None, losses: list[Any] | None, return_outputs: bool = True, loss_kwargs: dict[str, Any] | None = None) -> Any: raise NotImplementedError def step(self, *args: Any, target: Any = None, losses: list[Any] | None = None, return_outputs: bool = True, loss_kwargs: dict[str, Any] | None = None, arg_mbs: Any = None, kwarg_mbs: Any = None, target_mbs: Any = None, **kwargs: Any) -> Any: arg_mbs, kwarg_mbs, target_mbs = self._get_microbatch_inputs(args, kwargs, target, arg_mbs, kwarg_mbs, target_mbs) self._initialize_for_step(args, kwargs, arg_mbs, kwarg_mbs, target, loss_kwargs) self._losses = [] return self._step_microbatches(arg_mbs or [], kwarg_mbs or [], target_mbs, losses, return_outputs, loss_kwargs) def _initialize_for_step( self, args: tuple[Any, ...], kwargs: dict[str, Any], arg_mbs: list[Any], kwarg_mbs: list[Any], target: Any, loss_kwargs: Any, ) -> None: init_args = tuple(arg_mbs[0]) if arg_mbs else args init_kwargs = dict(kwarg_mbs[0]) if kwarg_mbs else kwargs if isinstance(self, PipelineScheduleSingle): self._initialize_stage(init_args, init_kwargs, target, loss_kwargs) elif self._stages: self._initialize_stages(init_args, init_kwargs, target, loss_kwargs) def eval(self, *args: Any, target: Any = None, losses: list[Any] | None = None, arg_mbs: Any = None, kwarg_mbs: Any = None, target_mbs: Any = None, **kwargs: Any) -> Any: old = self._has_backward self._has_backward = False try: return self.step(*args, target=target, losses=losses, arg_mbs=arg_mbs, kwarg_mbs=kwarg_mbs, target_mbs=target_mbs, **kwargs) finally: self._has_backward = old def _check_inputs( self, arg_mbs: Any = None, kwarg_mbs: Any = None, target_mbs: Any = None, losses: Any = None, ) -> tuple[list[Any], list[Any]]: def check_type_and_len(value: Any, name: str) -> None: if not isinstance(value, list): raise TypeError(f"{name} must be a list but got a {type(value)}") if len(value) != self._n_microbatches: raise ValueError( f"Expecting {self._n_microbatches} {name} but got {len(value)}" ) if arg_mbs is not None: check_type_and_len(arg_mbs, "arg_mbs") else: arg_mbs = [()] * self._n_microbatches if kwarg_mbs is not None: check_type_and_len(kwarg_mbs, "kwarg_mbs") else: kwarg_mbs = [{}] * self._n_microbatches if target_mbs is not None: check_type_and_len(target_mbs, "target_mbs") if losses is not None and not isinstance(losses, list): raise TypeError(f"losses must be a list but got a {type(losses)}") return arg_mbs, kwarg_mbs def _compute_loss(self, output: Any, target: Any, loss_kwargs: dict[str, Any] | None = None) -> Any: return self._loss_fn(output, target, **(loss_kwargs or {})) def _split_inputs( self, args: tuple[Any, ...], kwargs: dict[str, Any] ) -> tuple[list[tuple[Any, ...]], list[dict[str, Any]]]: if args or kwargs: return split_args_kwargs_into_chunks( args, kwargs, self._n_microbatches, self._args_chunk_spec, self._kwargs_chunk_spec, ) return [()] * self._n_microbatches, [{}] * self._n_microbatches def _get_microbatch_inputs( self, args: tuple[Any, ...], kwargs: dict[str, Any], target: Any, arg_mbs: Any, kwarg_mbs: Any, target_mbs: Any, ) -> tuple[list[Any], list[Any], list[Any] | None]: pre_split = any( value is not None for value in (arg_mbs, kwarg_mbs, target_mbs) ) if not pre_split: args_split, kwargs_split = self._split_inputs(args, kwargs) target_split = ( list(_split_tensor(target, TensorChunkSpec(0), self._n_microbatches)) if target is not None else None ) return args_split, kwargs_split, target_split if args: raise ValueError( "When using pre-split inputs, pass pre-split positional inputs " "through arg_mbs=... instead of positional args." ) if kwargs: names = ", ".join(sorted(kwargs)) raise ValueError( f"Unexpected keyword arguments with pre-split inputs: {names}. " "Pass pre-split keyword inputs through kwarg_mbs=." ) if target is not None: raise ValueError( "When using pre-split inputs, pass pre-split targets through " "target_mbs=... instead of target=." ) arg_mbs, kwarg_mbs = self._check_inputs(arg_mbs, kwarg_mbs, target_mbs) for mb_index, (arg_mb, kwarg_mb) in enumerate( zip(arg_mbs, kwarg_mbs, strict=True) ): if not isinstance(arg_mb, tuple): raise TypeError( "arg_mbs must be a list of tuples, but " f"arg_mbs[{mb_index}] is a {type(arg_mb)}" ) if not isinstance(kwarg_mb, dict): raise TypeError( "kwarg_mbs must be a list of dicts, but " f"kwarg_mbs[{mb_index}] is a {type(kwarg_mb)}" ) return arg_mbs, kwarg_mbs, target_mbs def _merge_outputs(self, output_chunks: list[Any]) -> Any: if self._output_merge_spec is None: output_spec = _default_merge_spec(output_chunks[0]) else: output_spec = self._output_merge_spec return merge_chunks(output_chunks, output_spec) [docs] class PipelineScheduleSingle(_PipelineSchedule): def __init__(self, stage: Any, n_microbatches: int, loss_fn: Any = None, args_chunk_spec: Any = None, kwargs_chunk_spec: Any = None, output_merge_spec: Any = None, scale_grads: bool = True) -> None: super().__init__(n_microbatches, loss_fn, args_chunk_spec, kwargs_chunk_spec, output_merge_spec, scale_grads) self._stage = stage self._stages = [stage] self._num_stages = int(stage.num_stages) self._stage_forward_initialized = False self._stage_backward_initialized = False self.pipeline_order = self._get_pipeline_order() self._stage.has_backward = self._has_backward def _initialize_stage(self, args: Any, kwargs: Any, target: Any = None, loss_kwargs: Any = None) -> None: ( self._stage_forward_initialized, self._stage_backward_initialized, ) = self._initialize_pp_stages( [self._stage], args, kwargs, target, self._stage_forward_initialized, self._stage_backward_initialized, loss_kwargs, ) def step( self, *args: Any, target: Any = None, losses: list[Any] | None = None, return_outputs: bool = True, loss_kwargs: dict[str, Any] | None = None, arg_mbs: Any = None, kwarg_mbs: Any = None, target_mbs: Any = None, **kwargs: Any, ) -> Any: if self._has_backward and not tp.is_grad_enabled(): raise RuntimeError( "step() requires gradients to be enabled for backward computation" ) self._stage.has_backward = self._has_backward self._stage.clear_runtime_states() args_split, kwargs_split, targets_split = self._get_microbatch_inputs( args, kwargs, target, arg_mbs, kwarg_mbs, target_mbs, ) self._losses = [] self._initialize_stage( tuple(args_split[0]) if args_split else args, dict(kwargs_split[0]) if kwargs_split else kwargs, targets_split[0] if targets_split else None, loss_kwargs, ) self._step_microbatches( args_split, kwargs_split, targets_split, losses, return_outputs, loss_kwargs=loss_kwargs, ) if self._stage.is_last and return_outputs and self._stage.output_chunks: return self._merge_outputs(self._stage.output_chunks) return None def _step_microbatches(self, arg_mbs: list[tuple[Any, ...]], kwarg_mbs: list[dict[str, Any]], target_mbs: list[Any] | None, losses: list[Any] | None, return_outputs: bool = True, loss_kwargs: dict[str, Any] | None = None) -> Any: self._stage.clear_runtime_states() outputs = [] forward_sends = [] for index, (args, kwargs) in enumerate(zip(arg_mbs, kwarg_mbs)): for work in _run_p2p(self._stage.get_fwd_recv_ops(index)): work.wait() output = self._stage.forward_one_chunk( index, args, kwargs, save_forward_output=return_outputs, ) outputs.append(output) forward_sends.extend(_run_p2p(self._stage.get_fwd_send_ops(index))) for work in forward_sends: work.wait() if self._has_backward: backward_sends = [] for index in range(len(outputs)): for work in _run_p2p(self._stage.get_bwd_recv_ops(index)): work.wait() loss = self._maybe_compute_loss(self._stage, outputs[index], target_mbs, index, loss_kwargs) if self._stage.is_last: if loss is not None: self._stage.backward_one_chunk(index, loss=loss) else: self._stage.backward_one_chunk( index, last_backward=index == len(outputs) - 1, ) backward_sends.extend(_run_p2p(self._stage.get_bwd_send_ops(index))) for work in backward_sends: work.wait() if self._scale_grads: self._stage.scale_grads(self._n_microbatches) if losses is not None and self._stage.is_last: losses.extend(self._losses) if not return_outputs or not self._stage.is_last: return None return self._merge_outputs(outputs) def _get_pipeline_order(self) -> dict[int, list[_Action]]: actions: list[_Action | None] = [ _Action(self._stage.stage_index, _ComputationType.FORWARD, index) for index in range(self._n_microbatches) ] if self._has_backward: actions.extend( _Action(self._stage.stage_index, _ComputationType.FULL_BACKWARD, index) for index in range(self._n_microbatches) ) actions = _add_reduce_grad(actions, self._n_microbatches) return {int(self._stage.group_rank): actions} def _batch_p2p(operations: list[Any], desc: str | None = None) -> list[Any]: del desc if not operations: return [] operations_by_group: dict[str, list[Any]] = defaultdict(list) for operation in operations: group = operation.group group_name = getattr(group, "group_name", None) operations_by_group[str(group_name if group_name is not None else group)].append( operation ) if len(operations_by_group) > 1: works: list[Any] = [] for _, group_operations in sorted(operations_by_group.items()): works.extend(_batch_p2p(group_operations)) return works operation_types = {operation.op for operation in operations} if operation_types == {dist.isend}: return [ work for work in ( operation.op( operation.tensor, group=operation.group, tag=operation.tag, group_dst=operation.group_peer, ) for operation in operations ) if work is not None ] if operation_types == {dist.irecv}: return [ work for work in ( operation.op( operation.tensor, group=operation.group, tag=operation.tag, group_src=operation.group_peer, ) for operation in operations ) if work is not None ] return dist.batch_isend_irecv(operations) def _sorted_batch_p2p( operations: list[Any], desc: str | None = None ) -> dict[int, list[Any]]: del desc operations_by_peer: dict[int, list[Any]] = defaultdict(list) works_by_peer: dict[int, list[Any]] = {} for operation in operations: operations_by_peer[int(operation.peer)].append(operation) for peer, peer_operations in sorted(operations_by_peer.items()): works_by_peer[peer] = _batch_p2p(peer_operations) return works_by_peer def _wait_batch_p2p(works: list[Any]) -> None: for work in works: work.wait() def _run_p2p(operations: list[Any]) -> list[Any]: return [ work for peer_works in _sorted_batch_p2p(operations).values() for work in peer_works ] def _normalize_stage_args(value: Any) -> tuple[Any, ...]: if isinstance(value, tuple): return value if isinstance(value, list): return tuple(value) return (value,) class _ScheduleForwardOnly(PipelineScheduleSingle): def _step_microbatches(self, *args: Any, **kwargs: Any) -> Any: arg_mbs = kwargs.pop("arg_mbs", args[0] if args else None) kwarg_mbs = kwargs.pop("kwarg_mbs", args[1] if len(args) > 1 else None) target_mbs = kwargs.pop("target_mbs", args[2] if len(args) > 2 else None) losses = kwargs.pop("losses", args[3] if len(args) > 3 else None) return_outputs = kwargs.pop( "return_outputs", args[4] if len(args) > 4 else True ) if target_mbs is not None or losses is not None: raise RuntimeError("forward-only schedule does not support loss computation") arg_mbs, kwarg_mbs = self._check_inputs( arg_mbs, kwarg_mbs, target_mbs, losses ) self._initialize_stage(arg_mbs[0], kwarg_mbs[0]) self._stage.clear_runtime_states() send_works: list[Any] = [] for index in range(self._n_microbatches): _wait_batch_p2p( _batch_p2p( self._stage.get_fwd_recv_ops(index), desc="fwd_recv" ) ) self._stage.forward_one_chunk( index, arg_mbs[index], kwarg_mbs[index], save_forward_output=return_outputs, ) send_works.extend( _batch_p2p( self._stage.get_fwd_send_ops(index), desc="fwd_send" ) ) _wait_batch_p2p(send_works) if not return_outputs or not self._stage.is_last: return None return self._merge_outputs(self._stage.output_chunks) [docs] class ScheduleGPipe(PipelineScheduleSingle): """Execute all forward microbatches before draining their backwards.""" def _step_microbatches( self, arg_mbs: list[tuple[Any, ...]], kwarg_mbs: list[dict[str, Any]], target_mbs: list[Any] | None, losses: list[Any] | None, return_outputs: bool = True, loss_kwargs: dict[str, Any] | None = None, ) -> Any: arg_mbs, kwarg_mbs = self._check_inputs( arg_mbs, kwarg_mbs, target_mbs, losses ) outputs: list[Any] = [] forward_sends: list[Any] = [] for index in range(self._n_microbatches): _wait_batch_p2p( [ work for works in _sorted_batch_p2p( self._stage.get_fwd_recv_ops(index), desc="fwd_recv" ).values() for work in works ] ) output = self._stage.forward_one_chunk( index, arg_mbs[index], kwarg_mbs[index], save_forward_output=return_outputs, ) outputs.append(output) forward_sends.extend( work for works in _sorted_batch_p2p( self._stage.get_fwd_send_ops(index), desc="fwd_send" ).values() for work in works ) self._maybe_compute_loss( self._stage, output, target_mbs, index, loss_kwargs ) _wait_batch_p2p(forward_sends) backward_sends: list[Any] = [] if self._has_backward: for index in range(self._n_microbatches): _wait_batch_p2p( [ work for works in _sorted_batch_p2p( self._stage.get_bwd_recv_ops(index), desc="bwd_recv" ).values() for work in works ] ) self._stage.backward_one_chunk( index, loss=self._maybe_get_loss(self._stage, index), last_backward=index == self._n_microbatches - 1, ) backward_sends.extend( work for works in _sorted_batch_p2p( self._stage.get_bwd_send_ops(index), desc="bwd_send" ).values() for work in works ) _wait_batch_p2p(backward_sends) self._stage.perform_reduce_grad( self._n_microbatches if self._scale_grads else 1 ) self._update_losses(self._stage, losses) if not return_outputs or not self._stage.is_last: return None return self._merge_outputs(outputs) def _get_pipeline_order(self) -> dict[int, list[_Action | None]]: group_size = int(self._stage.group_size) if group_size != self._num_stages: raise ValueError("GPipe requires one stage per pipeline rank") pipeline_order: dict[int, list[_Action | None]] = {} for rank in range(group_size): actions: list[_Action | None] = [None] * rank actions.extend( _Action(rank, _ComputationType.FORWARD, microbatch) for microbatch in range(self._n_microbatches) ) if self._has_backward: actions.extend( [None] * (3 * (group_size - 1 - rank)) ) actions.extend( _Action(rank, _ComputationType.FULL_BACKWARD, microbatch) for microbatch in range(self._n_microbatches) ) pipeline_order[rank] = _add_reduce_grad( actions, self._n_microbatches ) else: pipeline_order[rank] = actions return pipeline_order [docs] class Schedule1F1B(PipelineScheduleSingle): """Overlap forward and backward microbatches after warmup.""" def __init__(self, stage: Any, n_microbatches: int, loss_fn: Any = None, args_chunk_spec: Any = None, kwargs_chunk_spec: Any = None, output_merge_spec: Any = None, scale_grads: bool = True) -> None: super().__init__( stage, n_microbatches, loss_fn, args_chunk_spec, kwargs_chunk_spec, output_merge_spec, scale_grads, ) if self._has_backward and n_microbatches < self._num_stages: raise ValueError("1F1B requires at least one microbatch per stage") def _step_microbatches( self, arg_mbs: list[tuple[Any, ...]], kwarg_mbs: list[dict[str, Any]], target_mbs: list[Any] | None, losses: list[Any] | None, return_outputs: bool = True, loss_kwargs: dict[str, Any] | None = None, ) -> Any: arg_mbs, kwarg_mbs = self._check_inputs( arg_mbs, kwarg_mbs, target_mbs, losses ) warmup = min(self._n_microbatches, self._num_stages - self._stage.stage_index) outputs: list[Any] = [] fwd_mb_index = 0 bwd_mb_index = 0 send_work: list[Any] = [] fwd_sends: list[Any] = [] for _ in range(warmup): _wait_batch_p2p( _batch_p2p( self._stage.get_fwd_recv_ops(fwd_mb_index), desc="fwd_recv" ) ) output = self._stage.forward_one_chunk( fwd_mb_index, arg_mbs[fwd_mb_index], kwarg_mbs[fwd_mb_index], save_forward_output=return_outputs, ) outputs.append(output) _wait_batch_p2p(send_work) fwd_sends = self._stage.get_fwd_send_ops(fwd_mb_index) if not self._has_backward or fwd_mb_index != warmup - 1: send_work = _batch_p2p(fwd_sends, desc="fwd_send") self._maybe_compute_loss( self._stage, output, target_mbs, fwd_mb_index, loss_kwargs ) fwd_mb_index += 1 if not self._has_backward: for fwd_mb_index in range(fwd_mb_index, self._n_microbatches): _wait_batch_p2p( _batch_p2p( self._stage.get_fwd_recv_ops(fwd_mb_index), desc="fwd_recv", ) ) output = self._stage.forward_one_chunk( fwd_mb_index, arg_mbs[fwd_mb_index], kwarg_mbs[fwd_mb_index], save_forward_output=return_outputs, ) outputs.append(output) _wait_batch_p2p(send_work) send_work = _batch_p2p( self._stage.get_fwd_send_ops(fwd_mb_index), desc="fwd_send" ) self._maybe_compute_loss( self._stage, output, target_mbs, fwd_mb_index, loss_kwargs ) _wait_batch_p2p(send_work) else: while True: _wait_batch_p2p( _batch_p2p( fwd_sends + self._stage.get_bwd_recv_ops(bwd_mb_index), desc="fwd_send_bwd_recv", ) ) self._stage.backward_one_chunk( bwd_mb_index, loss=self._maybe_get_loss(self._stage, bwd_mb_index), last_backward=bwd_mb_index == self._n_microbatches - 1, ) bwd_sends = self._stage.get_bwd_send_ops(bwd_mb_index) bwd_mb_index += 1 if fwd_mb_index == self._n_microbatches: break _wait_batch_p2p( _batch_p2p( bwd_sends + self._stage.get_fwd_recv_ops(fwd_mb_index), desc="bwd_send_fwd_recv", ) ) output = self._stage.forward_one_chunk( fwd_mb_index, arg_mbs[fwd_mb_index], kwarg_mbs[fwd_mb_index], save_forward_output=return_outputs, ) outputs.append(output) self._maybe_compute_loss( self._stage, output, target_mbs, fwd_mb_index, loss_kwargs ) fwd_sends = self._stage.get_fwd_send_ops(fwd_mb_index) fwd_mb_index += 1 send_work = _batch_p2p(bwd_sends, desc="bwd_send") while bwd_mb_index < self._n_microbatches: _wait_batch_p2p( _batch_p2p( self._stage.get_bwd_recv_ops(bwd_mb_index), desc="bwd_recv", ) ) self._stage.backward_one_chunk( bwd_mb_index, loss=self._maybe_get_loss(self._stage, bwd_mb_index), last_backward=bwd_mb_index == self._n_microbatches - 1, ) _wait_batch_p2p(send_work) send_work = _batch_p2p( self._stage.get_bwd_send_ops(bwd_mb_index), desc="bwd_send" ) bwd_mb_index += 1 _wait_batch_p2p(send_work) self._stage.perform_reduce_grad( self._n_microbatches if self._scale_grads else 1 ) self._update_losses(self._stage, losses) if not return_outputs or not self._stage.is_last: return None return self._merge_outputs(outputs) def _get_pipeline_order(self) -> dict[int, list[_Action | None]]: group_size = int(self._stage.group_size) if group_size != self._num_stages: raise ValueError("1F1B requires one stage per pipeline rank") pipeline_order: dict[int, list[_Action | None]] = {} for rank in range(group_size): actions: list[_Action | None] = [None] * rank warmup = min(self._n_microbatches, group_size - 1 - rank) actions.extend( _Action(rank, _ComputationType.FORWARD, microbatch) for microbatch in range(warmup) ) actions.extend([None] * (2 * (group_size - 1 - rank))) next_forward = warmup next_backward = 0 while next_forward < self._n_microbatches: actions.append( _Action(rank, _ComputationType.FORWARD, next_forward) ) next_forward += 1 actions.append( _Action(rank, _ComputationType.FULL_BACKWARD, next_backward) ) next_backward += 1 while next_backward < self._n_microbatches: if rank != group_size - 1: actions.append(None) actions.append( _Action(rank, _ComputationType.FULL_BACKWARD, next_backward) ) next_backward += 1 pipeline_order[rank] = _add_reduce_grad( actions, self._n_microbatches ) return pipeline_order [docs] class PipelineScheduleMulti(_PipelineSchedule): def __init__(self, stages: list[Any], n_microbatches: int, loss_fn: Any = None, args_chunk_spec: Any = None, kwargs_chunk_spec: Any = None, output_merge_spec: Any = None, use_full_backward: bool | None = None, scale_grads: bool = True, backward_requires_autograd: bool = True) -> None: if not stages: raise ValueError("at least one pipeline stage is required") super().__init__(n_microbatches, loss_fn, args_chunk_spec, kwargs_chunk_spec, output_merge_spec, scale_grads) self._stages = list(stages) self.use_full_backward = use_full_backward self.backward_requires_autograd = backward_requires_autograd self._backward_requires_autograd = backward_requires_autograd self._num_stages = int(stages[0].num_stages) self.pp_group_size = int(stages[0].group_size) if self._num_stages <= 0 or self.pp_group_size <= 0: raise ValueError("pipeline dimensions must be positive") if any(int(stage.num_stages) != self._num_stages for stage in stages): raise ValueError("all pipeline stages must use the same stage count") if len({int(stage.stage_index) for stage in stages}) != len(stages): raise ValueError("a pipeline stage cannot be listed more than once") self.rank = int(stages[0].group_rank) self.stage_index_to_group_rank = { index: index % self.pp_group_size for index in range(self._num_stages) } for stage in self._stages: stage.stage_index_to_group_rank = dict(self.stage_index_to_group_rank) self._stages_forward_initialized = False self._stages_backward_initialized = False self.pipeline_order: dict[int, list[_Action | None]] = {} if use_full_backward is not None: logger.warning( "use_full_backward is no longer supported; omit it from the schedule" ) def _initialize_stages(self, args: Any, kwargs: Any, target: Any = None, loss_kwargs: Any = None) -> None: reinit_for_mode_switch = self._stages_forward_initialized and ( self._has_backward != self._stages_backward_initialized ) forward_initialized_before = self._stages_forward_initialized ( self._stages_forward_initialized, self._stages_backward_initialized, ) = self._initialize_pp_stages( self._stages, args, kwargs, target, self._stages_forward_initialized, self._stages_backward_initialized, loss_kwargs, ) if self._stages_forward_initialized and ( not forward_initialized_before or reinit_for_mode_switch ): self._validate_adjacent_stage_communication() def step( self, *args: Any, target: Any = None, losses: list[Any] | None = None, return_outputs: bool = True, loss_kwargs: dict[str, Any] | None = None, arg_mbs: Any = None, kwarg_mbs: Any = None, target_mbs: Any = None, **kwargs: Any, ) -> Any: if ( self._has_backward and self._backward_requires_autograd and not tp.is_grad_enabled() ): raise RuntimeError( "step() requires gradients to be enabled for backward computation" ) for stage in self._stages: stage.has_backward = self._has_backward stage.clear_runtime_states() args_split, kwargs_split, targets_split = self._get_microbatch_inputs( args, kwargs, target, arg_mbs, kwarg_mbs, target_mbs, ) self._losses = [] self._initialize_stages( tuple(args_split[0]) if args_split else args, dict(kwargs_split[0]) if kwargs_split else kwargs, targets_split[0] if targets_split else None, loss_kwargs, ) self._step_microbatches( args_split, kwargs_split, targets_split, losses, return_outputs, loss_kwargs=loss_kwargs, ) if return_outputs: for stage in self._stages: if stage.is_last and stage.output_chunks: return self._merge_outputs(stage.output_chunks) return None def _step_microbatches(self, arg_mbs: list[tuple[Any, ...]], kwarg_mbs: list[dict[str, Any]], target_mbs: list[Any] | None, losses: list[Any] | None, return_outputs: bool = True, loss_kwargs: dict[str, Any] | None = None) -> Any: if self._stages_are_local(): return self._step_local_stages( arg_mbs, kwarg_mbs, target_mbs, losses, return_outputs, loss_kwargs, ) return self._step_distributed_stages( arg_mbs, kwarg_mbs, target_mbs, losses, return_outputs, loss_kwargs, ) def _stages_are_local(self) -> bool: if not self._stages: return True if not dist.is_initialized(): return True if len(self._stages) == 1: stage = self._stages[0] return stage.num_stages <= 1 or stage.group_size <= 1 first_rank = self._stages[0].group_rank return all(stage.group_rank == first_rank for stage in self._stages) def _step_local_stages( self, arg_mbs: list[tuple[Any, ...]], kwarg_mbs: list[dict[str, Any]], target_mbs: list[Any] | None, losses: list[Any] | None, return_outputs: bool, loss_kwargs: dict[str, Any] | None, ) -> Any: for stage in self._stages: stage.clear_runtime_states() outputs = [] for index, (args, kwargs) in enumerate(zip(arg_mbs, kwarg_mbs)): value = self._stages[0].forward_one_chunk(index, args, kwargs) for stage in self._stages[1:]: stage.set_local_fwd_input(value, index) value = stage.forward_one_chunk( index, _normalize_stage_args(value), {} ) outputs.append(value) loss = self._maybe_compute_loss(self._stages[-1], value, target_mbs, index, loss_kwargs) if self._has_backward: for index in reversed(range(len(outputs))): loss = self._maybe_get_loss(self._stages[-1], index) next_grad = self._stages[-1].backward_one_chunk(index, loss=loss) for stage_index in range(len(self._stages) - 2, -1, -1): stage = self._stages[stage_index] stage.set_local_bwd_input(next_grad, index) next_grad = stage.backward_one_chunk(index) if self._has_backward and self._scale_grads: for stage in self._stages: stage.scale_grads(self._n_microbatches) if losses is not None: losses.extend(self._losses) return self._merge_outputs(outputs) if return_outputs else None def _step_distributed_stages( self, arg_mbs: list[tuple[Any, ...]], kwarg_mbs: list[dict[str, Any]], target_mbs: list[Any] | None, losses: list[Any] | None, return_outputs: bool, loss_kwargs: dict[str, Any] | None, ) -> Any: for stage in self._stages: stage.clear_runtime_states() outputs: list[Any] = [] send_works: list[Any] = [] for index, (args, kwargs) in enumerate(zip(arg_mbs, kwarg_mbs)): for stage in self._stages: for work in _run_p2p(stage.get_fwd_recv_ops(index)): work.wait() output = stage.forward_one_chunk(index, args, kwargs) if stage.is_last: outputs.append(output) loss = self._maybe_compute_loss( stage, output, target_mbs, index, loss_kwargs ) send_works.extend(_run_p2p(stage.get_fwd_send_ops(index))) for work in send_works: work.wait() if self._has_backward: for index in reversed(range(self._n_microbatches)): for stage in reversed(self._stages): for work in _run_p2p(stage.get_bwd_recv_ops(index)): work.wait() loss = self._maybe_get_loss(stage, index) stage.backward_one_chunk(index, loss=loss) for work in _run_p2p(stage.get_bwd_send_ops(index)): work.wait() if self._scale_grads: for stage in self._stages: stage.scale_grads(self._n_microbatches) if losses is not None: losses.extend(self._losses) if not return_outputs or not outputs: return None return self._merge_outputs(outputs) def _validate_adjacent_stage_communication(self) -> None: def check_stage_indices( stage_index: int, direction: str, actual: set[int], expected: set[int], ) -> None: non_adjacent = actual - expected if non_adjacent: raise RuntimeError( f"stage {stage_index} has non-adjacent {direction} stages " f"{sorted(non_adjacent)}; expected only {sorted(expected)}" ) for stage in self._stages: stage_index = stage.stage_index forward_sources = { int(getattr(info, "source")) for info in stage.args_recv_info.get(0, ()) if getattr(info, "source", None) is not None } expected_sources = set() if stage.is_first else {stage_index - 1} check_stage_indices( stage_index, "forward receive", forward_sources, expected_sources, ) forward_destinations = { int(destination) for destinations in stage.act_send_info.values() for destination in destinations if destination is not None } expected_destinations = set() if stage.is_last else {stage_index + 1} check_stage_indices( stage_index, "forward send", forward_destinations, expected_destinations, ) def _validate_and_set_stage_mapping(self, actions: Any) -> None: self.stage_index_to_group_rank = _validate_schedule( actions, self.pp_group_size, self._num_stages, self._n_microbatches, ) for stage in self._stages: stage.stage_index_to_group_rank = dict(self.stage_index_to_group_rank) def _dump_csv(self, filename: str) -> None: with open(filename, "w", newline="", encoding="utf-8") as stream: writer = csv.writer(stream) for rank in sorted(self.pipeline_order): writer.writerow(self.pipeline_order[rank]) def _load_csv( self, filename: str, format: Literal["compute_only", "compute_comms"] = "compute_only", ) -> dict[int, list[_Action | None]]: if format != "compute_only": raise AssertionError(f"format must be compute_only, got {format}") with open(filename, newline="", encoding="utf-8") as stream: actions = { rank: [_Action.from_str(value) for value in row] for rank, row in enumerate(csv.reader(stream)) } self.pipeline_order = actions self._validate_and_set_stage_mapping(actions) return actions def _get_pipeline_order(self) -> dict[int, list[_Action | None]]: owned: dict[int, list[int]] = {rank: [] for rank in range(self.pp_group_size)} for stage_index in range(self._num_stages): rank = self.stage_index_to_group_rank[stage_index] owned.setdefault(rank, []).append(stage_index) pipeline_order: dict[int, list[_Action | None]] = {} for rank in range(self.pp_group_size): actions: list[_Action | None] = [None] * rank for stage_index in owned.get(rank, ()): actions.extend( _Action(stage_index, _ComputationType.FORWARD, microbatch) for microbatch in range(self._n_microbatches) ) if self._has_backward: actions.extend([None] * (2 * max(0, self.pp_group_size - 1 - rank))) for stage_index in reversed(owned.get(rank, ())): actions.extend( _Action(stage_index, _ComputationType.FULL_BACKWARD, microbatch) for microbatch in reversed(range(self._n_microbatches)) ) actions = _add_reduce_grad(actions, self._n_microbatches) pipeline_order[rank] = actions return pipeline_order @dataclass class _PipelineContext: schedule_ref: _PipelineSchedule arg_mbs: list[tuple[Any, ...]] | None = None kwarg_mbs: list[dict[str, Any]] | None = None target_mbs: list[Any] | None = None losses: list[Any] | None = None class _CustomFunctionProtocol(Protocol): def __call__(self, action: _Action, ctx: _PipelineContext) -> None: ... class _PipelineScheduleRuntime(PipelineScheduleMulti): def __init__(self, *args: Any, **kwargs: Any) -> None: self._defer_pp_recv = bool(kwargs.pop("defer_pp_recv", False)) max_active_stages = kwargs.pop("max_active_stages", 3) self._max_active_stages = 3 if max_active_stages is None else int(max_active_stages) super().__init__(*args, **kwargs) self._comp_type_to_function_map: dict[_ComputationType, Callable[..., Any]] = {} self.backward_counter: Counter[int] = Counter() self.bwd_recv_ops: dict[tuple[int, int], list[Any]] = {} self.fwd_recv_ops: dict[tuple[int, int], list[Any]] = {} self.unshard_ops: dict[int, list[UnshardHandle]] = defaultdict(list) self.unsharded_stages: set[int] = set() self.pipeline_order_with_comms: dict[int, list[_Action | None]] | None = None def register_custom_function( self, computation_type: _ComputationType, custom_function: _CustomFunctionProtocol, ) -> None: supported = { _ComputationType.FORWARD, _ComputationType.FULL_BACKWARD, _ComputationType.BACKWARD_INPUT, _ComputationType.BACKWARD_WEIGHT, _ComputationType.OVERLAP_F_B, _ComputationType.UNSHARD, _ComputationType.RESHARD, _ComputationType.REDUCE_GRAD, } if computation_type not in supported: raise ValueError(f"invalid computation type {computation_type}") if computation_type in self._comp_type_to_function_map: logger.warning( "computation type %s is already registered; replacing it", computation_type, ) self._comp_type_to_function_map[computation_type] = custom_function def _prepare_schedule_with_comms( self, actions: dict[int, list[_Action | None]], format: Literal["compute_only", "compute_comms"] = "compute_only", ) -> None: super()._validate_and_set_stage_mapping(actions) if format == "compute_comms": lowered: dict[int, list[_Action | None]] = {} for rank, rank_actions in actions.items(): if any(action is None for action in rank_actions): raise AssertionError("communication schedules cannot contain empty actions") lowered[rank] = list(rank_actions) self.pipeline_order_with_comms = lowered return if format != "compute_only": raise NotImplementedError(f"{format=} is not implemented") for rank, rank_actions in actions.items(): for index, action in enumerate(rank_actions): if action is not None and not action.is_compute_op: raise ValueError( f"expected compute-only action at rank {rank}, position {index}: {action}" ) lowered = { rank: _add_unshard_reshard( rank_actions, max_active_stages=self._max_active_stages ) for rank, rank_actions in actions.items() } lowered = { rank: _add_reduce_grad(rank_actions, self._n_microbatches) for rank, rank_actions in lowered.items() } lowered = _add_send_recv( lowered, stage_to_rank=lambda stage: self.stage_index_to_group_rank[stage], num_stages=self._num_stages, ) if self._defer_pp_recv: lowered = _defer_recv_ops( lowered, stage_to_rank=lambda stage: self.stage_index_to_group_rank[stage], ) self.pipeline_order_with_comms = lowered def _load_csv( self, filename: str, format: Literal["compute_only", "compute_comms"] = "compute_only", ) -> None: if format == "compute_only": actions = super()._load_csv(filename) self.pipeline_order = actions self._prepare_schedule_with_comms(actions) return if format != "compute_comms": raise NotImplementedError(f"{format=} is not implemented") with open(filename, newline="", encoding="utf-8") as stream: actions = { rank: [_Action.from_str(value) for value in row] for rank, row in enumerate(csv.reader(stream)) } self._prepare_schedule_with_comms(actions, format=format) def _dump_csv( self, filename: str, format: Literal["compute_only", "compute_comms"] = "compute_comms", ) -> None: if format == "compute_only": actions = self.pipeline_order elif format == "compute_comms": actions = self.pipeline_order_with_comms else: raise NotImplementedError(f"{format=} is not implemented") with open(filename, "w", newline="", encoding="utf-8") as stream: writer = csv.writer(stream) for rank in sorted(actions): writer.writerow(actions[rank]) def _simulate(self) -> Any: return _simulate_comms_compute( self.pipeline_order_with_comms, lambda stage: self.stage_index_to_group_rank[stage], self._num_stages, ) def _assert_unsharded(self, stage: Any) -> None: if not isinstance(stage.submod, FSDPModule): return stage_index = stage.stage_index if stage_index in self.unshard_ops: for handle in self.unshard_ops[stage_index]: handle.wait() del self.unshard_ops[stage_index] self.unsharded_stages.add(stage_index) if stage_index not in self.unsharded_stages: raise AssertionError(f"attempted to compute on sharded stage {stage_index}") def _step_microbatches( self, arg_mbs: list[tuple[Any, ...]] | None = None, kwarg_mbs: list[dict[str, Any]] | None = None, target_mbs: list[Any] | None = None, losses: list[Any] | None = None, return_outputs: bool = True, loss_kwargs: dict[str, Any] | None = None, ) -> None: arg_mbs, kwarg_mbs = self._check_inputs(arg_mbs, kwarg_mbs, target_mbs, losses) first_target = target_mbs[0] if target_mbs is not None else None self._initialize_stages(arg_mbs[0], kwarg_mbs[0], first_target, loss_kwargs) stage_index_to_stage = { stage.stage_index: stage for stage in self._stages } if self.pipeline_order_with_comms is None: raise AssertionError( "must prepare a schedule with communication actions before execution" ) self.fwd_recv_ops.clear() self.bwd_recv_ops.clear() self.unshard_ops.clear() self.unsharded_stages.clear() send_ops: list[list[Any]] = [] def perform_action(action: _Action) -> None: computation_type = action.computation_type microbatch = action.microbatch_index if microbatch is None and computation_type not in { _ComputationType.UNSHARD, _ComputationType.RESHARD, _ComputationType.REDUCE_GRAD, }: raise AssertionError(f"{action=} is missing a microbatch index") stage = stage_index_to_stage[action.stage_index] stage_index = action.stage_index stage_uses_fsdp = isinstance(stage.submod, FSDPModule) next_local = stage_index + 1 in stage_index_to_stage previous_local = stage_index - 1 in stage_index_to_stage mb_index = -1 if microbatch is None else microbatch if computation_type is _ComputationType.SEND_F: send_ops.append(_batch_p2p(stage.get_fwd_send_ops(mb_index))) elif computation_type is _ComputationType.SEND_B: send_ops.append(_batch_p2p(stage.get_bwd_send_ops(mb_index))) elif computation_type is _ComputationType.RECV_F: key = (stage_index, mb_index) if key in self.fwd_recv_ops: raise AssertionError(f"forward receive repeated for {key}") self.fwd_recv_ops[key] = _batch_p2p(stage.get_fwd_recv_ops(mb_index)) elif computation_type is _ComputationType.RECV_B: key = (stage_index, mb_index) if key in self.bwd_recv_ops: raise AssertionError(f"backward receive repeated for {key}") self.bwd_recv_ops[key] = _batch_p2p(stage.get_bwd_recv_ops(mb_index)) elif computation_type is _ComputationType.UNSHARD: if stage_uses_fsdp: if stage_index in self.unsharded_stages or stage_index in self.unshard_ops: raise AssertionError(f"unsharding stage {stage_index} twice") for submodule in stage.submod.modules(): if isinstance(submodule, FSDPModule): handle = cast(UnshardHandle, submodule.unshard(async_op=True)) self.unshard_ops[stage_index].append(handle) elif computation_type is _ComputationType.RESHARD: if stage_uses_fsdp: if stage_index not in self.unsharded_stages: raise AssertionError(f"resharding stage {stage_index} without unsharding") if stage_index in self.unshard_ops: raise AssertionError(f"resharding stage {stage_index} before unshard completion") for submodule in stage.submod.modules(): if isinstance(submodule, FSDPModule): submodule.reshard() self.unsharded_stages.remove(stage_index) elif computation_type is _ComputationType.FORWARD: self._assert_unsharded(stage) if not stage.is_first and not previous_local: key = (stage_index, mb_index) if key not in self.fwd_recv_ops: raise AssertionError(f"forward action {action} has no receive") _wait_batch_p2p(self.fwd_recv_ops.pop(key)) output = stage.forward_one_chunk( mb_index, arg_mbs[mb_index], kwarg_mbs[mb_index], save_forward_output=return_outputs, ) self._maybe_compute_loss(stage, output, target_mbs, mb_index, loss_kwargs) if next_local: stage_index_to_stage[stage_index + 1].set_local_fwd_input(output, mb_index) elif computation_type is _ComputationType.FULL_BACKWARD: self._assert_unsharded(stage) if not stage.is_last and not next_local: key = (stage_index, mb_index) if key not in self.bwd_recv_ops: raise AssertionError(f"backward action {action} has no receive") _wait_batch_p2p(self.bwd_recv_ops.pop(key)) self.backward_counter[stage_index] += 1 last_backward = self.backward_counter[stage_index] == self._n_microbatches stage.backward_one_chunk( mb_index, loss=self._maybe_get_loss(stage, mb_index), full_backward=True, last_backward=last_backward, ) if previous_local: stage_index_to_stage[stage_index - 1].set_local_bwd_input( stage.get_local_bwd_output(mb_index), mb_index ) elif computation_type is _ComputationType.BACKWARD_INPUT: self._assert_unsharded(stage) if not stage.is_last and not next_local: key = (stage_index, mb_index) if key not in self.bwd_recv_ops: raise AssertionError(f"backward action {action} has no receive") _wait_batch_p2p(self.bwd_recv_ops.pop(key)) stage.backward_one_chunk( mb_index, loss=self._maybe_get_loss(stage, mb_index), full_backward=False, last_backward=False, ) if previous_local: stage_index_to_stage[stage_index - 1].set_local_bwd_input( stage.get_local_bwd_output(mb_index), mb_index ) elif computation_type is _ComputationType.BACKWARD_WEIGHT: self._assert_unsharded(stage) self.backward_counter[stage_index] += 1 last_backward = self.backward_counter[stage_index] == self._n_microbatches stage.backward_weight_one_chunk(mb_index, last_backward=last_backward) elif computation_type is _ComputationType.REDUCE_GRAD: scale = self._n_microbatches if self._scale_grads else 1 stage.perform_reduce_grad(scale) else: raise ValueError(f"unknown or unsupported action {action}") self.backward_counter.clear() for time_step, action in enumerate(self.pipeline_order_with_comms[self.rank]): if action is None: continue try: custom = self._comp_type_to_function_map.get(action.computation_type) if custom is not None: custom( action, _PipelineContext(self, arg_mbs, kwarg_mbs, target_mbs, losses), ) elif action.computation_type is _ComputationType.OVERLAP_F_B: if action.sub_actions is None: raise AssertionError("overlap action must contain sub-actions") for sub_action in action.sub_actions: perform_action(sub_action) else: perform_action(action) except Exception: logger.error( "pipeline runtime failed at step %d on action %s", time_step, action, ) logger.error( "%s", _format_pipeline_order( self.pipeline_order_with_comms, error_step_number=time_step, ), ) raise while send_ops: _wait_batch_p2p(send_ops.pop()) if self.unshard_ops: raise AssertionError("unused unshard operations") self._update_losses(self._stages, losses) [docs] class ScheduleLoopedBFS(_PipelineScheduleRuntime): def __init__(self, stages: list[Any], n_microbatches: int, loss_fn: Any = None, output_merge_spec: Any = None, scale_grads: bool = True, backward_requires_autograd: bool = True, defer_pp_recv: bool = False, max_active_stages: int | None = None) -> None: super().__init__( stages, n_microbatches, loss_fn, output_merge_spec=output_merge_spec, scale_grads=scale_grads, backward_requires_autograd=backward_requires_autograd, defer_pp_recv=defer_pp_recv, max_active_stages=max_active_stages, ) self.defer_pp_recv = self._defer_pp_recv self.max_active_stages = self._max_active_stages self.pipeline_order = { rank: self._calculate_single_rank_operations(rank) for rank in range(self.pp_group_size) } self._prepare_schedule_with_comms(self.pipeline_order) def _calculate_single_rank_operations(self, rank: int) -> list[_Action | None]: local_stage_count = len(self._stages) stage_indices = range( rank, self.pp_group_size * local_stage_count, self.pp_group_size, ) rank_actions: list[_Action | None] = [None for _ in range(rank)] for stage_index in stage_indices: rank_actions.extend( _Action(stage_index, _ComputationType.FORWARD, microbatch) for microbatch in range(self._n_microbatches) ) rank_actions.extend( [None] * (2 * (self.pp_group_size - 1 - rank)) ) for stage_index in reversed(stage_indices): rank_actions.extend( _Action( stage_index, _ComputationType.FULL_BACKWARD, microbatch, ) for microbatch in reversed(range(self._n_microbatches)) ) return rank_actions [docs] class ScheduleInterleaved1F1B(_PipelineScheduleRuntime): def __init__( self, stages: list[Any], n_microbatches: int, loss_fn: Callable[..., Any] | None = None, args_chunk_spec: tuple[TensorChunkSpec, ...] | None = None, kwargs_chunk_spec: dict[str, TensorChunkSpec] | None = None, output_merge_spec: Any = None, scale_grads: bool = True, backward_requires_autograd: bool = True, defer_pp_recv: bool = False, max_active_stages: int = 3, ) -> None: self.pp_group_size = stages[0].group_size super().__init__( stages=stages, n_microbatches=n_microbatches, loss_fn=loss_fn, args_chunk_spec=args_chunk_spec, kwargs_chunk_spec=kwargs_chunk_spec, output_merge_spec=output_merge_spec, scale_grads=scale_grads, backward_requires_autograd=backward_requires_autograd, defer_pp_recv=defer_pp_recv, max_active_stages=max_active_stages, ) self.n_local_stages = len(stages) self.rank = stages[0].group_rank self.number_of_rounds = max(1, n_microbatches // self.pp_group_size) self.microbatches_per_round = n_microbatches // self.number_of_rounds if n_microbatches % self.number_of_rounds != 0: raise ValueError( "Interleaved 1F1B requires the microbatch count to be a multiple " f"of the round count ({self.number_of_rounds}), got {n_microbatches}" ) self.pipeline_order = { rank: self._calculate_single_rank_operations(rank) for rank in range(self.pp_group_size) } self._prepare_schedule_with_comms(self.pipeline_order) def _calculate_single_rank_operations(self, rank: int) -> list[_Action | None]: warmup_ops = _get_warmup_ops( rank, self.n_local_stages, self.microbatches_per_round, self.pp_group_size, self._n_microbatches, multiply_factor=2, ) microbatch_ops = self.n_local_stages * self._n_microbatches forward_backward_ops = microbatch_ops - warmup_ops cooldown_ops = microbatch_ops - forward_backward_ops def forward_stage_index(step: int) -> int: local_index = (step // self.microbatches_per_round) % self.n_local_stages return local_index * self.pp_group_size + rank def backward_stage_index(step: int) -> int: local_index = ( self.n_local_stages - 1 - ((step - warmup_ops) // self.microbatches_per_round) % self.n_local_stages ) return local_index * self.pp_group_size + rank logger.debug( "rank %s: warmup=%s steady=%s cooldown=%s", rank, warmup_ops, forward_backward_ops, cooldown_ops, ) return _get_1f1b_rank_ops( self.n_local_stages, self.pp_group_size, warmup_ops, forward_backward_ops, cooldown_ops, rank, forward_stage_index, backward_stage_index, ) [docs] class ScheduleInterleavedZeroBubble(_PipelineScheduleRuntime): def __init__( self, stages: list[Any], n_microbatches: int, loss_fn: Callable[..., Any] | None = None, args_chunk_spec: tuple[TensorChunkSpec, ...] | None = None, kwargs_chunk_spec: dict[str, TensorChunkSpec] | None = None, output_merge_spec: Any = None, scale_grads: bool = True, backward_requires_autograd: bool = True, defer_pp_recv: bool = False, max_active_stages: int = 3, ) -> None: _check_torch_compile_compatibility(stages, self.__class__.__name__) self.pp_group_size = stages[0].group_size super().__init__( stages=stages, n_microbatches=n_microbatches, loss_fn=loss_fn, args_chunk_spec=args_chunk_spec, kwargs_chunk_spec=kwargs_chunk_spec, output_merge_spec=output_merge_spec, scale_grads=scale_grads, backward_requires_autograd=backward_requires_autograd, defer_pp_recv=defer_pp_recv, max_active_stages=max_active_stages, ) self.n_local_stages = len(stages) self.rank = stages[0].group_rank self.number_of_rounds = max(1, n_microbatches // self.pp_group_size) self.microbatches_per_round = n_microbatches // self.number_of_rounds if n_microbatches % self.number_of_rounds != 0: raise ValueError( "Zero bubble requires the microbatch count to be a multiple of " f"the round count ({self.number_of_rounds}), got {n_microbatches}" ) self.pipeline_order = { rank: self._calculate_single_rank_operations(rank) for rank in range(self.pp_group_size) } self.pipeline_order = self._add_bubbles_to_actions( self.n_local_stages * self.pp_group_size ) self._prepare_schedule_with_comms(self.pipeline_order) def _calculate_single_rank_operations(self, rank: int) -> list[_Action | None]: warmup_ops = _get_warmup_ops( rank, self.n_local_stages, self.microbatches_per_round, self.pp_group_size, self._n_microbatches, multiply_factor=1, ) microbatch_ops = self.n_local_stages * self._n_microbatches forward_backward_ops = microbatch_ops - warmup_ops cooldown_ops = microbatch_ops - forward_backward_ops def forward_stage_index(step: int) -> int: local_index = (step // self.microbatches_per_round) % self.n_local_stages return local_index * self.pp_group_size + rank def backward_stage_index(step: int) -> int: local_index = ( self.n_local_stages - 1 - ((step - warmup_ops) // self.microbatches_per_round) % self.n_local_stages ) return local_index * self.pp_group_size + rank return _get_1f1b_rank_ops( self.n_local_stages, self.pp_group_size, warmup_ops, forward_backward_ops, cooldown_ops, rank, forward_stage_index, backward_stage_index, rank, enable_zero_bubble=True, ) def _add_bubbles_to_actions( self, num_stages_global: int ) -> dict[int, list[_Action | None]]: actions = self.pipeline_order def need_bubble( stage: int, operation: _ComputationType, microbatch: int | None, seen_ops: set[tuple[int, _ComputationType, int]], ) -> bool: if operation is _ComputationType.FORWARD: return stage != 0 and ( stage - 1, operation, microbatch, ) not in seen_ops if operation is _ComputationType.FULL_BACKWARD: if stage == num_stages_global - 1: return ( stage, _ComputationType.FORWARD, microbatch, ) not in seen_ops return ( stage + 1, operation, microbatch, ) not in seen_ops return False seen_ops: set[tuple[int, _ComputationType, int]] = set() result: dict[int, list[_Action | None]] = { rank: [] for rank in range(self.pp_group_size) } next_pointer = {rank: 0 for rank in range(self.pp_group_size)} bubbles_added = {rank: 0 for rank in range(self.pp_group_size)} total_bubbles_added = 0 while True: should_stop = True temporary_seen: set[tuple[int, _ComputationType, int]] = set() for rank in range(self.pp_group_size): timestamp = next_pointer[rank] if timestamp >= len(actions[rank]): continue should_stop = False action = actions[rank][timestamp] if action is None: result[rank].append(None) next_pointer[rank] += 1 continue stage_index, operation, microbatch, _ = action if not need_bubble( stage_index, operation, microbatch, seen_ops, ): result[rank].append(action) if microbatch is not None: temporary_seen.add((stage_index, operation, microbatch)) next_pointer[rank] += 1 else: result[rank].append(None) bubbles_added[rank] += 1 seen_ops.update(temporary_seen) if should_stop: break if total_bubbles_added > 0: logger.warning( "non-zero bubbles added: total=%s by-rank=%s", total_bubbles_added, bubbles_added, ) return result [docs] class ScheduleZBVZeroBubble(_PipelineScheduleRuntime): def __init__( self, stages: list[Any], n_microbatches: int, loss_fn: Callable[..., Any] | None = None, args_chunk_spec: tuple[TensorChunkSpec, ...] | None = None, kwargs_chunk_spec: dict[str, TensorChunkSpec] | None = None, output_merge_spec: Any = None, scale_grads: bool = True, backward_requires_autograd: bool = True, defer_pp_recv: bool = False, max_active_stages: int = 3, ) -> None: _check_torch_compile_compatibility(stages, self.__class__.__name__) self.pp_group_size = stages[0].group_size super().__init__( stages=stages, n_microbatches=n_microbatches, loss_fn=loss_fn, args_chunk_spec=args_chunk_spec, kwargs_chunk_spec=kwargs_chunk_spec, output_merge_spec=output_merge_spec, scale_grads=scale_grads, backward_requires_autograd=backward_requires_autograd, defer_pp_recv=defer_pp_recv, max_active_stages=max_active_stages, ) self.stage_index_to_group_rank = generate_stage_to_rank_mapping( self.pp_group_size, self._num_stages, style="v", ) for stage in self._stages: stage.stage_index_to_group_rank = self.stage_index_to_group_rank self.n_local_stages = len(stages) if self.n_local_stages != 2: raise ValueError( "ZBV requires exactly two stages per rank, " f"got {self.n_local_stages}" ) self.rank = stages[0].group_rank self.num_stages = stages[0].num_stages self.pipeline_order = { rank: self._calculate_single_rank_operations(rank) for rank in range(self.pp_group_size) } self._prepare_schedule_with_comms(self.pipeline_order) def _calculate_single_rank_operations(self, rank: int) -> list[_Action | None]: microbatch_count = max(2 * self.pp_group_size - 1, self._n_microbatches) rank_actions: list[_Action | None] = [None for _ in range(rank)] forward_chunk0 = 0 forward_chunk1 = 0 backward_chunk0 = 0 backward_chunk1 = 0 warmup_first = 2 * (self.pp_group_size - rank) - 1 stage_chunk0 = rank stage_chunk1 = self.num_stages - 1 - rank for _ in range(warmup_first): rank_actions.append( _Action( stage_chunk0, _ComputationType.FORWARD, forward_chunk0, ) ) forward_chunk0 += 1 warmup_second = rank for _ in range(warmup_second): rank_actions.append( _Action(stage_chunk1, _ComputationType.FORWARD, forward_chunk1) ) forward_chunk1 += 1 rank_actions.append( _Action(stage_chunk0, _ComputationType.FORWARD, forward_chunk0) ) forward_chunk0 += 1 warmup_third = self.pp_group_size - rank for _ in range(warmup_third): rank_actions.append( _Action(stage_chunk1, _ComputationType.FORWARD, forward_chunk1) ) forward_chunk1 += 1 rank_actions.append( _Action(stage_chunk1, _ComputationType.BACKWARD_INPUT, backward_chunk1) ) rank_actions.append( _Action(stage_chunk1, _ComputationType.BACKWARD_WEIGHT, backward_chunk1) ) backward_chunk1 += 1 while forward_chunk1 < forward_chunk0 or forward_chunk0 < microbatch_count: if forward_chunk0 < microbatch_count: rank_actions.append( _Action(stage_chunk0, _ComputationType.FORWARD, forward_chunk0) ) forward_chunk0 += 1 rank_actions.append( _Action(stage_chunk0, _ComputationType.BACKWARD_INPUT, backward_chunk0) ) rank_actions.append( _Action(stage_chunk0, _ComputationType.BACKWARD_WEIGHT, backward_chunk0) ) backward_chunk0 += 1 rank_actions.append( _Action(stage_chunk1, _ComputationType.FORWARD, forward_chunk1) ) forward_chunk1 += 1 rank_actions.append( _Action(stage_chunk1, _ComputationType.BACKWARD_INPUT, backward_chunk1) ) rank_actions.append( _Action(stage_chunk1, _ComputationType.BACKWARD_WEIGHT, backward_chunk1) ) backward_chunk1 += 1 weight_chunk0 = backward_chunk0 weight_chunk1 = backward_chunk1 for _ in range(rank): rank_actions.append( _Action(stage_chunk0, _ComputationType.BACKWARD_INPUT, backward_chunk0) ) backward_chunk0 += 1 rank_actions.append( _Action(stage_chunk1, _ComputationType.BACKWARD_INPUT, backward_chunk1) ) backward_chunk1 += 1 for _ in range(self.pp_group_size - rank): rank_actions.append( _Action(stage_chunk0, _ComputationType.BACKWARD_INPUT, backward_chunk0) ) backward_chunk0 += 1 rank_actions.append( _Action(stage_chunk0, _ComputationType.BACKWARD_WEIGHT, weight_chunk0) ) weight_chunk0 += 1 while weight_chunk1 < backward_chunk1: rank_actions.append( _Action(stage_chunk1, _ComputationType.BACKWARD_WEIGHT, weight_chunk1) ) weight_chunk1 += 1 while weight_chunk0 < backward_chunk0: rank_actions.append( _Action(stage_chunk0, _ComputationType.BACKWARD_WEIGHT, weight_chunk0) ) weight_chunk0 += 1 if not (weight_chunk0 == backward_chunk0 == forward_chunk0): raise AssertionError( "stage chunk 0 action counts do not match" ) if not (weight_chunk1 == backward_chunk1 == forward_chunk1): raise AssertionError( "stage chunk 1 action counts do not match" ) return [ action if action is not None and action.microbatch_index is not None and action.microbatch_index < self._n_microbatches else None for action in rank_actions ] [docs] class ScheduleDualPipeV(_PipelineScheduleRuntime): """Run the bidirectional local stage schedule.""" def __init__( self, stages: list[Any], n_microbatches: int, loss_fn: Callable[..., Any] | None = None, args_chunk_spec: tuple[TensorChunkSpec, ...] | None = None, kwargs_chunk_spec: dict[str, TensorChunkSpec] | None = None, output_merge_spec: Any = None, scale_grads: bool = True, backward_requires_autograd: bool = True, defer_pp_recv: bool = False, max_active_stages: int = 3, ) -> None: _check_torch_compile_compatibility(stages, self.__class__.__name__) self.pp_group_size = stages[0].group_size super().__init__( stages=stages, n_microbatches=n_microbatches, loss_fn=loss_fn, args_chunk_spec=args_chunk_spec, kwargs_chunk_spec=kwargs_chunk_spec, output_merge_spec=output_merge_spec, scale_grads=scale_grads, backward_requires_autograd=backward_requires_autograd, defer_pp_recv=defer_pp_recv, max_active_stages=max_active_stages, ) self.stage_index_to_group_rank = generate_stage_to_rank_mapping( self.pp_group_size, self._num_stages, style="v", ) for stage in self._stages: stage.stage_index_to_group_rank = self.stage_index_to_group_rank self.n_local_stages = len(stages) if self.n_local_stages != 2: raise ValueError( "ZBV requires exactly 2 stages per rank, but got " f"{self.n_local_stages}." ) if n_microbatches < self._num_stages: raise ValueError( "DualPipeV requires at least as many microbatches as stages, but got " f"{n_microbatches} microbatches and {self._num_stages} stages." ) self.rank = stages[0].group_rank self.num_stages = stages[0].num_stages self.pipeline_order = { rank: self._calculate_single_rank_operations(rank) for rank in range(self.pp_group_size) } self._prepare_schedule_with_comms(self.pipeline_order) def _calculate_single_rank_operations(self, rank: int) -> list[_Action | None]: actions: list[_Action | None] = [] counters: dict[tuple[int, _ComputationType], int] = {} weight_queue: list[tuple[int, int]] = [] num_ranks = self.pp_group_size num_chunks = self._n_microbatches rank_to_stages = generate_rank_to_stage_mapping( num_ranks, num_ranks * 2, style="v" ) stage0_index, stage1_index = rank_to_stages[rank] def increment_backward_counts(stage_index: int) -> None: input_key = (stage_index, _ComputationType.BACKWARD_INPUT) weight_key = (stage_index, _ComputationType.BACKWARD_WEIGHT) counters[input_key] = counters.get(input_key, 0) + 1 counters[weight_key] = counters.get(weight_key, 0) + 1 def add_overlap_f_b( forward_stage: int, backward_stage: int, ) -> None: forward_key = (forward_stage, _ComputationType.FORWARD) backward_key = (backward_stage, _ComputationType.BACKWARD_INPUT) forward_mb = counters.get(forward_key, 0) backward_mb = counters.get(backward_key, 0) sub_actions = ( _Action(forward_stage, _ComputationType.FORWARD, forward_mb), _Action(backward_stage, _ComputationType.FULL_BACKWARD, backward_mb), ) actions.append( _Action(-1, _ComputationType.OVERLAP_F_B, None, sub_actions) ) counters[forward_key] = forward_mb + 1 increment_backward_counts(backward_stage) def add_action( stage_index: int, computation_type: _ComputationType, ) -> None: key = ( (stage_index, computation_type) if computation_type != _ComputationType.FULL_BACKWARD else (stage_index, _ComputationType.BACKWARD_INPUT) ) mb_index = counters.get(key, 0) actions.append(_Action(stage_index, computation_type, mb_index)) if computation_type == _ComputationType.FULL_BACKWARD: increment_backward_counts(stage_index) else: if computation_type == _ComputationType.BACKWARD_INPUT: weight_queue.append((stage_index, mb_index)) counters[key] = mb_index + 1 def add_weight_action_if_pending() -> None: if not weight_queue: return actual_stage_index, weight_mb_index = weight_queue.pop(0) actions.append( _Action( actual_stage_index, _ComputationType.BACKWARD_WEIGHT, weight_mb_index, ) ) weight_key = (actual_stage_index, _ComputationType.BACKWARD_WEIGHT) counters[weight_key] = counters.get(weight_key, 0) + 1 step_1 = (num_ranks - rank - 1) * 2 for _ in range(step_1): add_action(stage0_index, _ComputationType.FORWARD) step_2 = rank + 1 for _ in range(step_2): add_action(stage0_index, _ComputationType.FORWARD) add_action(stage1_index, _ComputationType.FORWARD) step_3 = num_ranks - rank - 1 for _ in range(step_3): add_action(stage1_index, _ComputationType.BACKWARD_INPUT) add_weight_action_if_pending() add_action(stage1_index, _ComputationType.FORWARD) step_4 = num_chunks - num_ranks * 2 + rank + 1 for index in range(step_4): if index == 0 and rank == num_ranks - 1: add_action(stage0_index, _ComputationType.FORWARD) add_action(stage1_index, _ComputationType.FULL_BACKWARD) else: add_overlap_f_b(stage0_index, stage1_index) add_overlap_f_b(stage1_index, stage0_index) step_5 = num_ranks - rank - 1 for _ in range(step_5): add_action(stage1_index, _ComputationType.FULL_BACKWARD) add_overlap_f_b(stage1_index, stage0_index) step_6 = rank + 1 enable_zb = False for index in range(step_6): if index == step_6 // 2 and rank % 2 == 1: enable_zb = True comp_type = ( _ComputationType.BACKWARD_INPUT if enable_zb else _ComputationType.FULL_BACKWARD ) add_action(stage1_index, comp_type) if index == step_6 // 2 and rank % 2 == 0: enable_zb = True comp_type = ( _ComputationType.BACKWARD_INPUT if enable_zb else _ComputationType.FULL_BACKWARD ) add_action(stage0_index, comp_type) step_7 = num_ranks - rank - 1 for _ in range(step_7): add_weight_action_if_pending() comp_type = ( _ComputationType.BACKWARD_INPUT if enable_zb else _ComputationType.FULL_BACKWARD ) add_action(stage0_index, comp_type) step_8 = rank + 1 for _ in range(step_8): add_weight_action_if_pending() return actions def _requires_reduce_grad(action_type: _ComputationType) -> bool: return action_type in { _ComputationType.BACKWARD_WEIGHT, _ComputationType.FULL_BACKWARD, } def _add_reduce_grad( actions: list[_Action | None], n_microbatches: int ) -> list[_Action | None]: if n_microbatches <= 0: raise ValueError("n_microbatches must be positive") result: list[_Action | None] = [] counts: dict[int, int] = defaultdict(int) for action in actions: if action is None: result.append(None) continue result.append(action) leaves = action.sub_actions or (action,) for leaf in leaves: if not _requires_reduce_grad(leaf.computation_type): continue counts[leaf.stage_index] += 1 if counts[leaf.stage_index] == n_microbatches: result.append( _Action(leaf.stage_index, _ComputationType.REDUCE_GRAD, None) ) return result def _split_backward_pipeline_order( num_stages: int, pp_group_size: int, n_microbatches: int, stage_to_rank: dict[int, int], ) -> dict[int, list[_Action | None]]: owned: dict[int, list[int]] = {rank: [] for rank in range(pp_group_size)} for stage_index in range(num_stages): owned[int(stage_to_rank[stage_index])].append(stage_index) result: dict[int, list[_Action | None]] = {} for rank in range(pp_group_size): actions: list[_Action | None] = [None] * rank for stage_index in owned[rank]: actions.extend( _Action(stage_index, _ComputationType.FORWARD, microbatch) for microbatch in range(n_microbatches) ) for stage_index in reversed(owned[rank]): actions.extend( _Action(stage_index, _ComputationType.BACKWARD_INPUT, microbatch) for microbatch in reversed(range(n_microbatches)) ) actions.extend( _Action(stage_index, _ComputationType.BACKWARD_WEIGHT, microbatch) for microbatch in reversed(range(n_microbatches)) ) result[rank] = _add_reduce_grad(actions, n_microbatches) return result def _add_unshard_reshard( compute_actions: list[_Action | None], max_active_stages: int = 3 ) -> list[_Action]: if max_active_stages <= 0: raise ValueError("max_active_stages must be positive") active: set[int] = set() result: list[_Action] = [] def next_stage_indices( count: int, actions: list[_Action | None] ) -> list[int]: seen: set[int] = set() stages: list[int] = [] for action in actions: if action is None: continue leaves = action.sub_actions or (action,) for leaf in leaves: if leaf.stage_index not in seen: seen.add(leaf.stage_index) stages.append(leaf.stage_index) if len(stages) >= count: break return stages actions = list(compute_actions) for index, action in enumerate(actions): if action is None: continue upcoming = next_stage_indices(max_active_stages, actions[index:]) for stage_index in [stage for stage in active if stage not in upcoming]: active.remove(stage_index) result.append(_Action(stage_index, _ComputationType.RESHARD)) for stage_index in upcoming: if stage_index not in active: active.add(stage_index) result.append(_Action(stage_index, _ComputationType.UNSHARD)) result.append(action) for stage_index in list(active): result.append(_Action(stage_index, _ComputationType.RESHARD)) return result def _merge_bw(compute_actions: list[_Action | None]) -> list[_Action]: pending = list(compute_actions) result: list[_Action] = [] while pending: action = pending.pop(0) if action is None: continue while pending and pending[0] is None: pending.pop(0) if ( action.computation_type is _ComputationType.BACKWARD_INPUT and pending and pending[0] is not None and pending[0].computation_type is _ComputationType.BACKWARD_WEIGHT and action.stage_index == pending[0].stage_index and action.microbatch_index == pending[0].microbatch_index ): result.append( _Action( action.stage_index, _ComputationType.FULL_BACKWARD, action.microbatch_index, ) ) pending.pop(0) else: result.append(action) return result def _add_send_recv( compute_actions: dict[int, list[_Action | None]], stage_to_rank: Any, num_stages: int, ) -> dict[int, list[_Action | None]]: if callable(stage_to_rank): rank_of = stage_to_rank else: rank_of = lambda stage: stage_to_rank[int(stage)] remaining = { int(rank): list(actions) for rank, actions in compute_actions.items() } result: dict[int, list[_Action | None]] = {rank: [] for rank in remaining} previous: dict[int, set[_Action]] = {rank: set() for rank in remaining} def leaves(action: _Action) -> tuple[_Action, ...]: return action.sub_actions or (action,) def communicates(action: _Action) -> bool: if action.computation_type is _ComputationType.FORWARD: return ( action.stage_index < num_stages - 1 and rank_of(action.stage_index) != rank_of(action.stage_index + 1) ) if action.computation_type in { _ComputationType.BACKWARD_INPUT, _ComputationType.FULL_BACKWARD, }: return ( action.stage_index > 0 and rank_of(action.stage_index) != rank_of(action.stage_index - 1) ) return False def ready(action: _Action, done: set[_Action]) -> bool: if ( action.computation_type is _ComputationType.FORWARD and action.stage_index > 0 ): return ( _Action( action.stage_index, _ComputationType.RECV_F, action.microbatch_index, ) in done or _Action( action.stage_index - 1, _ComputationType.FORWARD, action.microbatch_index, ) in done ) if ( action.computation_type in {_ComputationType.BACKWARD_INPUT, _ComputationType.FULL_BACKWARD} and action.stage_index < num_stages - 1 ): return ( _Action( action.stage_index, _ComputationType.RECV_B, action.microbatch_index, ) in done or _Action( action.stage_index + 1, _ComputationType.BACKWARD_INPUT, action.microbatch_index, ) in done or _Action( action.stage_index + 1, _ComputationType.FULL_BACKWARD, action.microbatch_index, ) in done ) return True while remaining: progress = False for rank in sorted(tuple(remaining)): actions = remaining[rank] if not actions: del remaining[rank] continue action = actions[0] if action is None: result[rank].append(None) actions.pop(0) progress = True if not actions: del remaining[rank] continue action_leaves = leaves(action) if not all(ready(leaf, previous[rank]) for leaf in action_leaves): continue result[rank].append(action) for leaf in action_leaves: previous[rank].add(leaf) if not communicates(leaf): continue is_forward = leaf.computation_type is _ComputationType.FORWARD send_kind = ( _ComputationType.SEND_F if is_forward else _ComputationType.SEND_B ) recv_kind = ( _ComputationType.RECV_F if is_forward else _ComputationType.RECV_B ) peer_stage = leaf.stage_index + 1 if is_forward else leaf.stage_index - 1 send = _Action(leaf.stage_index, send_kind, leaf.microbatch_index) recv = _Action(peer_stage, recv_kind, leaf.microbatch_index) result[rank].append(send) previous[rank].add(send) peer_rank = int(rank_of(peer_stage)) if peer_rank not in result: raise ValueError(f"stage mapping points to unknown rank {peer_rank}") result[peer_rank].append(recv) previous[peer_rank].add(recv) actions.pop(0) progress = True if not actions: del remaining[rank] if not progress: raise ValueError("malformed pipeline schedule") return result def _defer_recv_ops( actions: list[_Action | None] | dict[int, list[_Action | None]], stage_to_rank: Any, ) -> list[_Action | None] | dict[int, list[_Action | None]]: rank_of = stage_to_rank if callable(stage_to_rank) else lambda stage: stage_to_rank[int(stage)] was_list = isinstance(actions, list) by_rank = {0: list(actions)} if was_list else { int(rank): list(rank_actions) for rank, rank_actions in actions.items() } result: dict[int, list[_Action | None]] = {} recv_types = {_ComputationType.RECV_F, _ComputationType.RECV_B} send_types = {_ComputationType.SEND_F, _ComputationType.SEND_B} def recv_peer(action: _Action) -> int: peer_stage = action.stage_index - 1 if action.computation_type is _ComputationType.RECV_F else action.stage_index + 1 return int(rank_of(peer_stage)) def send_peer(action: _Action) -> int: peer_stage = action.stage_index + 1 if action.computation_type is _ComputationType.SEND_F else action.stage_index - 1 return int(rank_of(peer_stage)) for rank, rank_actions in by_rank.items(): deferred: dict[tuple[int, _ComputationType, int | None], _Action] = {} output: list[_Action | None] = [] for action in rank_actions: if action is None: output.append(None) continue if action.computation_type in recv_types: key = (action.stage_index, action.computation_type, action.microbatch_index) deferred[key] = action continue if action.computation_type in send_types: peer = send_peer(action) if rank < peer: for key in tuple(deferred): if recv_peer(deferred[key]) == peer: output.append(deferred.pop(key)) leaves = action.sub_actions or (action,) for leaf in leaves: if leaf.computation_type is _ComputationType.FORWARD: key = (leaf.stage_index, _ComputationType.RECV_F, leaf.microbatch_index) elif leaf.computation_type in { _ComputationType.FULL_BACKWARD, _ComputationType.BACKWARD_INPUT, }: key = (leaf.stage_index, _ComputationType.RECV_B, leaf.microbatch_index) else: continue if key in deferred: output.append(deferred.pop(key)) output.append(action) if deferred: raise ValueError("every receive action must have a consuming compute action") result[rank] = output return result[0] if was_list else result def _validate_schedule( actions: Any, pp_group_size: int, num_stages: int, num_microbatches: int, ) -> dict[int, int]: if pp_group_size <= 0 or num_stages <= 0 or num_microbatches <= 0: raise ValueError("pipeline dimensions must be positive") if not isinstance(actions, dict) or len(actions) != pp_group_size: raise ValueError("schedule must provide one action list per rank") if set(actions) != set(range(pp_group_size)): raise ValueError("schedule ranks must be contiguous") stage_actions: dict[int, dict[_ComputationType, set[int]]] = { stage: { _ComputationType.FORWARD: set(), _ComputationType.BACKWARD_INPUT: set(), _ComputationType.BACKWARD_WEIGHT: set(), _ComputationType.FULL_BACKWARD: set(), } for stage in range(num_stages) } stage_index_to_rank: dict[int, int] = {} seen: set[tuple[int, _ComputationType, int | None]] = set() compute_types = { _ComputationType.FORWARD, _ComputationType.FULL_BACKWARD, _ComputationType.BACKWARD_INPUT, _ComputationType.BACKWARD_WEIGHT, } communication_types = { _ComputationType.SEND_F, _ComputationType.RECV_F, _ComputationType.SEND_B, _ComputationType.RECV_B, } def process_action(action: _Action, rank: int, step: int) -> None: if action.sub_actions is not None: if action.computation_type is not _ComputationType.OVERLAP_F_B: raise ValueError("only overlap actions may contain sub-actions") if not action.sub_actions: raise ValueError("an overlap action must contain sub-actions") for sub_action in action.sub_actions: if not isinstance(sub_action, _Action): raise TypeError("sub-actions must be _Action instances") process_action(sub_action, rank, step) return stage = action.stage_index kind = action.computation_type microbatch = action.microbatch_index if not 0 <= stage < num_stages: raise ValueError("action stage is outside the pipeline") if kind not in compute_types | communication_types | { _ComputationType.UNSHARD, _ComputationType.RESHARD, _ComputationType.REDUCE_GRAD, }: raise ValueError(f"unsupported pipeline action {kind!r}") if kind in compute_types | communication_types: if microbatch is None or not 0 <= microbatch < num_microbatches: raise ValueError("action microbatch is outside the schedule") elif microbatch is not None: raise ValueError("non-compute actions cannot carry a microbatch") previous_rank = stage_index_to_rank.get(stage) if previous_rank is not None and previous_rank != rank: raise ValueError( f"stage {stage} is assigned to ranks {previous_rank} and {rank}" ) stage_index_to_rank[stage] = rank if kind not in compute_types: return key = (stage, kind, microbatch) if key in seen: raise ValueError("a compute action occurs more than once") seen.add(key) if kind is _ComputationType.FORWARD: stage_actions[stage][kind].add(microbatch) return if kind is _ComputationType.FULL_BACKWARD: if microbatch not in stage_actions[stage][_ComputationType.FORWARD]: raise ValueError("backward ran before its forward") stage_actions[stage][kind].add(microbatch) return if kind is _ComputationType.BACKWARD_INPUT: if microbatch not in stage_actions[stage][_ComputationType.FORWARD]: raise ValueError("backward input ran before its forward") stage_actions[stage][kind].add(microbatch) return if microbatch not in stage_actions[stage][_ComputationType.BACKWARD_INPUT]: raise ValueError("backward weight ran before its input backward") stage_actions[stage][kind].add(microbatch) for rank, rank_actions in actions.items(): if not isinstance(rank_actions, list): raise TypeError(f"actions for rank {rank} must be a list") for step, action in enumerate(rank_actions): if action is None: continue if not isinstance(action, _Action): raise TypeError("schedule entries must be _Action instances") process_action(action, rank, step) for stage in range(num_stages): counts = stage_actions[stage] if len(counts[_ComputationType.FORWARD]) != num_microbatches: raise ValueError("schedule is missing a forward action") if len(counts[_ComputationType.BACKWARD_INPUT]) != len( counts[_ComputationType.BACKWARD_WEIGHT] ): raise ValueError("input and weight backward counts must match") if len(counts[_ComputationType.FULL_BACKWARD]) + len( counts[_ComputationType.BACKWARD_INPUT] ) != num_microbatches: raise ValueError("schedule is missing a backward action") if len(stage_index_to_rank) != num_stages: raise ValueError("schedule does not assign every pipeline stage") return stage_index_to_rank def _get_1f1b_rank_ops( n_local_stages: int, pp_group_size: int, warmup_ops: int, fwd_bwd_ops: int, cooldown_ops: int, rank: int, forward_stage_index: Any, backward_stage_index: Any, num_1f1b_microbatches: int = 0, enable_zero_bubble: bool = False, ) -> list[_Action | None]: if min(n_local_stages, pp_group_size) <= 0: raise ValueError("pipeline dimensions must be positive") if rank < 0 or rank >= pp_group_size: raise ValueError("rank is outside the pipeline group") if min(warmup_ops, fwd_bwd_ops, cooldown_ops) < 0: raise ValueError("operation counts must be non-negative") forward_counts: dict[int, int] = defaultdict(int) backward_counts: dict[int, int] = defaultdict(int) weight_counts: dict[int, int] = defaultdict(int) result: list[_Action | None] = [None] * rank backward_ids: list[int] = [] total_ops = warmup_ops + fwd_bwd_ops + cooldown_ops post_warmup = ( n_local_stages * pp_group_size + 2 * (pp_group_size - 1 - rank) - warmup_ops - rank ) if enable_zero_bubble: post_warmup = pp_group_size - rank - 1 for operation in range(total_ops): if operation < warmup_ops: stage = int(forward_stage_index(operation)) microbatch = forward_counts[stage] forward_counts[stage] += 1 result.append(_Action(stage, _ComputationType.FORWARD, microbatch)) if operation == warmup_ops - 1: result.extend([None] * max(0, post_warmup)) continue if operation < warmup_ops + fwd_bwd_ops: stage = int(forward_stage_index(operation)) microbatch = forward_counts[stage] forward_counts[stage] += 1 result.append(_Action(stage, _ComputationType.FORWARD, microbatch)) backward_stage = int(backward_stage_index(operation)) microbatch = backward_counts[backward_stage] backward_counts[backward_stage] += 1 backward_kind = ( _ComputationType.BACKWARD_INPUT if enable_zero_bubble else _ComputationType.FULL_BACKWARD ) result.append( _Action(backward_stage, backward_kind, microbatch) ) backward_ids.append(operation) if ( enable_zero_bubble and operation - warmup_ops >= num_1f1b_microbatches ): weight_index = sum(weight_counts.values()) weight_stage = int(backward_stage_index(backward_ids[weight_index])) microbatch = weight_counts[weight_stage] weight_counts[weight_stage] += 1 result.append( _Action( weight_stage, _ComputationType.BACKWARD_WEIGHT, microbatch, ) ) continue if not enable_zero_bubble: result.append(None) backward_stage = int(backward_stage_index(operation)) microbatch = backward_counts[backward_stage] backward_counts[backward_stage] += 1 backward_kind = ( _ComputationType.BACKWARD_INPUT if enable_zero_bubble else _ComputationType.FULL_BACKWARD ) result.append(_Action(backward_stage, backward_kind, microbatch)) backward_ids.append(operation) if ( enable_zero_bubble and operation - warmup_ops >= num_1f1b_microbatches ): weight_index = sum(weight_counts.values()) weight_stage = int(backward_stage_index(backward_ids[weight_index])) microbatch = weight_counts[weight_stage] weight_counts[weight_stage] += 1 result.append( _Action( weight_stage, _ComputationType.BACKWARD_WEIGHT, microbatch, ) ) while enable_zero_bubble and sum(weight_counts.values()) < len(backward_ids): weight_index = sum(weight_counts.values()) weight_stage = int(backward_stage_index(backward_ids[weight_index])) microbatch = weight_counts[weight_stage] weight_counts[weight_stage] += 1 result.append( _Action(weight_stage, _ComputationType.BACKWARD_WEIGHT, microbatch) ) return result def _get_warmup_ops( rank: int, n_local_stages: int, microbatches_per_round: int, pp_group_size: int, n_microbatches: int, multiply_factor: int = 2, ) -> int: warmups_last_stage = (n_local_stages - 1) * microbatches_per_round warmup_ops = warmups_last_stage + multiply_factor * (pp_group_size - 1 - rank) return min(warmup_ops, n_microbatches * n_local_stages) def get_schedule_class(schedule_name: str) -> type[_PipelineSchedule]: mapping = { "GPipe": ScheduleGPipe, "1F1B": Schedule1F1B, "Interleaved1F1B": ScheduleInterleaved1F1B, "LoopedBFS": ScheduleLoopedBFS, "InterleavedZeroBubble": ScheduleInterleavedZeroBubble, "ZBVZeroBubble": ScheduleZBVZeroBubble, "DualPipeV": ScheduleDualPipeV, "PipelineScheduleSingle": PipelineScheduleSingle, "PipelineScheduleMulti": PipelineScheduleMulti, } if not isinstance(schedule_name, str): raise TypeError("schedule name must be a string") normalized = {name.lower(): cls for name, cls in mapping.items()} try: return normalized[schedule_name.lower()] except KeyError as exc: raise ValueError(f"unknown pipeline schedule {schedule_name!r}") from exc def _simulate_comms_compute(pipeline_order: Any, stage_to_rank: Any, num_stages: int) -> Any: if not isinstance(pipeline_order, dict): raise TypeError("pipeline_order must be a rank-to-actions mapping") if callable(stage_to_rank): rank_of = stage_to_rank else: rank_of = lambda stage: stage_to_rank[int(stage)] pending = { int(rank): [action for action in actions if action is not None] for rank, actions in pipeline_order.items() } schedule: dict[int, list[_Action | None]] = { rank: [] for rank in sorted(pending) } completed: dict[int, set[_Action]] = { rank: set() for rank in pending } def leaves(action: _Action) -> tuple[_Action, ...]: return action.sub_actions or (action,) def ready_leaf(action: _Action, owner: int) -> bool: if action.stage_index < 0 or action.stage_index >= num_stages: raise ValueError("action stage is outside the pipeline") owner = int(rank_of(action.stage_index)) if action.stage_index >= 0 else owner done = completed[owner] kind = action.computation_type stage = action.stage_index microbatch = action.microbatch_index if kind is _ComputationType.FORWARD: if stage == 0: return True return ( _Action(stage, _ComputationType.RECV_F, microbatch) in done or _Action(stage - 1, _ComputationType.FORWARD, microbatch) in done ) if kind in { _ComputationType.BACKWARD_INPUT, _ComputationType.FULL_BACKWARD, }: if stage == num_stages - 1: return True return ( _Action(stage, _ComputationType.RECV_B, microbatch) in done or _Action(stage + 1, _ComputationType.BACKWARD_INPUT, microbatch) in done or _Action(stage + 1, _ComputationType.FULL_BACKWARD, microbatch) in done ) if kind is _ComputationType.SEND_F: return _Action(stage, _ComputationType.FORWARD, microbatch) in done if kind is _ComputationType.RECV_F: peer = stage - 1 return _Action(peer, _ComputationType.SEND_F, microbatch) in completed[int(rank_of(peer))] if kind is _ComputationType.SEND_B: return ( _Action(stage, _ComputationType.BACKWARD_INPUT, microbatch) in done or _Action(stage, _ComputationType.FULL_BACKWARD, microbatch) in done ) if kind is _ComputationType.RECV_B: peer = stage + 1 return _Action(peer, _ComputationType.SEND_B, microbatch) in completed[int(rank_of(peer))] if kind is _ComputationType.BACKWARD_WEIGHT: return True if kind in { _ComputationType.UNSHARD, _ComputationType.RESHARD, _ComputationType.REDUCE_GRAD, }: return True raise ValueError(f"unsupported pipeline action {kind!r}") def ready(action: _Action, owner: int) -> bool: return all(ready_leaf(leaf, owner) for leaf in leaves(action)) def mark_completed(action: _Action, owner: int) -> None: completed[owner].add(action) for leaf in leaves(action): leaf_owner = int(rank_of(leaf.stage_index)) completed[leaf_owner].add(leaf) while pending: progress = False for rank in sorted(tuple(pending)): actions = pending[rank] if not actions: del pending[rank] continue action = actions[0] if not ready(action, rank): schedule[rank].append(None) continue schedule[rank].append(action) mark_completed(action, rank) actions.pop(0) progress = True if not actions: del pending[rank] for rank in sorted(tuple(pending)): if not schedule[rank] or schedule[rank][-1] is not None: continue action = pending[rank][0] if ready(action, rank): schedule[rank][-1] = action mark_completed(action, rank) pending[rank].pop(0) if not pending[rank]: del pending[rank] progress = True if not progress: raise ValueError("pipeline schedule cannot make progress") return schedule def _dump_chrometrace(schedule: Any, filename: str) -> None: events: list[dict[str, Any]] = [] for rank in sorted(schedule): for timestep, action in enumerate(schedule[rank]): if action is None: continue events.append( { "name": str(action), "cat": ( "computation" if action.computation_type in { _ComputationType.FORWARD, _ComputationType.FULL_BACKWARD, _ComputationType.BACKWARD_WEIGHT, } else "communication" ), "ph": "X", "pid": rank, "tid": rank, "ts": timestep, "dur": 1, } ) import json with open(filename, "w", encoding="utf-8") as stream: json.dump({"traceEvents": events}, stream) def _check_torch_compile_compatibility(stages: Any, schedule_name: str) -> None: del stages, schedule_name def _default_merge_spec(value: Any) -> Any: if isinstance(value, tuple): return tuple(_default_merge_spec(item) for item in value) if isinstance(value, list): return [_default_merge_spec(item) for item in value] if isinstance(value, dict): return {key: _default_merge_spec(item) for key, item in value.items()} return TensorChunkSpec(0) if hasattr(value, "shape") else None ```