Activation Checkpointing: Gradient Memory

Michael BrenndoerferJanuary 22, 202656 min read

Part of Language AI Handbook

How activation checkpointing trades compute for memory by discarding and recomputing activations.

Choose your expertise level to adjust how many terms are explained. Beginners see more tooltips, experts see fewer to maintain reading flow. Hover over underlined terms for instant definitions.

Article links

Make inline references clickable

Activation Checkpointing

Training large language models is fundamentally a memory-constrained problem. A single forward pass through a transformer stores activations for every layer so that backpropagation can compute gradients. These stored tensors are not small: for a 7B-parameter model with a batch size of 4 and a sequence length of 2048, the activations alone can occupy over 30 GB of GPU memory, often more than the model weights themselves. When memory runs out, you cannot increase batch size, cannot train larger models, and cannot fit more context without buying more hardware.

Activation checkpointing (also called gradient checkpointing or rematerialization) solves this problem by trading computation for memory. Instead of storing every activation produced during the forward pass, the technique stores only a subset of them at strategic points in the network, called "checkpoints." During the backward pass, when gradients need to flow through a layer whose activations were discarded, those activations are recomputed on the fly from the nearest checkpoint. You pay a one-time recomputation cost but avoid storing the full activation buffer. For most architectures, this exchange is highly favorable: recomputing one forward pass per checkpoint adds roughly 33% more compute, but memory usage drops by a factor proportional to the number of layers between checkpoints.

The idea predates modern deep learning, tracing back to reverse-mode automatic differentiation research in the 1980s. In their foundational work on adjoint methods, researchers noticed that storing all intermediate values in a long chain of computations creates prohibitive memory requirements. The "checkpointing" strategy they developed for scientific computing is essentially the same concept: save only selected states, recompute the rest on demand. When deep learning rediscovered this idea for neural networks in the mid-2010s, tools like the Theano-based memory-efficient backpropagation library and later Chen et al.'s 2016 paper "Training Deep Nets with Sublinear Memory Cost" gave it renewed attention. Checkmate and similar analysis frameworks later formalized the optimal checkpoint placement problem as a graph-theoretic optimization. Today, frameworks like PyTorch expose activation checkpointing as a single function call via torch.utils.checkpoint.checkpoint, making it accessible without manual implementation.

This chapter explains how activation checkpointing works, why the memory-compute tradeoff is so favorable, how to choose which layers to checkpoint, and how to implement selective checkpointing to minimize overhead while maximizing memory savings.

The Memory Problem Activation Checkpointing Solves

To understand why activation checkpointing is necessary, you first need a precise picture of where memory goes during training. GPU memory is a scarce, non-expandable resource. Unlike CPU RAM, you cannot simply add more of it mid-run. When an allocation fails, the entire training process crashes, wasting potentially hours of compute time accumulated to that point.

Activation Memory During the Forward Pass

As we examined in the Memory Management chapter, total GPU memory during training decomposes into four categories: model parameters, optimizer states, gradients, and activations. For large models with large batch sizes or long sequences, activations dominate.

The reason activations are so expensive comes down to the structure of the computation graph. Neural network training uses automatic differentiation, which traces through every operation performed during the forward pass and records it in a graph structure. Later, during backpropagation, the system walks this graph in reverse to compute gradients. For the gradient computation at each node to work, the system needs to know the values that existed at that node during the forward pass. So every intermediate tensor produced during the forward pass gets retained in memory until its gradient has been computed during the backward pass. For a deep network with many layers, "until the backward pass computes its gradient" effectively means "until the very end of the backward pass," because gradients flow from the output layer back to the first layer.

Consider a single transformer layer with hidden dimension dd, sequence length TT, and batch size BB. The inputs and outputs of each sublayer (attention, feed-forward, layer norm) all need to be stored for the backward pass. The activations for a single transformer layer scale approximately as:

Memorylayer≈12⋅B⋅T⋅d⋅bytes\text{Memory}_{\text{layer}} \approx 12 \cdot B \cdot T \cdot d \cdot \text{bytes}

where:

  • BB is the batch size (number of sequences processed in parallel)
  • TT is the sequence length (tokens per sequence)
  • dd is the model hidden dimension
  • the factor of 12 accounts for storing query, key, and value matrices, attention weight matrices, intermediate feed-forward activations, and residual stream tensors, each contributing a B×T×dB \times T \times d or B×T×4dB \times T \times 4d tensor

This factor of 12 deserves more explanation because it surprises many practitioners who expect activations to be small. Within the attention sublayer alone, you store: the pre-norm input (one [B,T,d][B, T, d] tensor), the query projection QQ ([B,T,d][B, T, d]), the key projection KK ([B,T,d][B, T, d]), the value projection VV ([B,T,d][B, T, d]), and the attention weight matrix A=softmax(QKT/dk)A = \text{softmax}(QK^T / \sqrt{d_k}) ([B,H,T,T][B, H, T, T]). The attention weights alone contribute a tensor proportional to T2T^2, which is why long-context training is so memory-intensive. Within the feed-forward sublayer, the intermediate activation has shape [B,T,4d][B, T, 4d], four times the hidden dimension width.

For a 32-layer model with d=4096d = 4096, T=2048T = 2048, B=4B = 4, and using float16 (2 bytes per value), total activation memory across all layers is:

Memorytotal=32×12×B×T×d×2 bytes=32×12×4×2048×4096×2≈64 GB\begin{aligned} \text{Memory}_{\text{total}} &= 32 \times 12 \times B \times T \times d \times 2 \text{ bytes} \\ &= 32 \times 12 \times 4 \times 2048 \times 4096 \times 2 \\ &\approx 64 \text{ GB} \end{aligned}

That is 64 GB just for activations, on a GPU that might have 40-80 GB total. The parameters themselves (at 16-bit precision) occupy roughly 14 GB for a 7B-parameter model, and the Adam optimizer states (maintained at 32-bit precision for numerical stability) add another 56 GB. You can see immediately why activation memory is not a footnote but the primary driver of training memory constraints. Without some intervention, a 7B model trained with Adam simply cannot fit in a single 80 GB A100.

Why Backpropagation Needs Activations

Backpropagation requires activations because gradient computation is not invertible in general. To compute gradients with respect to the inputs of a layer, you must know what values those inputs had during the forward pass.

The chain rule makes this explicit. The gradient flowing back through a layer depends on the activation at that layer's input or output. For a nonlinear activation function ff, the gradient ∂L/∂x\partial \mathcal{L} / \partial x at some position xx requires knowing the value of xx itself:

∂L∂x=∂L∂f(x)⋅f′(x)\frac{\partial \mathcal{L}}{\partial x} = \frac{\partial \mathcal{L}}{\partial f(x)} \cdot f'(x)

where:

  • L\mathcal{L} is the training loss
  • f′(x)f'(x) is the derivative of the activation function, which depends on xx directly, not just on f(x)f(x)

For the GELU activation function specifically, f′(x)f'(x) has no simple closed form in terms of f(x)f(x) alone, so xx must be retained separately from f(x)f(x). You cannot reconstruct xx from f(x)f(x) without inverting GELU, which is expensive and numerically unstable.

The same dependency appears in every major component of a transformer. The attention mechanism requires storing the attention weight matrix A=softmax(QKT/dk)A = \text{softmax}(QK^T / \sqrt{d_k}) to compute gradient updates for the query, key, and value projection matrices WQW_Q, WKW_K, and WVW_V. The gradient of the softmax output with respect to its input is a Jacobian matrix, and computing it requires knowing the softmax output values themselves. The layer normalization sublayer must store its input mean and variance to correctly backpropagate through the normalization step.

Every nonlinear operation creates an activation dependency in the computation graph that must be resolved before the backward pass can continue. Without some form of checkpointing, the entire forward pass activation buffer must remain in memory until the corresponding backward pass gradient is computed, meaning activations from layer 1 must survive in GPU memory until after layer 32's backward pass completes.

The Compute-Memory Ledger for Training

It is helpful to think of training as maintaining a ledger of resources. On one side you have GPU memory capacity. On the other side you have compute budget in terms of FLOPs per second. The standard forward-backward cycle uses memory inefficiently: it reserves activation memory during the forward pass, holds it idle while the backward pass works layer by layer from the end, and only releases memory as each layer's gradients are computed.

