Know What You Need: Efficient Activation Checkpointing
Abstract
Large Language Model training is increasingly constrained by memory. Activation checkpointing reduces the activation memory footprint by discarding selected activations during the forward pass and rematerializing them during backpropagation, trading memory savings for additional FLOPs. Although essential for scaling models, this overhead makes rematerializing linear layers of transformers prohibitively expensive. We introduce a novel method for approximate rematerialization that uses information from the initial forward pass to selectively recompute only a fraction of the original layer components. These components are either chosen greedily, or sketched, yielding unbiased gradient estimates. Through ablation studies, we show that our method closely matches the original gradient. When applied to FFN layers, attention layers, or both, our method reduces the computational cost for the rematerialization and backward pass of linear layers by 33\% on a Llama-488M, with perplexity increases of only 1.7\%, 1.2\%, and 2.5\%.