core.transformer.residual_recompute#

Ordered output-discard replay for residual-stream operations.

Wide-residual maps are static model parameters, so replay only needs the residual-stream inputs and branch outputs already present in the autograd graph. Within each replay block, cheap reads, connected norms, and non-terminal writes are reconstructed in forward order during backward.

Module Contents#

Classes#

ResidualStreamRecomputeContext

One layer’s immutable view of a shared residual-stream replay block.

Functions#

residual_stream_recompute_enabled

Return whether selective residual-stream replay is active for this forward.

build_residual_stream_recompute_plan

Partition local physical layers into independent ordered replay blocks.

checkpoint_residual_read

Checkpoint a residual read while retaining its carried stream as state.

checkpoint_residual_write

Checkpoint one residual write while preserving standard module hooks.

Data#

_R

API#

core.transformer.residual_recompute._R#

‘TypeVar(…)’

class core.transformer.residual_recompute.ResidualStreamRecomputeContext#

One layer’s immutable view of a shared residual-stream replay block.

manager: megatron.core.tensor_parallel.random.CheckpointWithoutOutputManager#

None

is_block_end: bool#

None

checkpoint(
function: collections.abc.Callable[..., core.transformer.residual_recompute._R],
*args: Any,
fp8: bool = False,
) → core.transformer.residual_recompute._R#

Run one cheap operation and register it for ordered replay.

finalize(hidden_states: torch.Tensor) → None#

Discard this block’s registered outputs once its live boundary exists.

core.transformer.residual_recompute.residual_stream_recompute_enabled(
config: megatron.core.transformer.transformer_config.TransformerConfig,
training: bool,
) → bool#

Return whether selective residual-stream replay is active for this forward.

core.transformer.residual_recompute.build_residual_stream_recompute_plan(
num_layers: int,
block_size: int | None,
*,
atomic_layer_pairs: collections.abc.Sequence[tuple[int, int]] = (),
) → list[core.transformer.residual_recompute.ResidualStreamRecomputeContext]#

Partition local physical layers into independent ordered replay blocks.

Shortcut MoE’s predecessor (e.g. attention or Mamba) and paired MoE layer share replay-owned intermediates, so a replay block must not end between them.

Parameters:
  • num_layers – Local physical layer count, before shortcut wrapper grouping.

  • block_size – Target maximum physical layers per block; None uses the entire local stack. A pair stays intact even when block_size=1, forming a two-layer block.

  • atomic_layer_pairs – Disjoint adjacent pairs of zero-based local physical layer indices that must share one replay manager.

Returns:

One context per physical layer, sharing a manager within each block. Only the block’s final layer has is_block_end=True.

For example, four layers with block_size=3 and pairs (0, 1), (2, 3) form blocks [0, 1] and [2, 3], not [0, 1, 2] and [3].

core.transformer.residual_recompute.checkpoint_residual_read(
connection: megatron.core.transformer.residual_connection.ResidualConnection,
hidden_states: torch.Tensor,
context: core.transformer.residual_recompute.ResidualStreamRecomputeContext,
*,
fp32_residual_connection: bool,
branch_input_dtype: torch.dtype | None = None,
) → tuple[torch.Tensor, megatron.core.transformer.residual_connection.ResidualConnectionState]#

Checkpoint a residual read while retaining its carried stream as state.

An optional branch-input dtype request is captured by the replay closure so eager execution and backward reconstruction produce the same branch dtype.

core.transformer.residual_recompute.checkpoint_residual_write(
connection: megatron.core.transformer.residual_connection.ResidualConnection,
branch_output: megatron.core.transformer.residual_connection.ResidualBranchOutput,
state: megatron.core.transformer.residual_connection.ResidualConnectionState,
context: core.transformer.residual_recompute.ResidualStreamRecomputeContext,
*,
dropout_probability: float,
training: bool,
) → torch.Tensor#

Checkpoint one residual write while preserving standard module hooks.