This idle reservation is the waste that activation checkpointing targets. The forward pass activations from early layers sit in memory for the entire duration of the backward pass through later layers. If you have 32 layers, layer 1's activations are reserved while layers 32, 31, 30... 3, 2 all compute their gradients, finally being used only when the backward pass reaches layer 1. For a 32-layer model, early-layer activations are idle for roughly 97% of the backward pass. Discarding them and recomputing on demand turns idle memory into compute work, which is a much more efficient use of the available hardware budget.

The Checkpointing Mechanism

Activation checkpointing redefines which activations must survive the full forward-to-backward journey. Understanding the mechanism requires tracing through both the forward and backward passes in detail.

Basic Checkpointing Algorithm

The algorithm proceeds in three phases:

Segmented forward pass. The network is divided into segments, typically one segment per transformer layer or one per block of layers. At the boundary between segments, the activation tensor is saved to memory. Within each segment, intermediate activations are computed normally but immediately discarded after use, rather than being held for the backward pass.

Standard backward pass (outer). Backpropagation proceeds from the output layer toward the input. When it reaches a segment boundary, it has the activation at that boundary available (because it was checkpointed), and it triggers a recomputation pass for that segment.

Local recomputation. For each segment being backpropagated through, a mini-forward pass is re-run from that segment's saved input activation to its output. This regenerates all the intermediate activations within the segment. Backpropagation through the segment can then proceed normally using these regenerated activations, after which they are discarded again.

The memory at any point in backward is bounded by the number of checkpoint boundaries plus the memory for one segment's activations (the one currently being recomputed). If you have LL layers total and checkpoint every kk layers, you store L/kL/k checkpoint boundary tensors and hold at most kk layers of intermediate activations in memory at once during recomputation.

The behavior is like a rolling window: you always have the boundary tensors pinned, and the intermediate activations within one segment are temporarily materialized, then released as you move to the next segment. The maximum memory footprint is not LL times a layer's activation memory, but rather L/k+kL/k + k times a layer's activation memory.

Checkpoint Boundary

A checkpoint boundary is a point in the computation graph where an activation tensor is explicitly saved to memory for use during the backward pass. Activations computed between two boundaries are discarded after the forward pass and recomputed during backpropagation.

Memory Analysis: The Tradeoff in Numbers

Let's work through the memory savings concretely. With LL layers total, storing all activations requires memory proportional to LL. Without checkpointing, this scales as O(L)O(L) because every layer's activations coexist in memory simultaneously during the backward pass.

With checkpointing at every kk layers, the memory for activations is:

Memorycheckpointed=Lk⋅Mckpt+k⋅Mlayer\text{Memory}_{\text{checkpointed}} = \frac{L}{k} \cdot M_{\text{ckpt}} + k \cdot M_{\text{layer}}

where:

  • L/kL/k is the number of checkpoint boundaries stored simultaneously
  • MckptM_{\text{ckpt}} is the memory for one checkpoint boundary tensor (the residual stream at a layer boundary)
  • kk is the number of layers between consecutive checkpoints
  • MlayerM_{\text{layer}} is the activation memory for one full layer (the cost paid during recomputation)

The first term represents the checkpoint storage cost (boundaries stored during the entire forward pass), and the second term represents the recomputation cost (intermediate activations held during one backward segment pass).

To find the optimal kk that minimizes total memory, take the derivative with respect to kk and set it to zero:

ddk[L⋅Mckptk+k⋅Mlayer]=−L⋅Mckptk2+Mlayer=0\frac{d}{dk}\left[\frac{L \cdot M_{\text{ckpt}}}{k} + k \cdot M_{\text{layer}}\right] = -\frac{L \cdot M_{\text{ckpt}}}{k^2} + M_{\text{layer}} = 0

Solving for kk:

k∗=L⋅MckptMlayerk^* = \sqrt{\frac{L \cdot M_{\text{ckpt}}}{M_{\text{layer}}}}

When Mckpt≈MlayerM_{\text{ckpt}} \approx M_{\text{layer}} (checkpoint tensors are similar in size to per-layer activations), the optimal is k∗=Lk^* = \sqrt{L}, giving total memory proportional to L\sqrt{L} rather than LL. For 32 layers and equal-sized tensors, the reduction is about 2.8x; the improvement can be larger when checkpoint boundaries are smaller than full-layer activations. The formula says something intuitive: you want to make the two terms in the memory expression roughly equal in magnitude, balancing the cost of storing boundaries against the cost of holding one segment in memory during recomputation.

However, in practice the Mckpt≈MlayerM_{\text{ckpt}} \approx M_{\text{layer}} assumption is often violated. A checkpoint boundary tensor is just the residual stream at one layer's output, with shape [B,T,d][B, T, d]. The full per-layer activation memory includes the attention weight matrix ([B,H,T,T][B, H, T, T]), which can be many times larger. This makes Mckpt≪MlayerM_{\text{ckpt}} \ll M_{\text{layer}}, shifting the optimal kk to smaller values and making denser checkpointing more attractive.

The visualization below shows how activation memory scales with model depth under different checkpointing strategies. The baseline and per-layer checkpointing curves both grow linearly, but per-layer checkpointing has a much smaller slope because it stores only boundary tensors. Spacing checkpoints according to the analytical optimum grows as O(L)O\left(\sqrt{L}\right) and produces the lowest activation-memory curve under the stated assumptions.

Out[3]:
Visualization
Line chart comparing memory scaling under three checkpointing strategies across 4 to 64 model layers.
Activation memory scaling with model depth under three checkpointing strategies. The baseline and per-layer checkpointing curves both grow linearly, although checkpointing reduces the slope. Optimal checkpoint spacing grows with the square root of depth and gives the lowest memory curve under the stated activation-size assumptions. All curves are normalized relative to a single-layer baseline.

Computational Overhead

The recomputation adds a second forward pass per segment during backpropagation. In a standard training step, the relative compute cost is:

Costcheckpointed=Cfwd+Crecompute+Cbwd\text{Cost}_{\text{checkpointed}} = C_{\text{fwd}} + C_{\text{recompute}} + C_{\text{bwd}}

where:

  • CfwdC_{\text{fwd}} is the cost of the forward pass
  • CrecomputeC_{\text{recompute}} is the cost of the recomputed forward pass during backward (equal to CfwdC_{\text{fwd}} for full per-layer checkpointing)
  • CbwdC_{\text{bwd}} is the cost of the backward pass, approximately 2×Cfwd2 \times C_{\text{fwd}} (it must compute gradients for both weights and inputs at each layer)

For full per-layer checkpointing:

Costbaseline=Cfwd+2Cfwd=3CfwdCostcheckpointed=Cfwd+Cfwd+2Cfwd=4Cfwd\begin{aligned} \text{Cost}_{\text{baseline}} &= C_{\text{fwd}} + 2 C_{\text{fwd}} = 3 C_{\text{fwd}} \\ \text{Cost}_{\text{checkpointed}} &= C_{\text{fwd}} + C_{\text{fwd}} + 2 C_{\text{fwd}} = 4 C_{\text{fwd}} \end{aligned}

The overhead ratio is 4/3≈1.334/3 \approx 1.33, confirming the often-cited "33% more compute" figure. This trade reduces activation memory by 10-20x for 33% more compute, which in practice means you can double or triple the batch size, fit a much larger model on the same hardware, or increase the sequence length substantially.

Why is the overhead only 33% and not 100%? Because the backward pass already costs twice the forward pass (to compute weight gradients and input gradients). The recomputation adds one forward pass on top of a baseline cost that already included two forward passes worth of work. If the backward pass were free, the overhead would be 100%. The expensive backward pass is what makes recomputation relatively cheap by comparison.

Selective checkpointing, covered below, reduces the 33% overhead by only recomputing the most expensive layers and leaving cheap activations in memory.

A Concrete Worked Example

Let's trace through exactly what happens during a two-layer forward and backward pass with checkpointing at the layer boundary. This makes the memory states concrete.

Setup: Two transformer layers L1L_1 and L2L_2 with a checkpoint at their boundary. The input to L1L_1 is x0x_0, the output of L1L_1 (and input to L2L_2) is x1x_1, and the output of L2L_2 is x2x_2.

Forward pass, without checkpointing:

  1. Compute x1=L1(x0)x_1 = L_1(x_0). Store all of L1L_1's internal activations (attention weights, FFN intermediate, layer norm statistics).
  2. Compute x2=L2(x1)x_2 = L_2(x_1). Store all of L2L_2's internal activations.
  3. Compute loss L(x2)\mathcal{L}(x_2).

