Gradient checkpointing
Gradient checkpointing stores fewer activations by recomputing them when needed, exchanging extra computation for lower memory use.
Reverse-mode differentiation saves intermediate values from the forward pass for use in backward. For large models, that residual memory is often the main memory problem.
Checkpointing keeps a function’s inputs but not its internal tensors. During backward, the function is invoked again and the missing intermediates are recomputed for gradient calculation.
This is a memory-versus-compute trade-off. Selective activation checkpointing can save chosen operations while recomputing others, which gives finer control than rematerializing every operation in a region.
Backpropagation needs intermediate activations, and keeping all of them can dominate training memory.
Keep boundaries, then rebuild the middle.
- 1 · forwardRun the selected function during the ordinary forward pass.
- 2 · saveKeep its inputs while not retaining the internal tensors required only for backward.
- 3 · recomputeInvoke the function again during backward to rebuild those intermediates.
- 4 · differentiateUse the rebuilt values to compute gradients and continue backpropagation.
Checkpoint placement chooses a point on the memory-versus-compute trade-off.
| Who | What they ask | What it works with |
|---|---|---|
| Large-model trainer | “Can this network fit without reducing its depth?” | Peak activation memory |
| Fine-tuning engineer | “Can a larger batch fit on the same accelerator?” | Memory saved and recomputation time |
| Compiler engineer | “Which intermediates are worth saving?” | Operation cost and rematerialization policy |
- It reduces the number of forward tensors retained for backward.
- It can let deeper models fit within a memory limit.
- Selective policies can avoid recomputing selected operations while recomputing others.
- The original method focuses on memory for intermediate feature maps and gradients.
- It adds computation because discarded intermediates must be produced again.
- A backward invocation that differs from the forward invocation can produce errors or incorrect gradients.
Sources used
This explainer is written in original language. The links below support its factual claims.
- paperTraining Deep Nets with Sublinear Memory Cost, Chen et al. · read 28 Sept 2026
- docstorch.utils.checkpoint, PyTorch · read 28 Sept 2026
- docsGradient checkpointing with jax.checkpoint, JAX · read 28 Sept 2026
- docstf.recompute_grad, TensorFlow · read 28 Sept 2026
- officialCurrent and New Activation Checkpointing Techniques in PyTorch, PyTorch Foundation · read 28 Sept 2026