Gradient Checkpointing (Activation Recomputation)
梯度检查点(激活重计算)AdvancedStoring fewer intermediate activations during the forward pass and recomputing them during backward, trading compute for memory.
During training, the intermediate results (activations) produced by the forward pass are normally kept in memory until backpropagation is done with them, which is one of the biggest consumers of GPU memory. Gradient checkpointing saves activations only at a handful of checkpoint locations and discards the rest, recomputing them with a fresh forward pass from the nearest checkpoint whenever backpropagation needs them. Tianqi Chen and colleagues' 2016 paper “Training Deep Nets with Sublinear Memory Cost” systematically introduced this: an n-layer network needs only about O(√n) memory for activations, at the cost of roughly one extra forward pass per small batch; in the paper, a 1,000-layer residual network's memory usage dropped from 48GB to 7GB, with about a 30% increase in running time. PyTorch's torch.utils.checkpoint implements this, and it's commonly turned on together with mixed precision and gradient accumulation when training large models and VLAs.
ExampleFine-tuning a Transformer policy by wrapping each Transformer layer in torch.utils.checkpoint noticeably lowers memory usage, at the cost of each training step running somewhat slower.
- Also called
- Activation Checkpointing
- Related
- Backpropagation · GPU Memory (VRAM) · Gradient Accumulation · Mixed-Precision Training · Fully Sharded Data Parallel (FSDP) · DeepSpeed
- Sources
- Training Deep Nets with Sublinear Memory Cost (arXiv 1604.06174)
torch.utils.checkpoint (PyTorch 文档) (Chinese)