At this point, both L1L_1's and L2L_2's internal activations are in memory simultaneously.

Backward pass, without checkpointing:

  1. Compute ∂L/∂x2\partial \mathcal{L} / \partial x_2 from the loss.
  2. Use L2L_2's stored activations to compute ∂L/∂x1\partial \mathcal{L} / \partial x_1 and the weight gradients for L2L_2's parameters. Discard L2L_2's activations.
  3. Use L1L_1's stored activations to compute ∂L/∂x0\partial \mathcal{L} / \partial x_0 and weight gradients for L1L_1's parameters. Discard L1L_1's activations.

Peak memory: Memory(x0)+Memory(L1 internals)+Memory(L2 internals)+Memory(x1)+Memory(x2)\text{Memory}(x_0) + \text{Memory}(L_1 \text{ internals}) + \text{Memory}(L_2 \text{ internals}) + \text{Memory}(x_1) + \text{Memory}(x_2).

Forward pass, with checkpointing at x1x_1:

  1. Compute x1=L1(x0)x_1 = L_1(x_0). Save x1x_1 as a checkpoint. Discard L1L_1's internal activations.
  2. Compute x2=L2(x1)x_2 = L_2(x_1). Store L2L_2's internal activations.
  3. Compute loss L(x2)\mathcal{L}(x_2).

At this point, only x1x_1 (the checkpoint) and L2L_2's activations are in memory. L1L_1's internals have been discarded.

Backward pass, with checkpointing:

  1. Compute ∂L/∂x2\partial \mathcal{L} / \partial x_2 from the loss.
  2. Use L2L_2's stored activations to compute ∂L/∂x1\partial \mathcal{L} / \partial x_1 and weight gradients for L2L_2's parameters. Discard L2L_2's activations.
  3. Re-run L1L_1's forward pass starting from the saved x0x_0 to reconstruct L1L_1's internal activations. Discard the checkpoint x1x_1.
  4. Use the recomputed L1L_1 activations to compute ∂L/∂x0\partial \mathcal{L} / \partial x_0 and weight gradients for L1L_1. Discard L1L_1's activations.

Peak memory: Memory(x0)+Memory(x1)+Memory(L2 internals)\text{Memory}(x_0) + \text{Memory}(x_1) + \text{Memory}(L_2 \text{ internals}) during step 5, and Memory(x0)+Memory(L1 internals)\text{Memory}(x_0) + \text{Memory}(L_1 \text{ internals}) during steps 6-7. L1L_1 and L2L_2 activations never coexist.

This is the fundamental gain: layers 1 and 2 no longer need to coexist in memory. For 32 layers, this prevents all 32 layers' activations from coexisting. Instead, you hold at most one layer's worth of internal activations at a time during recomputation.

Checkpoint Selection

Which activations to checkpoint is not arbitrary. The goal is to minimize total memory while minimizing recomputation cost. Several factors determine whether a layer is a good checkpoint candidate.

Checkpoint Granularity

The natural granularity for checkpointing in transformer models is the transformer layer. Each transformer block is a closed computation unit: it takes one residual stream tensor as input and produces one residual stream tensor as output. Checkpointing at layer boundaries stores only these residual tensors, which are small relative to the internal activations of the layer.

For a model with hidden dimension d=4096d = 4096, batch size B=4B = 4, and sequence length T=2048T = 2048, the residual stream tensor at each layer boundary has shape [B,T,d]=[4,2048,4096][B, T, d] = [4, 2048, 4096]. At float16 (2 bytes), this is:

4×2048×4096×2≈67 MB4 \times 2048 \times 4096 \times 2 \approx 67 \text{ MB}

Storing 32 such tensors for a 32-layer model requires only about 2.1 GB, compared to the 64 GB needed for all intermediate activations: a 30x reduction for the checkpoint storage cost alone.

The reason the layer boundary tensor is so much smaller than the full layer activations is that it does not include the attention weight matrix. For the same configuration with 32 attention heads and sequence length 2048, the attention weight matrix has shape [B,H,T,T]=[4,32,2048,2048][B, H, T, T] = [4, 32, 2048, 2048], occupying:

4×32×20482×2≈1.07 GB per layer4 \times 32 \times 2048^2 \times 2 \approx 1.07 \text{ GB per layer}

Summed across 32 layers, attention weights alone account for about 34 GB. The residual stream checkpoints are tiny by comparison.

Choosing Within a Layer

Within a single transformer layer, several activations are particularly expensive to store:

Attention weight matrices. The attention computation produces a weight matrix with shape [B,H,T,T][B, H, T, T], where HH is the number of attention heads. The memory cost scales quadratically with sequence length:

Memoryattn=B×H×T2×2 bytes\text{Memory}_{\text{attn}} = B \times H \times T^2 \times 2 \text{ bytes}

For B=4B = 4, H=32H = 32, T=4096T = 4096, this is 4×32×40962×2≈44 \times 32 \times 4096^2 \times 2 \approx 4 GB per layer. For long-context models, this term dominates layer activation memory. Because the attention weights can be recomputed from the stored Q and K tensors (one matrix multiply to form QKTQK^T and one softmax), they are excellent candidates for not storing.

Feed-forward intermediate activations. The FFN sublayer expands the hidden dimension by a factor of 4, producing an intermediate tensor with shape [B,T,4d][B, T, 4d]. At 4x the width of the residual stream, this is the largest single activation within the FFN sublayer. But unlike the attention weights, FFN intermediate activations grow only linearly with sequence length, so their relative importance increases for shorter sequences.

Layer norm statistics. Layer normalization stores per-token mean and variance for backpropagation. These have shape [B,T][B, T] (one scalar per token), making them negligible compared to the above.

The cost-benefit of checkpointing within a layer depends on the ratio of activation size to recomputation cost. Feed-forward activations are expensive to store but cheap to recompute (one matrix multiply and GELU). Attention weights are very expensive to store for long sequences and also cheap to recompute from the stored Q and K tensors.

The visualization below shows how attention weight memory grows quadratically with sequence length, while feed-forward memory grows only linearly. For sequences beyond 2K tokens, attention memory completely dominates.

Out[4]:
Visualization
Stacked area chart showing memory contributions from attention weights, FFN, and residuals across sequence lengths from 512 to 8192.
Per-layer activation memory breakdown by component as sequence length increases, for a model with d=4096, H=32 heads, and B=4. Attention weight memory grows quadratically (O(T^2)) and dominates at long sequences, while feed-forward and residual stream memory grow linearly. This explains why FlashAttention and attention checkpointing provide the largest memory savings for long-context models.

Optimal Checkpoint Placement

Finding the globally optimal checkpoint placement is an NP-hard problem in general, requiring solving a graph scheduling problem over the full computation graph. The Checkmate system (Jain et al., 2019) framed this as an integer linear program and showed that optimal placement can differ substantially from naive equal-spacing heuristics for irregular networks. However, for regular transformer architectures where all layers have identical structure, the heuristics work very well in practice.

The commonly used approaches are:

Equal spacing. Checkpoint every kk layers, with kk chosen based on the memory budget and acceptable overhead. Checkpointing every layer (k=1k = 1) is a common, simple default because it minimizes the within-segment activations, although the boundary-plus-segment model above can favor a larger kk when the accumulated boundary tensors matter. Full recomputation adds roughly 33% to the total forward-plus-backward compute cost.

Recomputation cost weighting. Weight checkpoints toward computation-cheap layers. If two adjacent layers have very different compute costs, it is cheaper to recompute the cheaper layer and checkpoint after it.

Memory pressure-aware scheduling. In pipeline-parallel settings (as discussed in the Pipeline Parallelism chapter), different pipeline stages may have different numbers of micro-batches in flight. Stages with more micro-batches need more aggressive checkpointing, and per-stage checkpoint policies can tune the balance independently. This is important for Megatron-LM and similar frameworks that use interleaved pipeline schedules.

Activation-size-proportional checkpointing. Rather than spacing by layer count, space checkpoints by cumulative activation memory. This produces uniform memory usage per segment even when layers vary in their activation footprint, which is useful for models with heterogeneous layer sizes.

In practice, most large-scale training runs use one of two configurations: full per-layer checkpointing (when memory is extremely tight) or selective checkpointing with FlashAttention (when a good balance of memory and compute is needed). The theoretical optimal for irregular networks is rarely computed in production because transformer layers are regular enough that equal-spacing gives near-optimal results.

