Activation checkpointing saves memory by rerunning a forward segment during backward; it is unrelated to saving training state on disk.
Activation checkpointing, recomputation and RNG parity
Separate two meanings of checkpoint
A durable training checkpoint stores model and optimizer state so a process can resume after interruption. Activation checkpointing deliberately avoids retaining selected intermediate tensors from a forward segment, then recomputes them during backward. It trades extra compute for lower saved-activation memory and cannot restore a crashed run. The resume lesson covers the disk artifact. Use distinct names in code and metrics so an operator does not mistake a memory optimization for recovery protection.
Choose a pure recomputable segment
The wrapped function must produce the same needed values when backward reruns it. Avoid untracked global counters, file writes, mutating batch-normalization state inside a recomputed region, or branches that depend on external state changed between calls. Store input tensors that define the segment; select a block large enough to save meaningful activations but not so large that recomputation dominates latency. The example wraps a simple feed-forward block with dropout and compares gradients against a plain forward under the same seed.
Preserve stochastic behavior deliberately
Dropout draws random masks. A checkpointed backward needs the same stochastic choices as the original forward for gradient parity, so use an RNG-preservation setting when that parity is required. Saving and restoring RNG state has overhead and can have device-specific limits; do not assume stochastic operators on arbitrary extra devices are automatically handled. Make the checkpoint implementation mode explicit and test fixed-input output and parameter gradients. Precision settings must also match in the comparison.
Measure saved memory and added time
On a CUDA device, reset peak statistics immediately before a representative train step and synchronize around timing. Record peak tensor-allocated and allocator-reserved bytes separately. The allocator may reserve blocks beyond live tensors, so a lower allocated peak may not produce a proportional drop in process-visible memory. Run several warm steps with the same batch shape and optimizer state already initialized. The profiler lesson distinguishes these measurements.
Check the release claim
Compare plain and checkpointed gradients and one optimizer step within a stated numerical tolerance, then evaluate task metrics after equal accepted-example budgets. If memory drops but update time rises enough to miss a training deadline, report that cost. If the larger batch made possible by checkpointing changes the effective optimization, it is a second experiment, not a clean parity test. The project measures both memory and model behavior.
Implementation
import torch
from torch import nn
from torch.utils.checkpoint import checkpoint
torch.manual_seed(47)
receipt_block = nn.Sequential(nn.Linear(12, 24), nn.ReLU(),
nn.Dropout(0.2), nn.Linear(24, 12))
base_inputs = torch.randn(3, 12)
plain_inputs = base_inputs.clone().requires_grad_(True)
torch.manual_seed(71)
plain_loss = receipt_block(plain_inputs).square().mean()
plain_loss.backward()
plain_gradients = [parameter.grad.clone() for parameter in receipt_block.parameters()]
receipt_block.zero_grad(set_to_none=True)
recomputed_inputs = base_inputs.clone().requires_grad_(True)
torch.manual_seed(71)
recomputed_loss = checkpoint(receipt_block, recomputed_inputs,
use_reentrant=False, preserve_rng_state=True).square().mean()
recomputed_loss.backward()
torch.testing.assert_close(plain_loss, recomputed_loss)
torch.testing.assert_close(plain_inputs.grad, recomputed_inputs.grad)
for expected_gradient, parameter in zip(plain_gradients, receipt_block.parameters()):
torch.testing.assert_close(expected_gradient, parameter.grad)Performance and operating cost
Let A be the saved activations of a wrapped block and F its forward cost. Checkpointing can avoid retaining much of A, but typically adds up to another F of recomputation during backward; exact savings depend on segment boundaries and implementation. Model parameters, optimizer moments and input tensors remain. RNG preservation adds bookkeeping. The demonstration checks parity on a small block but cannot quantify memory benefit on a larger network; measure peak allocated bytes and update time on the actual accelerator.
Common Mistakes
- Do not describe activation recomputation as a recoverable disk checkpoint.
- Do not checkpoint a function with untracked side effects or changing global state.
- Do not disable RNG preservation for a stochastic block and still expect exact gradient parity.
