nemo_automodel.components.distributed.recompute_replay
nemo_automodel.components.distributed.recompute_replay
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
Data
API
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:
Channel name used in error messages.
Wrap a checkpoint context_fn so this channel records in forward and replays in recompute.
Parameters:
Existing torch.utils.checkpoint context factory returning the
(forward, recompute) context managers, or None for no other contexts.
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.
Return the active (recorder, mode), or None outside a scope.
Bind recorder in mode ("record" or "replay") for the enclosed region.
Parameters:
Log shared by the forward and recompute of one checkpointed call;
None disables replay inside the scope.
"record" or "replay".
Bases: Generic[T]
Ordered log of one checkpointed call’s forward decisions, consumed on recompute.
The recorded decisions in forward order.
Append one forward decision.
Restart replay from the first record.
Return the next recorded decision, or None when replay outruns the log.
Bases: local
Thread-local (recorder, mode) binding; __init__ runs once per thread.