Selective Checkpointing

Full per-layer checkpointing is maximally conservative: it recomputes everything. Selective checkpointing applies checkpointing only to the parts of the computation that are most memory-expensive, leaving cheap activations in memory. The goal is to recover most of the memory benefit while adding less than the 33% compute overhead of full checkpointing.

The key insight is that not all activations are equally expensive to store, and not all recomputations are equally expensive in terms of compute. By identifying which activations are large but cheap to recompute, you can selectively discard exactly those tensors while keeping smaller or more expensive-to-recompute activations in memory.

Selective by Layer Type

The most common form of selective checkpointing distinguishes between the attention sublayer and the feed-forward sublayer within each transformer block.

Attention sublayers with long sequences produce the expensive [B,H,T,T][B, H, T, T] weight matrices. Recomputing attention is inexpensive because the stored Q, K, and V projections make the recompute path short: one matrix multiply and softmax to regenerate the attention weights. So attention sublayers are high-value checkpoint targets: large tensors to save, cheap to recompute.

Feed-forward sublayers have large intermediate tensors (shape [B,T,4d][B, T, 4d]) but also have short recomputation paths. The input to the FFN is available from the post-attention residual, and the intermediate activation is just one matrix multiply and GELU away. So FFN intermediates are also good candidates for checkpointing.

Layer norm activations (the mean and variance per token) are tiny by comparison and not worth checkpointing. The Q, K, and V projection matrices themselves are moderate in size and are needed for the attention recomputation, so they are typically kept.

A selective policy might checkpoint at the beginning of each attention block and each FFN block, skipping layer norm intermediate states and small projection activations that are cheap to store relative to their recomputation cost. This configuration typically achieves 70-80% of full-checkpointing memory savings at only 10-15% compute overhead rather than 33%.

Selective by Sequence Length

In models that process variable-length inputs, activation memory scales as O(T2)O(T^2) for attention weight matrices but only O(T)O(T) for feed-forward activations. Selective checkpointing can trigger more aggressively for long sequences and relax for short ones, keeping the memory budget balanced without imposing fixed overhead regardless of sequence length.

Concretely, you might implement a threshold: if the sequence length exceeds some value (say, 2048 tokens), enable attention checkpointing; if below, skip it. This is useful for fine-tuning datasets with highly variable sequence lengths, where most short sequences would waste compute overhead if uniformly checkpointed.

FlashAttention as Implicit Checkpointing

FlashAttention (Dao et al., 2022) implements a tiled attention algorithm that never materializes the full [B,H,T,T][B, H, T, T] attention weight matrix in memory. Instead, it tiles the computation into blocks that fit in the GPU's SRAM (fast on-chip memory), computes softmax numerically stably in blocks, and recomputes the weight tiles during the backward pass from the stored query and key tiles. This is conceptually identical to activation checkpointing applied specifically to attention weights, but implemented at the CUDA kernel level for maximum efficiency.

The SRAM-level implementation makes FlashAttention's "checkpointing" much faster than explicit framework-level checkpointing. SRAM access is orders of magnitude faster than HBM (high-bandwidth memory), so recomputing from Q and K tiles that fit in SRAM adds minimal overhead. By contrast, framework-level checkpointing must read the checkpoint tensor from HBM, adding memory bandwidth cost on top of compute cost.

When using FlashAttention, the O(T2)O(T^2) attention memory cost is already handled. Remaining activation memory is O(T)O(T), and further checkpointing provides smaller marginal benefit. In such settings, practitioners often apply checkpointing selectively to only the feed-forward intermediate activations, accepting the smaller memory cost of other activations to reduce compute overhead.

The combined FlashAttention plus selective FFN checkpointing configuration achieves roughly 80% of the memory savings of full per-layer checkpointing at only about 12% compute overhead. This is the configuration used in most modern large-scale training runs.

Quantized Activations

An orthogonal technique, sometimes combined with checkpointing, stores activation tensors at lower precision. If intermediate activations are quantized to int8 before being saved as checkpoints (and dequantized before reuse), checkpoint memory is halved compared to float16. The accuracy impact depends on the layer and gradient magnitude, but for many training scenarios the quantization error in saved activations is small relative to gradient noise. This approach is sometimes called "activation compression."

Quantized activation storage is particularly appealing for the residual stream checkpoint tensors. These large [B,T,d][B, T, d] tensors are saved at layer boundaries and held for the duration of the backward pass. Quantizing them to int8 reduces their storage cost by 2x with minimal impact on gradient quality, since the gradient magnitudes in the residual stream are typically much larger than the quantization error.

The memory-compute tradeoff for different selective strategies can be visualized as a frontier curve. Each point represents a different checkpointing policy, from no checkpointing (maximum memory, minimum compute overhead) to full per-layer checkpointing (minimum memory, maximum compute). Selective strategies occupy the interior of this frontier.

Out[5]:
Visualization
Scatter plot showing memory reduction vs compute overhead for different checkpointing strategies, illustrating the tradeoff frontier.
Illustrative memory-compute tradeoff frontier for activation checkpointing strategies on a 32-layer transformer. Moving right along the frontier saves memory at the cost of more recomputation. Under these estimates, FlashAttention with selective FFN-only checkpointing achieves near-full-checkpointing memory savings at substantially lower compute overhead.

Historical Development and the Path to Modern Checkpointing

Activation checkpointing did not emerge from deep learning research: it arrived from scientific computing. Researchers working on optimal control and parameter estimation in the 1980s faced a structurally identical problem. They needed to compute gradients of an objective function through long sequences of operations, and storing all intermediate states was prohibitively expensive on the hardware of the time. Andreas Griewank's 1992 paper "Achieving logarithmic growth of temporal and spatial complexity in reverse automatic differentiation" established the theoretical foundation, showing that a divide-and-conquer checkpoint placement strategy could reduce the temporal complexity of gradient computation to O(nlog⁡n)O(n \log n) at the cost of O(log⁡n)O(\log n) memory, where nn is the number of time steps. This is the same trade that modern activation checkpointing exploits, applied to the layer dimension rather than the time dimension.

The first direct application to neural network training appeared in the context of recurrent networks. Training long recurrent neural networks with backpropagation through time (BPTT) required storing a hidden state for every time step, which made very long sequences impractical. Researchers applied checkpointing along the time axis: save hidden states every kk steps, recompute intermediate states on demand during the backward pass. This "truncated BPTT with checkpointing" was used in practice well before the transformer era.

Chen et al.'s 2016 paper "Training Deep Nets with Sublinear Memory Cost" brought the idea to feedforward networks and proved the O(n)O(\sqrt{n}) memory bound for equal-spacing checkpointing with k=nk = \sqrt{n}. This paper was influential because it showed that the technique applied to arbitrary feedforward architectures, not just recurrent ones, and provided a clean theoretical analysis of the tradeoff. The paper also demonstrated that memory savings of 10x or more were achievable for networks with tens of layers at less than 33% compute overhead, which matched what practitioners were finding empirically.

PyTorch incorporated torch.utils.checkpoint in version 0.4.0 (2018), and the API has been stable since. The initial implementation used the reentrant approach, which had correctness limitations. The non-reentrant implementation added in later versions fixed these issues and is now the recommended interface. Jax, the other major deep learning framework used for large-scale training, provides jax.checkpoint (also called jax.remat) with similar semantics and has been used extensively in training models like PaLM and Gemini.

The scaling of transformer models through 2020-2023 made activation checkpointing essential rather than optional. Models like GPT-3 (175 billion parameters), Megatron-Turing NLG (530 billion parameters), and the Chinchilla family could not have been trained within practical memory budgets without combining activation checkpointing with tensor and pipeline parallelism. The Megatron-LM codebase, which underlies most large-scale transformer training at NVIDIA and Microsoft, has activation checkpointing enabled by default for any model above a modest size threshold.

The discovery of FlashAttention in 2022 represents the most significant recent advance in the checkpointing ecosystem. By fusing the attention computation into a single CUDA kernel that tiles across SRAM, FlashAttention achieves the memory profile of activation checkpointing for the attention sublayer while running at or above the speed of standard attention on hardware with favorable compute-to-bandwidth ratios. FlashAttention-2 and FlashAttention-3 further extended these gains with improved parallelism and hardware-specific optimizations. The practical effect was to make long-context training tractable: 128K-token context windows that would have required petabytes of activation memory with naive attention are now trainable on standard GPU clusters.

