GPU Memory and Speed

Gradient checkpointing

Gradient checkpointing throws away most in-between results during the forward pass and recomputes them during backward — trading roughly a third more compute for a large memory saving.

On this page 5
  1. Why it exists
  2. How it works
  3. A real example you have seen
  4. Remember this
  5. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

Gradient checkpointing saves memory by throwing away rough work and redoing it later, instead of keeping every page.

Imagine solving a long maths assignment on a small desk. Normally you keep every sheet of working, because checking your answers later needs them. The desk overflows. The trick: keep only one sheet per chapter, bin the rest, and when checking time comes, redo the working between saved sheets.

You pay with some repeated effort. You win a mostly empty desk.

Why it exists

From the memory breakdown lesson: the biggest memory tenant during training is usually activations — the in-between results of every layer, kept for the correction step.

A deep model is a chain of layers, and normally every link of the chain stores its output. Checkpointing keeps outputs only at a few chosen links — the checkpoints — and forgets everything between them. During the correction pass, it recomputes each forgotten stretch from the nearest checkpoint, uses it, and forgets it again.

How it works

normal:        A -> B -> C -> D -> E -> F -> G -> H
stored:        *    *    *    *    *    *    *    *    (8 saved)

checkpointed:  A -> B -> C -> D -> E -> F -> G -> H
stored:        *              *              *          (3 saved)
                    ...redo B,C when their turn comes...

The forward pass runs a second time in pieces, so a step takes noticeably longer. In exchange, batch sizes that crashed before now fit.

A real example you have seen

Every "fine-tune a large language model on one GPU" tutorial you have seen relies on this, usually as a single flag called gradient_checkpointing=True. Without it, models of that size cannot store a full set of activations on one card at any useful batch size.

Remember this

  • Checkpointing forgets most in-between results and recomputes them on demand.
  • It trades roughly a third more compute for a large activation saving.
  • It is the standard move when the model fits but the batch does not.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

The first demo runs anywhere, CPU included. The memory measurement needs a GPU; those numbers were captured with torch 2.5.1 on an NVIDIA RTX A6000.

Watch the recomputation happen

The cleanest proof that forward runs twice — a counter:

ckpt_counter.py
import torch
import torch.nn as nn
from torch.utils.checkpoint import checkpoint

calls = {"count": 0}

class Block(nn.Module):
    def __init__(self):
        super().__init__()
        self.layer = nn.Linear(64, 64)
    def forward(self, x):
        calls["count"] += 1
        return torch.relu(self.layer(x))

torch.manual_seed(0)
block = Block()
x = torch.randn(8, 64, requires_grad=True)

calls["count"] = 0
block(x).sum().backward()
print("normal:        forward ran", calls["count"], "time(s)")

calls["count"] = 0
checkpoint(block, x, use_reentrant=False).sum().backward()
print("checkpointed:  forward ran", calls["count"], "time(s)")
Output
normal:        forward ran 1 time(s)
checkpointed:  forward ran 2 time(s)

checkpoint(block, x) runs the block without storing its internal activations. When backward() reaches that region, the block runs forward again — the counter catches it red-handed — and this time the activations are kept, used, and dropped.

use_reentrant=False selects the modern implementation. Pass it explicitly: the old default has sharp edges around unused inputs and RNG state, and newer PyTorch versions warn until you choose.

Measure the saving

For a stack of layers, checkpoint_sequential splits it into segments and stores only segment boundaries:

ckpt_mem.py
import torch
import torch.nn as nn
from torch.utils.checkpoint import checkpoint_sequential

if not torch.cuda.is_available():
    raise SystemExit("needs a GPU to measure memory")

def peak_mb(use_ckpt):
    torch.manual_seed(0)
    layers = nn.Sequential(*[
        nn.Sequential(nn.Linear(1024, 1024), nn.ReLU()) for _ in range(16)
    ]).cuda()
    x = torch.randn(4096, 1024, device="cuda", requires_grad=True)
    torch.cuda.reset_peak_memory_stats()
    if use_ckpt:
        out = checkpoint_sequential(layers, segments=4, input=x, use_reentrant=False)
    else:
        out = layers(x)
    out.sum().backward()
    return torch.cuda.max_memory_allocated() / 1024**2

