> For clean Markdown of any page, append .md to the page URL.
> For a complete documentation index, see https://docs.nvidia.com/nemo/automodel/llms.txt.
> For AI client integration (Claude Code, Cursor, etc.), connect to the MCP server at https://docs.nvidia.com/nemo/automodel/_mcp/server.

# 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

| Name                                                                                                         | Description                                                                      |
| ------------------------------------------------------------------------------------------------------------ | -------------------------------------------------------------------------------- |
| [`RecomputeReplay`](#nemo_automodel-components-distributed-recompute_replay-RecomputeReplay)                 | One named record/replay channel with its own thread-local binding.               |
| [`RecomputeReplayRecorder`](#nemo_automodel-components-distributed-recompute_replay-RecomputeReplayRecorder) | Ordered log of one checkpointed call's forward decisions, consumed on recompute. |
| [`_Binding`](#nemo_automodel-components-distributed-recompute_replay-_Binding)                               | Thread-local `(recorder, mode)` binding; `__init__` runs once per thread.        |

### Data

[`CheckpointContextFn`](#nemo_automodel-components-distributed-recompute_replay-CheckpointContextFn)

[`T`](#nemo_automodel-components-distributed-recompute_replay-T)

### API

```python
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()`

---

```python
nemo_automodel.components.distributed.recompute_replay.RecomputeReplay.checkpoint_context_fn(
    context_fn: nemo_automodel.components.distributed.recompute_replay.CheckpointContextFn | None,
    on_record_exit: collections.abc.Callable[[RecomputeReplayRecorder[T]], None] | None = None
) -> nemo_automodel.components.distributed.recompute_replay.CheckpointContextFn
```

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] | None` — default: 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.

```python
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.

```python
nemo_automodel.components.distributed.recompute_replay.RecomputeReplay.scope(
    recorder: nemo_automodel.components.distributed.recompute_replay.RecomputeReplayRecorder[nemo_automodel.components.distributed.recompute_replay.T] | None,
    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"`.

---

```python
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`

---

```python
nemo_automodel.components.distributed.recompute_replay.RecomputeReplayRecorder.__len__() -> int
```

```python
nemo_automodel.components.distributed.recompute_replay.RecomputeReplayRecorder.record(
    entry: nemo_automodel.components.distributed.recompute_replay.T
) -> None
```

Append one forward decision.

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

Restart replay from the first record.

```python
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.

```python
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`

---

```python
nemo_automodel.components.distributed.recompute_replay.CheckpointContextFn = Callable[[], tuple[AbstractContextManager, AbstractContextManager]]
```

```python
nemo_automodel.components.distributed.recompute_replay.T = TypeVar('T')
```