The next frontier is automatic checkpoint placement. Current frameworks require the user to specify where to checkpoint, either at the per-layer granularity or via segment counts. Automatic differentiation systems that can analyze the computation graph and determine optimal checkpoint placement without user intervention would eliminate this manual tuning step. Research systems like Checkmate have demonstrated that optimal placement can provide 20-30% better memory-compute tradeoffs than equal-spacing heuristics for heterogeneous networks. Incorporating this into production training frameworks remains an active research direction.

How PyTorch Implements Activation Checkpointing

Understanding the implementation mechanics helps you debug issues and use the API correctly. The PyTorch implementation of checkpoint.checkpoint with use_reentrant=False works as follows.

The autograd.Function Wrapper

The non-reentrant implementation uses a custom torch.autograd.Function subclass. This class has two class methods, forward and backward, that define the computation during both passes. The forward method runs the user-provided function (the module to be checkpointed) but uses torch.no_grad() to prevent PyTorch from building an autograd graph during this run. This is why activations are not retained: with no_grad() active, PyTorch discards the computation graph as it is built.

The forward method saves the inputs to the checkpointed region (the boundary tensors), the function itself, and the RNG state. It does not save any intermediate activations.

During the backward method, PyTorch calls the saved function again on the saved inputs, this time with autograd enabled. This recomputation run builds a fresh computation graph for the checkpointed segment. The backward pass then traverses this fresh graph to compute gradients, after which the graph is immediately discarded.

RNG State Management

The most subtle aspect of the implementation is RNG state management. Suppose your checkpointed segment contains a dropout layer. Dropout samples a random mask during the forward pass. If the backward pass recomputes dropout with a different random mask, the gradients will be incorrect (they will reflect a different network configuration than the one that produced the forward pass output).

PyTorch addresses this by recording the RNG state on both CPU and CUDA immediately before the checkpointed region executes during the forward pass. The state includes the full state of the PyTorch random number generator (a Mersenne Twister on CPU and a device-specific generator on CUDA). During the recomputation in the backward pass, PyTorch restores this state before re-running the function, making sure that any random operations produce exactly the same results as the original forward pass.

This RNG state save-and-restore happens automatically and transparently. You do not need to do anything special when using dropout or other random operations inside a checkpointed region.

The preserve_rng_state Argument

The preserve_rng_state argument to checkpoint.checkpoint (defaulting to True) controls whether the RNG state is saved and restored. You might set this to False if your checkpointed region contains no random operations and you want to avoid the overhead of saving the RNG state. In practice, the overhead is negligible (saving a few hundred bytes), so leaving it at True is safe.

Memory Pinning and Device Management

When checkpoint boundary tensors are saved, they are saved as regular CUDA tensors in HBM. The PyTorch garbage collector will not free them until they are no longer referenced by any Python object or autograd graph. This is what ensures they survive until the backward pass needs them for recomputation.

One practical implication: the checkpoint boundary tensors are included in PyTorch's memory tracking. When you query torch.cuda.memory_allocated(), you will see the boundary tensors contributing to the reported allocation. This is correct behavior: they are truly allocated. The memory reduction from checkpointing shows up as a reduction in the total allocated memory, not a reduction in what is tracked.

Graph Building During Recomputation

A subtlety in the non-reentrant implementation: during the recomputation run in the backward pass, PyTorch builds a complete autograd graph for the checkpointed segment. This graph is then used to compute gradients within the segment. Once the segment's gradients have been computed, the graph is freed. This means the memory spike during recomputation includes the intermediate activations plus the autograd graph nodes, which add a small constant overhead per operation. For large segments with many operations, this overhead can be noticeable.

The reentrant implementation handled this differently: it re-entered the Python autograd machinery directly, which avoided building an explicit graph but had other correctness issues. The non-reentrant approach trades slightly higher peak memory during recomputation for correctness and compatibility with modern PyTorch features.

Interaction with Distributed Training

Activation checkpointing does not operate in isolation: it interacts with every other component of the training infrastructure stack. Understanding these interactions is essential for configuring large-scale training runs correctly.

Pipeline Parallelism

In pipeline-parallel training, the model is split across multiple GPUs (stages), with each stage holding a portion of the layers. Data flows through stages in sequence, and multiple micro-batches are in flight simultaneously to keep all stages busy. This creates a critical interaction with activation memory: each stage must hold activations for all micro-batches currently buffered in its pipeline.

If a stage has mm micro-batches in flight and each micro-batch has activation memory MM, the total activation memory for that stage is m×Mm \times M. For large pipeline schedules (many micro-batches, many stages), this can be enormous. Activation checkpointing within each stage reduces the per-micro-batch activation memory, but the multiplier mm still applies to the checkpointed memory.

The interleaved pipeline schedule (as used in Megatron-LM) assigns each stage multiple non-contiguous chunks of layers, which reduces the number of micro-batches needed to fill the pipeline and consequently reduces the activation memory multiplier. Combining the interleaved schedule with per-layer checkpointing provides both benefits.

Tensor Parallelism

In tensor-parallel training, each layer's weight matrices are split across multiple GPUs along a dimension (typically the hidden dimension). When using tensor parallelism, the activation tensors are also split: each GPU holds a shard of shape [B,T,d/tp][B, T, d/\text{tp}] where tp\text{tp} is the tensor parallel size. This directly reduces activation memory by a factor of tp\text{tp}, acting as another form of memory reduction that stacks multiplicatively with activation checkpointing.

When combining tensor and pipeline parallelism with activation checkpointing (the "3D parallelism" stack), the memory savings multiply. A checkpoint boundary tensor that would occupy 67 MB in the baseline might occupy only 4 MB per GPU in a tp=8\text{tp} = 8, pp=2\text{pp} = 2 configuration. Checkpointing every layer in this setting saves memory from a much smaller base, meaning you can often reduce the checkpoint frequency (fewer segments, less compute overhead) while still staying within the memory budget.

Gradient Accumulation

Gradient accumulation splits a large logical batch into multiple smaller micro-batches, runs forward and backward passes on each separately, accumulates the gradients, and only updates the optimizer once per logical batch. This is used to simulate large batch sizes when memory would not allow fitting the full batch in one pass.

The interaction with activation checkpointing is additive rather than multiplicative in the memory dimension: each micro-batch's activations are independent. You run forward-backward for micro-batch 1 (using checkpointed activations), then run forward-backward for micro-batch 2, and so on. Because micro-batches do not coexist in memory during the backward pass, gradient accumulation does not increase activation memory. Activation checkpointing reduces the memory within each micro-batch's pass, and gradient accumulation reduces the effective batch size per pass. Together, they allow training at large logical batch sizes with very limited per-pass memory.

Data Parallelism and Gradient Checkpointing

Data-parallel training replicates the model on multiple GPUs and splits the batch across them. Each GPU runs a complete forward-backward pass on its local batch shard. Gradient checkpointing applies independently per GPU, since each GPU's computation is identical but on different data. The memory savings are the same as in the single-GPU case. The only interaction is that data-parallel gradient synchronization (all-reduce) happens on the gradient tensors after the backward pass, which is after all activations have been freed.

Code Implementation

Let's implement activation checkpointing from scratch to understand the mechanics, then demonstrate PyTorch's built-in interface.

Setup and Dependencies

A Simple Transformer Block for Demonstration

We will build a minimal transformer-like block with enough structure to show meaningful activation memory differences.

In[7]:
Code
import torch
import torch.nn as nn
import torch.utils.checkpoint as checkpoint


class FeedForward(nn.Module):
    def __init__(self, d_model: int, d_ff: int):
        super().__init__()
        self.fc1 = nn.Linear(d_model, d_ff)
        self.act = nn.GELU()
        self.fc2 = nn.Linear(d_ff, d_model)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return self.fc2(self.act(self.fc1(x)))


class TransformerBlock(nn.Module):
    def __init__(self, d_model: int, n_heads: int, d_ff: int):
        super().__init__()
        self.norm1 = nn.LayerNorm(d_model)
        self.attn = nn.MultiheadAttention(d_model, n_heads, batch_first=True)
        self.norm2 = nn.LayerNorm(d_model)
        self.ffn = FeedForward(d_model, d_ff)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        # Self-attention with residual
        normed = self.norm1(x)
        attn_out, _ = self.attn(normed, normed, normed, need_weights=False)
        x = x + attn_out
        # Feed-forward with residual
        x = x + self.ffn(self.norm2(x))
        return x


