ReferenceFull Library ReferenceNemo AutomodelNemo AutomodelComponentsDistributednemo_automodel.components.distributed.recompute_replay

nemo_automodel.components.distributed.recompute_replay

View as Markdown

Record-and-replay of forward-pass decisions across activation-checkpoint recompute.

Activation checkpointing reruns a block during backward. Some of the block’s work is a deterministic decision that the recompute would reproduce bit for bit — a frozen indexer’s top-k routes, an expert-parallel dispatch layout — yet costs real host or communication time to recompute. A RecomputeReplay channel lets the checkpoint forward record such decisions and the recompute take them back in the same order, falling back to recomputing when the log runs out.

Each checkpointed call gets its own RecomputeReplayRecorder from RecomputeReplay.checkpoint_context_fn, so concurrent microbatches and repeated calls of one block never share a log. The binding is thread-local so pipeline-parallel worker threads stay independent.

This is distinct from nemo_automodel.components.moe.router_replay, which replays MoE top-k across two separate forward passes for RL (R3) through a process-global, per-gate registry.

Module Contents

Classes

NameDescription
RecomputeReplayOne named record/replay channel with its own thread-local binding.
RecomputeReplayRecorderOrdered log of one checkpointed call’s forward decisions, consumed on recompute.
_BindingThread-local (recorder, mode) binding; __init__ runs once per thread.

Data

CheckpointContextFn

T

API

class nemo_automodel.components.distributed.recompute_replay.RecomputeReplay(
name: str
)

Bases: Generic[T]

One named record/replay channel with its own thread-local binding.

Consumers create a module-level channel and consult current where the decision is made: record the fresh decision under "record", take a stored one under "replay". checkpoint_context_fn binds the channel around a torch.utils.checkpoint context factory.

Parameters:

name
str

Channel name used in error messages.

_binding
= _Binding()
nemo_automodel.components.distributed.recompute_replay.RecomputeReplay.checkpoint_context_fn(
on_record_exit: collections.abc.Callable[[RecomputeReplayRecorder[T]], None] | None = None

Wrap a checkpoint context_fn so this channel records in forward and replays in recompute.

Parameters:

context_fn
CheckpointContextFn | None

Existing torch.utils.checkpoint context factory returning the (forward, recompute) context managers, or None for no other contexts.

on_record_exit
Callable[[RecomputeReplayRecorder[T]], None] | NoneDefaults to None

Optional hook run on the recorder after the forward context exits, for work that must stay out of the checkpointed op trace.

Returns: CheckpointContextFn

A context factory that binds a fresh recorder per checkpoint call.

nemo_automodel.components.distributed.recompute_replay.RecomputeReplay.current() -> tuple[nemo_automodel.components.distributed.recompute_replay.RecomputeReplayRecorder[nemo_automodel.components.distributed.recompute_replay.T], str] | None

Return the active (recorder, mode), or None outside a scope.

nemo_automodel.components.distributed.recompute_replay.RecomputeReplay.scope(
mode: str
) -> collections.abc.Iterator[None]

Bind recorder in mode ("record" or "replay") for the enclosed region.

Parameters:

recorder
RecomputeReplayRecorder[T] | None

Log shared by the forward and recompute of one checkpointed call; None disables replay inside the scope.

mode
str

"record" or "replay".

class nemo_automodel.components.distributed.recompute_replay.RecomputeReplayRecorder()

Bases: Generic[T]

Ordered log of one checkpointed call’s forward decisions, consumed on recompute.

_cursor
= 0
_records
list[T] = []
records
tuple[T, ...]

The recorded decisions in forward order.

replay_misses
= 0
nemo_automodel.components.distributed.recompute_replay.RecomputeReplayRecorder.__len__() -> int
nemo_automodel.components.distributed.recompute_replay.RecomputeReplayRecorder.record(
) -> None

Append one forward decision.

nemo_automodel.components.distributed.recompute_replay.RecomputeReplayRecorder.rewind() -> None

Restart replay from the first record.

nemo_automodel.components.distributed.recompute_replay.RecomputeReplayRecorder.take() -> nemo_automodel.components.distributed.recompute_replay.T | None

Return the next recorded decision, or None when replay outruns the log.

class nemo_automodel.components.distributed.recompute_replay._Binding()

Bases: local

Thread-local (recorder, mode) binding; __init__ runs once per thread.

mode
str | None = None
recorder
RecomputeReplayRecorder | None = None
nemo_automodel.components.distributed.recompute_replay.CheckpointContextFn = Callable[[], tuple[AbstractContextManager, AbstractContextManager]]
nemo_automodel.components.distributed.recompute_replay.T = TypeVar('T')