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#
One layer’s immutable view of a shared residual-stream replay block. |
Functions#
Return whether selective residual-stream replay is active for this forward. |
|
Partition local physical layers into independent ordered replay blocks. |
|
Checkpoint a residual read while retaining its carried stream as state. |
|
Checkpoint one residual write while preserving standard module hooks. |
Data#
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,
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,
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]] = (),
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,
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,
Checkpoint one residual write while preserving standard module hooks.