class TransformerModel(nn.Module):
    def __init__(
        self,
        n_layers: int,
        d_model: int,
        n_heads: int,
        d_ff: int,
        use_checkpoint: bool = False,
    ):
        super().__init__()
        self.layers = nn.ModuleList(
            [TransformerBlock(d_model, n_heads, d_ff) for _ in range(n_layers)]
        )
        self.use_checkpoint = use_checkpoint

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        for layer in self.layers:
            if self.use_checkpoint:
                x = checkpoint.checkpoint(layer, x, use_reentrant=False)
            else:
                x = layer(x)
        return x

The key difference between the two models is in the forward pass: when use_checkpoint=True, each layer is wrapped with checkpoint.checkpoint(). This tells PyTorch to discard intermediate activations for that layer after the forward pass and recompute them during backpropagation. The use_reentrant=False argument is the recommended modern interface, which avoids limitations of the older reentrant implementation.

The use_reentrant=False mode is important for several reasons. The legacy reentrant mode used Python re-entry into autograd to simulate the checkpointing behavior, which created restrictions on what operations could be performed inside the checkpointed region. In particular, operations that used in-place modification of tensors or that registered custom backward hooks behaved incorrectly. The non-reentrant mode avoids these issues by using a cleaner implementation based on autograd.Function.

Measuring Memory and Time

In[8]:
Code
import time


def measure_training_step(
    n_layers: int,
    d_model: int,
    n_heads: int,
    d_ff: int,
    batch_size: int,
    seq_len: int,
    use_checkpoint: bool,
) -> dict:
    """Run one forward+backward step and record peak memory and wall time."""
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

    model = TransformerModel(
        n_layers, d_model, n_heads, d_ff, use_checkpoint=use_checkpoint
    ).to(device)
    optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)

    x = torch.randn(batch_size, seq_len, d_model, device=device)
    target = torch.randn(batch_size, seq_len, d_model, device=device)

    if device.type == "cuda":
        torch.cuda.reset_peak_memory_stats(device)
        torch.cuda.synchronize()

    t0 = time.perf_counter()
    optimizer.zero_grad()
    out = model(x)
    loss = nn.MSELoss()(out, target)
    loss.backward()
    optimizer.step()

    if device.type == "cuda":
        torch.cuda.synchronize()

    t1 = time.perf_counter()

    peak_mb = (
        torch.cuda.max_memory_allocated(device) / 1024**2
        if device.type == "cuda"
        else 0.0
    )

    return {
        "use_checkpoint": use_checkpoint,
        "peak_memory_mb": peak_mb,
        "time_ms": (t1 - t0) * 1000,
    }


# Configuration
N_LAYERS = 8
D_MODEL = 512
N_HEADS = 8
D_FF = 2048
BATCH_SIZE = 4
SEQ_LEN = 256

results_baseline = measure_training_step(
    N_LAYERS, D_MODEL, N_HEADS, D_FF, BATCH_SIZE, SEQ_LEN, use_checkpoint=False
)
results_ckpt = measure_training_step(
    N_LAYERS, D_MODEL, N_HEADS, D_FF, BATCH_SIZE, SEQ_LEN, use_checkpoint=True
)
Out[9]:
Console
Training Step Comparison
--------------------------------------------------

Baseline (no checkpointing) (CPU, no memory tracking)
  Peak memory:   0.0 MB
  Step time:     287.8 ms

With activation checkpointing (CPU, no memory tracking)
  Peak memory:   0.0 MB
  Step time:     300.5 ms

On a CUDA device, activation checkpointing typically reduces peak memory by 3-5x for this architecture, with a compute overhead of 1.2-1.4x. The exact ratio depends on layer size and the fraction of activation memory that dominates total usage. On CPU (where memory measurement is not per-operation), the timing overhead still reflects the recomputation cost.

Profiling Activation Memory Per Layer

In[10]:
Code
def profile_activation_memory_scaling(
    n_layers_list: list,
    d_model: int,
    n_heads: int,
    d_ff: int,
    batch_size: int,
    seq_len: int,
) -> dict:
    """Profile peak memory vs number of layers with and without checkpointing."""
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    results = {"n_layers": [], "baseline_mb": [], "ckpt_mb": []}

    for n_layers in n_layers_list:
        for use_ckpt, key in [(False, "baseline_mb"), (True, "ckpt_mb")]:
            model = TransformerModel(
                n_layers, d_model, n_heads, d_ff, use_checkpoint=use_ckpt
            ).to(device)
            x = torch.randn(batch_size, seq_len, d_model, device=device)
            target = torch.randn(batch_size, seq_len, d_model, device=device)

            if device.type == "cuda":
                torch.cuda.reset_peak_memory_stats(device)

            out = model(x)
            loss = nn.MSELoss()(out, target)
            loss.backward()

            peak = (
                torch.cuda.max_memory_allocated(device) / 1024**2
                if device.type == "cuda"
                else float(n_layers) * 50
            )
            results[key].append(peak)
            del model, out, loss, x, target

        results["n_layers"].append(n_layers)

    return results


n_layers_list = [2, 4, 6, 8, 10, 12]
scaling_results = profile_activation_memory_scaling(
    n_layers_list, d_model=512, n_heads=8, d_ff=2048, batch_size=4, seq_len=256
)
Out[11]:
Console
Memory Scaling with Number of Layers
  Layers   Baseline (MB)    Checkpointed (MB)    Reduction
------------------------------------------------------------
       2           100.0                100.0        1.00x
       4           200.0                200.0        1.00x
       6           300.0                300.0        1.00x
       8           400.0                400.0        1.00x
      10           500.0                500.0        1.00x
      12           600.0                600.0        1.00x

The table shows the linear-versus-sublinear memory scaling property: baseline memory grows roughly linearly with the number of layers, while checkpointed memory grows much more slowly. For larger layer counts, the reduction factor increases, which means the technique becomes more valuable as models scale.

Using PyTorch's checkpoint_sequential for Multiple Segments

For cases where you want checkpoint granularity at a level coarser than one layer (for example, every 4 layers), PyTorch provides checkpoint_sequential:

In[12]:
Code
class TransformerModelCoarseCheckpoint(nn.Module):
    def __init__(
        self,
        n_layers: int,
        d_model: int,
        n_heads: int,
        d_ff: int,
        checkpoint_segments: int = 2,
    ):
        super().__init__()
        self.layers = nn.Sequential(
            *[TransformerBlock(d_model, n_heads, d_ff) for _ in range(n_layers)]
        )
        self.checkpoint_segments = checkpoint_segments

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return checkpoint.checkpoint_sequential(
            self.layers, self.checkpoint_segments, x, use_reentrant=False
        )


# With 8 layers and 2 segments, checkpoints happen at the layer 4 boundary
model_coarse = TransformerModelCoarseCheckpoint(
    n_layers=8, d_model=512, n_heads=8, d_ff=2048, checkpoint_segments=2
)
x_demo = torch.randn(2, 64, 512)
out_coarse = model_coarse(x_demo)
Out[13]:
Console
Checkpoint Segments vs Memory and Time (8 layers)
  Segments   Peak Memory (MB)    Time (ms)                 Note
-----------------------------------------------------------------
         1                0.0        215.9        max recompute
         2                0.0        237.4           2 segments
         4                0.0        247.0           4 segments
         8                0.0        251.0 per-layer (min memory)

Fewer segments (coarser checkpointing) saves less memory but adds less compute overhead. More segments save more memory but add more compute. The number of segments is effectively the L/kL/k parameter from the earlier analysis: 1 segment means no checkpointing, LL segments means per-layer checkpointing.

Implementing Selective Checkpointing

Beyond the standard per-layer interface, you can implement selective checkpointing by wrapping only specific sublayers within a transformer block. This gives finer control over the memory-compute tradeoff:

In[14]:
Code
class SelectiveCheckpointTransformerBlock(nn.Module):
    """Transformer block with selective checkpointing of the FFN sublayer only."""

    def __init__(
        self, d_model: int, n_heads: int, d_ff: int, checkpoint_ffn: bool = True
    ):
        super().__init__()
        self.norm1 = nn.LayerNorm(d_model)
        self.attn = nn.MultiheadAttention(d_model, n_heads, batch_first=True)
        self.norm2 = nn.LayerNorm(d_model)
        self.ffn = FeedForward(d_model, d_ff)
        self.checkpoint_ffn = checkpoint_ffn

    def _ffn_forward(self, x: torch.Tensor) -> torch.Tensor:
        return self.ffn(self.norm2(x))

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        # Attention sublayer: store activations normally
        normed = self.norm1(x)
        attn_out, _ = self.attn(normed, normed, normed, need_weights=False)
        x = x + attn_out

        # FFN sublayer: optionally checkpoint
        if self.checkpoint_ffn:
            ffn_out = checkpoint.checkpoint(
                self._ffn_forward, x, use_reentrant=False
            )
        else:
            ffn_out = self._ffn_forward(x)

        x = x + ffn_out
        return x