print(f"without checkpointing: {peak_mb(False):7.1f} MB peak")
print(f"with 4 segments:       {peak_mb(True):7.1f} MB peak")
Output
without checkpointing:   420.3 MB peak
with 4 segments:         276.4 MB peak

A third of the peak gone, on a model this shallow. The saving grows with depth: the deeper the network, the larger the share of memory that is skippable in-between results. Real transformer training checkpoints each block and commonly cuts activation memory several-fold.

Common mistakes

Checkpointing modules with randomness or state, carelessly. Dropout runs twice; the modern implementation replays RNG so the two runs match, but custom modules with side effects (counters, caches, prints) will visibly execute twice. BatchNorm updates running statistics on each call — a checkpointed BatchNorm block updates them twice per step in older reentrant mode. Know what is inside the block you wrap.

Wrapping the whole model in one checkpoint. Then nothing in between is stored, and backward recomputes the entire forward — maximum time cost, and the peak is dominated by the single giant recomputed segment. Several mid-sized segments beat one huge one.

Expecting a speedup. This lesson's tool spends time to buy memory. Steps get roughly 20–35% slower. The win is indirect: the memory freed may allow a batch size with better GPU utilisation.

Passing inputs that do not require grad. With nothing to differentiate, older reentrant checkpointing silently skips gradient tracking for the segment. The use_reentrant=False path handles it correctly — one more reason to always pass the flag.

Try it yourself

In ckpt_mem.py, sweep segments over 2, 4, 8 and 16, printing the peak each time. Find the sweet spot, and notice it is not 16 — boundary storage starts to dominate.

What to learn next

Researcher — Mathematics and papers.

The compute–memory trade, quantified

Split an $L$-layer sequential network into $k$ segments. Stored activations drop from $O(L)$ to $O(k + L/k)$ — the boundaries plus one segment's interior, live during that segment's recomputation. Minimising over $k$ gives $k = \sqrt{L}$ and peak activation memory $O(\sqrt{L})$. The price is one extra forward pass: a step normally costs about 3 forward-equivalents (one forward, backward roughly twice that), and checkpointing makes it 4 — ~33% more step time.

Chen et al. (2016) formalised the $O(\sqrt{n})$ result; the fully general problem — optimal recomputation schedules on arbitrary DAGs under a memory budget — is treated by Griewank and Walther's binomial checkpointing (2000) and, for real networks, solved approximately by checkmate (Jain et al., 2020) as an ILP.

Selective recomputation

Uniform checkpointing recomputes cheap and expensive activations alike. Korthikanti et al. (2022) note the attention score tensors are enormous but nearly free to recompute, while linear-layer outputs are the reverse. Recomputing only attention internals captures most of the memory win at a few percent extra compute — the default in Megatron-style stacks, and largely subsumed by FlashAttention (see the sdpa lesson), which never materialises those tensors at all.

Composability notes: checkpointing composes cleanly with AMP (recomputation replays autocast state) and with FSDP, where apply_activation_checkpointing wraps blocks post-sharding. The non-reentrant implementation is built on saved-tensor hooks (torch.autograd.graph.saved_tensors_hooks), the same machinery usable for offloading activations to CPU instead of recomputing.

References

  • Chen, Xu, Zhang, Guestrin (2016), Training Deep Nets with Sublinear Memory Cost — the $O(\sqrt{n})$ schedule.
  • Griewank and Walther (2000), Algorithm 799: revolve — optimal binomial checkpointing from automatic differentiation.
  • Jain et al. (2020), Checkmate: Breaking the Memory Wall with Optimal Tensor Rematerialization, MLSys.
  • Korthikanti et al. (2022), Reducing Activation Recomputation in Large Transformer Models — selective recomputation.

What to learn next