Concepts

Gradient checkpointing

4 min readintermediateUpdated 28 Sept 2026
1 · In one line

Gradient checkpointing stores fewer activations by recomputing them when needed, exchanging extra computation for lower memory use.

1 · What it is

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.

2 · Why it exists

Backpropagation needs intermediate activations, and keeping all of them can dominate training memory.

Activation memoryAutodiff normally retains forward intermediates until their gradients are computed.
Model sizeMemory limits can prevent deeper models from fitting.
Trade-offRecomputation lowers saved activation memory but adds forward work during backward.
3 · How it works

Keep boundaries, then rebuild the middle.

  1. 1 · forwardRun the selected function during the ordinary forward pass.
  2. 2 · saveKeep its inputs while not retaining the internal tensors required only for backward.
  3. 3 · recomputeInvoke the function again during backward to rebuild those intermediates.
  4. 4 · differentiateUse the rebuilt values to compute gradients and continue backpropagation.

Checkpoint placement chooses a point on the memory-versus-compute trade-off.

4 · Where it's used
WhoWhat they askWhat 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
5 · What it solves, and what it doesn't
solves
  • 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.
doesn't solve
  • 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.
6 · Go deeper

Sources used

This explainer is written in original language. The links below support its factual claims.

  1. paperTraining Deep Nets with Sublinear Memory Cost, Chen et al. · read 28 Sept 2026
  2. docstorch.utils.checkpoint, PyTorch · read 28 Sept 2026
  3. docsGradient checkpointing with jax.checkpoint, JAX · read 28 Sept 2026
  4. docstf.recompute_grad, TensorFlow · read 28 Sept 2026
  5. officialCurrent and New Activation Checkpointing Techniques in PyTorch, PyTorch Foundation · read 28 Sept 2026