# Build a model that selectively checkpoints only FFN sublayers
class SelectiveCheckpointModel(nn.Module):
    def __init__(
        self,
        n_layers: int,
        d_model: int,
        n_heads: int,
        d_ff: int,
        checkpoint_ffn: bool = True,
    ):
        super().__init__()
        self.layers = nn.ModuleList(
            [
                SelectiveCheckpointTransformerBlock(
                    d_model, n_heads, d_ff, checkpoint_ffn=checkpoint_ffn
                )
                for _ in range(n_layers)
            ]
        )

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        for layer in self.layers:
            x = layer(x)
        return x


model_selective = SelectiveCheckpointModel(
    n_layers=8, d_model=512, n_heads=8, d_ff=2048, checkpoint_ffn=True
)
x_demo2 = torch.randn(2, 64, 512)
out_selective = model_selective(x_demo2)
Out[15]:
Console
Selective Checkpointing Comparison
Strategy                          Memory (MB)    Time (ms)
------------------------------------------------------------
No checkpointing                          0.0        213.1
FFN-only checkpointing                    0.0        229.7

The selective approach checkpoints only the FFN intermediate activations, which are the largest single activation in the feed-forward sublayer. The attention sublayer's activations are kept in memory. This is approximately half the memory savings of full checkpointing, but adds roughly half the compute overhead, showing the near-linear relationship between the fraction of activations checkpointed and both the memory benefit and the compute cost.

Key Parameters

The key configuration choices for activation checkpointing are:

  • use_reentrant: Set to False for the modern interface. The legacy True mode has restrictions on operations inside checkpointed regions and can cause subtle bugs with custom autograd functions.
  • Checkpoint granularity: Whether to checkpoint every layer (checkpoint.checkpoint per layer), every kk layers (checkpoint_sequential with L/kL/k segments), or specific sublayers within a layer.
  • Combining with FlashAttention: When using FlashAttention, attention weight materialization is already avoided, so checkpointing provides smaller marginal memory benefit for attention layers. Focus checkpointing on feed-forward layers in this case.
  • Gradient accumulation interaction: When using gradient accumulation over multiple micro-batches, activation memory is per-micro-batch. Checkpointing and gradient accumulation together are a common pattern: checkpointing reduces per-micro-batch activation memory, while gradient accumulation allows simulating larger batch sizes without increasing per-step peak memory.

Recomputation Correctness and Determinism

A subtle but important requirement for activation checkpointing is that the recomputed activations must be numerically identical to the activations produced during the original forward pass. If the recomputation produces different values, the gradients computed during backpropagation will be incorrect, leading to silent training failures that can be very difficult to diagnose.

For most operations in standard transformer architectures, recomputation is deterministic: given the same input tensor, the same layer produces the same output. However, several operations can introduce non-determinism:

Dropout. Dropout randomly zeroes a fraction of tensor elements during training. A naive recomputation would re-apply dropout with a different random mask, producing different activations and incorrect gradients. The correct behavior is to apply the same mask during recomputation as was applied during the original forward pass. PyTorch handles this automatically by saving the random number generator state at each checkpoint boundary and restoring it during recomputation.

Random number generation in custom layers. Any custom operation that draws from a random distribution during the forward pass must either be made deterministic or must save its random state alongside the checkpoint tensor. PyTorch's non-reentrant checkpointing saves the CPU and CUDA random states automatically, which covers most use cases.

Non-deterministic CUDA operations. Some CUDA operations (particularly certain reduction operations) are non-deterministic by default on NVIDIA GPUs, meaning they can produce slightly different results on different runs due to floating-point non-associativity. This rarely causes correctness issues for checkpointing, but can complicate debugging if you are comparing activations between the original forward pass and the recomputed version.

The PyTorch implementation handles the random state correctly by design. When you call checkpoint.checkpoint with use_reentrant=False, the implementation records the RNG state before the forward pass and restores it before the recomputation, making sure dropout and similar operations produce identical results.

Limitations and Practical Considerations

Activation checkpointing is not universally beneficial, and several practical constraints limit its applicability.

Compute Overhead on Memory-Bandwidth-Bound Workloads

The 33% compute overhead is a reasonable cost when the training bottleneck is GPU memory capacity. But not all training configurations are memory-bound. If you are training a small model on a large GPU where activation memory is a minor fraction of total memory, the added recomputation moves you from memory-bound to compute-bound, wasting time without enabling anything new.

The technique is most valuable when it is the enabling factor for a configuration that would otherwise cause an out-of-memory error or require a reduced batch size. A practical test: if you can already train at your desired batch size and sequence length without checkpointing, the overhead is pure waste. Enable checkpointing only when it allows you to increase batch size, fit a larger model, or extend sequence length in ways that improve training efficiency or model quality.

On modern hardware like H100 GPUs with HBM3 memory, the compute-to-memory-bandwidth ratio is much higher than on older A100 or V100 GPUs. This means recomputation is relatively cheaper in terms of wall time because the GPU executes matrix multiplications much faster relative to memory accesses. On H100 hardware, the effective overhead of full per-layer checkpointing is often closer to 20-25% rather than 33%, because memory bandwidth savings from smaller activation tensors accelerate the overall forward pass.

Incompatibility with Some Autograd Operations

The use_reentrant=False mode requires that all inputs to a checkpointed segment that require gradient computation are properly handled by PyTorch's autograd graph. Some operations do not play well with non-reentrant checkpointing, including certain custom CUDA extensions and operations that use in-place mutations on tensors that are returned as outputs.

When you encounter unexpected errors in checkpointed regions, the first diagnostic step is to run without checkpointing to isolate whether the error is checkpointing-specific. If the model runs correctly without checkpointing but fails with it, the issue is almost always one of: (1) in-place modification of a tensor that appears in the checkpoint's computation graph, (2) a custom autograd function that does not correctly implement its backward pass under checkpointing, or (3) operations that break the RNG state restoration.

Interaction with Compilation (torch.compile)

PyTorch 2.0's torch.compile performs whole-graph optimizations that can conflict with the boundaries imposed by checkpoint.checkpoint. The compiler may not be able to fuse operations across checkpoint boundaries as aggressively as it would without them. In some configurations, using torch.compile with checkpointing requires careful configuration of the fullgraph=False option or selective application of compilation to non-checkpointed regions.

The underlying tension is that torch.compile works best when it can see the entire computation graph and optimize across it globally. Checkpoint boundaries break the graph into segments that are individually compiled, preventing cross-boundary fusion. For some operations, this can negate a significant fraction of the compilation speedup. This is an active area of tooling improvement in PyTorch, and the interaction is expected to improve in future releases.

Double Memory Peak During Recomputation

During the backward pass, when a segment is being recomputed, there is a brief period where both the checkpoint input activation (from the saved boundary) and the newly recomputed intermediate activations coexist in memory. This creates a temporary memory spike above the steady-state checkpointed level.

For a checkpointed model with L/kL/k boundary tensors and kk-layer segments, the steady-state activation memory is approximately L/k⋅MckptL/k \cdot M_{\text{ckpt}} (boundaries) plus k⋅Mlayerk \cdot M_{\text{layer}} (current recomputation segment). During the recomputation spike, you additionally need the previous segment's boundary tensor MckptM_{\text{ckpt}} (to re-run from). So the peak is slightly higher than the steady state by one extra boundary tensor.

For models operating very close to the memory limit, this spike can still cause an out-of-memory error even with checkpointing enabled. The solution is to either increase checkpoint frequency (smaller segments, more boundaries, lower spike) or reduce the batch size slightly as a buffer.

Practical Configuration Guidance

For practitioners setting up activation checkpointing for a new training run, the following decision process covers most scenarios:

Start with a memory profile. Before enabling checkpointing, measure memory usage in your training configuration. Use torch.cuda.memory_allocated() and torch.cuda.max_memory_allocated() at different points in the training step to identify the peak. If the peak occurs during the forward pass (after many layers have been processed), activation memory is the culprit. If the peak occurs at the optimizer step, optimizer state memory is the constraint, and checkpointing will not help.

Enable FlashAttention first. If you are not already using FlashAttention, enable it before considering explicit activation checkpointing. FlashAttention provides 50-70% memory reduction for attention weights with near-zero compute overhead on modern hardware. After enabling it, re-profile to see how much activation memory remains.

Choose checkpointing granularity based on memory budget. If you are still memory-constrained after FlashAttention, enable per-layer checkpointing. Start with full per-layer checkpointing (k=1) to confirm correctness, then experiment with coarser granularity (k=2 or k=4) to recover some compute speed while staying within the memory budget.

Monitor training loss curves for anomalies. After enabling checkpointing, training loss should behave identically to the baseline (no checkpointing) run, since the gradients are mathematically equivalent. If training loss diverges or shows unusual spikes, disable checkpointing and compare: an anomaly that disappears without checkpointing indicates a compatibility issue with a specific operation in the model.

Account for the memory spike. If you are targeting a memory budget at the very limit of your GPU's capacity, leave a 10-15% buffer above your steady-state checkpointed memory usage to accommodate the recomputation spike during the backward pass. Running right at the limit risks out-of-memory errors that can be difficult to reproduce consistently.

Larger Models and Longer Contexts

Activation checkpointing has been a critical enabling technique throughout the scaling era of language models. GPT-3 (Brown et al., 2020) and subsequent models used gradient checkpointing to fit training within available GPU memory budgets. The training codebase for Megatron-LM (the framework used for GPT-3 and many subsequent large models) has included activation checkpointing as a core feature since its initial release.

The most significant recent development is the interaction between activation checkpointing and FlashAttention. By eliminating the O(T2)O(T^2) attention memory cost through kernel-level recomputation, FlashAttention made long-context training (sequences of 8K, 32K, or 128K tokens) feasible without proportionally increasing activation memory. For a sequence length of 128K tokens with 32 attention heads and batch size 1, the attention weight matrix alone would occupy:

1×32×(128×103)2×2≈1 TB1 \times 32 \times (128 \times 10^3)^2 \times 2 \approx 1 \text{ TB}

This is physically impossible to store. FlashAttention's kernel-level recomputation eliminates this memory requirement entirely, making 128K context training feasible on clusters of A100 GPUs.

The combination of FlashAttention (implicit attention checkpointing) plus selective FFN checkpointing plus mixed precision training (discussed in the next chapter) defines the current standard memory optimization stack for training large language models. This stack is what enabled models like Claude, GPT-4, and Gemini to train on very long documents at scale. Activation checkpointing alone was not sufficient to reach these context lengths, but it remains an indispensable component of the full stack.

Summary

Activation checkpointing addresses the fundamental memory bottleneck in deep network training by discarding intermediate activations during the forward pass and recomputing them on demand during backpropagation. The key concepts from this chapter are:

  • The memory-compute tradeoff. Full per-layer checkpointing reduces activation memory by roughly 10-20x at the cost of approximately 33% more compute, by running a second forward pass through each layer during backpropagation. The overhead is only 33% rather than 100% because the backward pass already costs twice the forward pass.

  • Why backpropagation needs activations. Gradient computation at each layer requires knowing the values that were present during the forward pass. Nonlinear operations like GELU cannot be inverted, so their inputs must be stored or recomputed. Every activation dependency creates a memory reservation that lasts until the corresponding backward pass gradient is computed.

  • Checkpoint placement. Transformer layer boundaries are natural checkpoint points because they have small, well-defined activation tensors (just the residual stream). Within a layer, attention weight matrices and FFN intermediate activations are the most expensive activations to store, making them the highest-priority targets for selective checkpointing.

  • Selective checkpointing. Applying checkpointing only to high-cost activations while retaining cheap activations in memory reduces compute overhead below 33% while recovering most of the memory benefit. The FlashAttention-plus-selective-FFN configuration achieves around 80% memory savings at roughly 12% compute overhead.

  • FlashAttention interaction. FlashAttention implements implicit attention-weight checkpointing at the kernel level, handling the O(T2)O(T^2) memory cost without explicit overhead. When FlashAttention is used, remaining checkpointing targets are primarily feed-forward layers.

  • Distributed training interactions. Activation checkpointing interacts with pipeline parallelism (reduces per-micro-batch memory at each stage), tensor parallelism (stacks multiplicatively), and gradient accumulation (each micro-batch's activations are independent). Understanding these interactions is necessary for configuring the full training infrastructure stack.

  • Determinism requirement. Recomputed activations must exactly match original forward pass activations. PyTorch handles this automatically by saving and restoring random number generator states at checkpoint boundaries, which is critical for operations like dropout.

  • PyTorch API. torch.utils.checkpoint.checkpoint(module, input, use_reentrant=False) implements per-layer checkpointing. checkpoint.checkpoint_sequential(modules, n_chunks, input) implements coarser segmentation for reduced overhead. The use_reentrant=False argument is required for the modern, correct implementation.

  • Limitations. Checkpointing adds a temporary memory spike during recomputation, can interact poorly with torch.compile fusing, and is most valuable when activation memory is the binding constraint for a configuration. On hardware with high compute-to-bandwidth ratios (like H100), the practical overhead is often lower than the theoretical 33%.

  • Historical significance. The technique originated in scientific computing in the 1980s and was formalized for neural networks by Chen et al. in 2016. It has been indispensable for training models at the scale of GPT-3 and beyond, and remains a core component of every large-scale language model training stack today.

Activation checkpointing is one of those techniques that becomes more important as models scale. A technique that saves 20x memory is helpful at small scales but essential at large scales. As language models continue to grow in size and context length, and as FlashAttention handles the attention memory problem, the remaining challenge is feed-forward memory scaling. Continuing advances in selective checkpointing, activation compression, and automated checkpoint placement will shape how practitioners manage this constraint for the next generation of models.

Quiz

Ready to test your understanding? Take this quick quiz to reinforce what you've learned about activation checkpointing.

Activation Checkpointing Quiz

Question 1 of 80 of 8 completed
What is the primary resource that activation checkpointing trades away to reduce GPU memory usage?

Comments

1 comment

  1. ARAVINDAN RaviMember

    It was really awesome to understand the entire techniques and practical considerations. Thanks a lot !!

    1. Michael BrenndoerferMember

      Very Kind of you - glad you found it useful!

Reference

Citation details

Cite or share this article.

BIBTEXAcademic
@misc{brenndoerfer2026activationcheckpointing, author = {Michael Brenndoerfer}, title = {Activation Checkpointing: Gradient Memory}, year = {2026}, url = {https://mbrenndoerfer.com/writing/activation-checkpointing-gradient-memory-selective-recomputation}, organization = {mbrenndoerfer.com}, note = {Accessed: 2026-09-30} }
APAAcademic
Michael Brenndoerfer (2026). Activation Checkpointing: Gradient Memory. Retrieved from https://mbrenndoerfer.com/writing/activation-checkpointing-gradient-memory-selective-recomputation
MLAAcademic
Michael Brenndoerfer. "Activation Checkpointing: Gradient Memory." 2026. Web. September 30, 2026. <https://mbrenndoerfer.com/writing/activation-checkpointing-gradient-memory-selective-recomputation>.
CHICAGOAcademic
Michael Brenndoerfer. "Activation Checkpointing: Gradient Memory." Accessed September 30, 2026. https://mbrenndoerfer.com/writing/activation-checkpointing-gradient-memory-selective-recomputation.
HARVARDAcademic
Michael Brenndoerfer (2026) 'Activation Checkpointing: Gradient Memory'. Available at: https://mbrenndoerfer.com/writing/activation-checkpointing-gradient-memory-selective-recomputation (Accessed: September 30, 2026).
SimpleBasic
Michael Brenndoerfer (2026). Activation Checkpointing: Gradient Memory. https://mbrenndoerfer.com/writing/activation-checkpointing-gradient-memory-selective-recomputation

About the author

Continue with the full handbook

This chapter is part of Language AI Handbook. Use the handbook page to browse the complete table of contents and continue reading in sequence.

Explore Language AI Handbook
Newsletter

Stay up to date

Get articles, book updates, and news delivered to your inbox.

No spam, unsubscribe anytime.

or

Join the community

Sign in to remove popups, track your reading progress, and join the discussion.