Gradient checkpointing, also called activation checkpointing or activation recomputation, is the trick that lets a model train at a sequence length or batch size that would otherwise run out of GPU memory. Instead of keeping every intermediate tensor from the forward pass until backward needs it, you keep only a few of them and recompute the rest during backward. You pay in arithmetic to save memory. On a modern GPU that trade is usually good, because memory capacity is the hard wall and compute is the soft one.

A word on naming, because this site has two other pages with "checkpointing" in the title. Those (LLM training checkpointing and the GPU checkpointing deep dive) are about saving model and optimizer state to storage so a job can resume. This page is about something entirely different: which activations stay in HBM during one training step. Nothing here touches disk.

By the end you will be able to estimate a transformer's activation memory from its shape, pick between full, selective and partial recomputation, write the PyTorch code for each, and read your utilisation numbers correctly once recomputation is switched on.

Why activations dominate

Backpropagation needs, for every operation, some of the values that operation saw in the forward pass. A matrix multiply needs its input to compute the weight gradient. A GELU needs its input to compute its local derivative. Softmax needs its output. Dropout needs its mask. Autograd records these "saved tensors" as the forward pass runs and releases each one only after the matching backward step consumes it. So at the moment the forward pass finishes, the saved tensors of every layer are alive at once. That is the activation peak.

Weights, gradients and optimizer state scale with parameter count and do not care about sequence length. Activations scale with batch size times sequence length times hidden size times layer count. Double the context and the activations double, or worse if attention scores are materialised. That is why long-context training and fine-tuning hit out-of-memory errors that adding more data-parallel GPUs does not fix. Sharding optimizer state with ZeRO or FSDP shrinks the parameter-shaped terms but leaves each GPU's activations exactly where they were.

Counting activation bytes

Korthikanti and colleagues at NVIDIA ("Reducing Activation Recomputation in Large Transformer Models", 2022) counted the saved bytes of a GPT-style transformer layer with 16-bit activations and no model parallelism. With sequence length s, micro-batch b, hidden size h and a attention heads, one layer stores about:

activation_bytes_per_layer = s * b * h * (34 + 5 * a * s / h)

The 34 covers the QKV and output projections, the MLP with its 4h expansion, layer norms and dropout masks. The 5as/h term is the attention score matrix, its softmax and its dropout mask, which grow with the square of the sequence. Fused attention kernels such as FlashAttention never materialise that matrix, so with them the per-layer figure falls to roughly 34sbh. Models with SwiGLU MLPs and RMSNorm (most current open models) have a somewhat different constant, so treat 34 as an estimate and confirm with the profiler.

If you keep only each layer's input and recompute everything inside it, the layer stores just that input: 2sbh bytes. During backward you also need one layer's full set of intermediates while you recompute and differentiate it, so the peak becomes roughly L * 2sbh + 34sbh instead of L * 34sbh.

Worked example: a 7B model at 4K context

Take a 7B-class model: 32 layers, hidden size 4096, 32 heads, training at sequence length 4096 with micro-batch 1. Then sbh = 4096 * 1 * 4096 = 16.8M.

ConfigurationPer layer32 layersNotes
No recompute, naive attentionsbh(34 + 160) = 3.25 GB~104 GB5as/h = 5*32*4096/4096 = 160; does not fit an 80 GB H100
No recompute, FlashAttention34sbh = 0.57 GB~18 GBFits, but leaves little room for larger batches
Full recompute, FlashAttention2sbh = 34 MB~1.1 GB + one layerPeak ~1.7 GB of activations
Recompute every other layermixed~9.7 GBHalf the extra compute of full recompute

The weights, gradients and Adam state for 7B parameters in mixed precision are around 16 bytes per parameter, about 112 GB, so this model already needs FSDP or ZeRO across several GPUs. Once that is sharded, activations become the term you control per GPU. Full recomputation turns 18 GB into under 2 GB, which you can spend on a micro-batch of 8 instead of 1, or on a 16K context. The arithmetic above is the habit to build: before you reach for checkpointing, compute the activation term, compare it with the free HBM after weights and optimizer state, and decide how much of it you actually need to remove.

What recomputation costs and how to schedule it

How much compute does this cost? For a transformer, backward is roughly twice the FLOPs of forward, so a step is about three forward-equivalents. Recomputing every layer adds one more forward: about 33% more arithmetic. In wall-clock terms the overhead is often 20 to 30%, because part of each step is communication and memory-bound work that recompute does not lengthen.

Chen, Xu, Zhang and Guestrin ("Training Deep Nets with Sublinear Memory Cost", 2016) showed that for a chain of n layers you can store a checkpoint every sqrt(n) layers and recompute the segment between checkpoints, giving O(sqrt(n)) activation memory for the price of one extra forward pass. In transformer practice the segments are whole blocks, and frameworks let you choose how many blocks to checkpoint. That knob is the most useful one: checkpoint only as many blocks as you need to fit, and run the rest normally.

Forward keeps only block inputs; backward replays one block at a timeBlock 1forward, no_grad insidesaved inputBlock 2forward, no_grad insidesaved inputBlock 3forward, no_grad insidesaved inputBlock 4forward, no_grad insidesaved inputBackward of block k: re-run forward of block k from its saved input (graph rebuilt) -> backward through it -> freeBwd block 4done, freedBwd block 3recompute + gradBwd block 2waitingBwd block 1waitingPeak activation memory ~ L saved inputs + one block's full intermediates, instead of L blocks' intermediates.Cost: roughly one extra forward pass per step when every block is checkpointed.
Full-block checkpointing. The forward pass keeps only each block's input; backward re-runs one block's forward, differentiates it and frees it before moving on.

Selective recomputation

Not all saved tensors are equally expensive to recreate. Korthikanti et al. noticed that the attention-score tensors (the 5as/h term) are large but cheap to recompute, while the outputs of the big matrix multiplies are smaller per byte saved but expensive to recompute. Selective recomputation stores the matmul outputs and recomputes only the attention core. They reported it removing most of the activation memory at a few percent extra FLOPs, combined with sequence parallelism to split the remaining 34 term across tensor-parallel ranks.

With FlashAttention the attention core is already recomputed inside the backward kernel, so the textbook form of selective recomputation is partly built in. The general idea survives as an op-level policy: save the outputs of matmuls (and perhaps attention), recompute pointwise ops such as activations, norms and dropout. PyTorch exposes this directly:

import functools
import torch
from torch.utils.checkpoint import (
    checkpoint, create_selective_checkpoint_contexts, CheckpointPolicy)

SAVE = {torch.ops.aten.mm.default, torch.ops.aten.addmm.default,
        torch.ops.aten._scaled_dot_product_flash_attention.default}

def policy(ctx, op, *args, **kwargs):
    # Keep expensive matmul and attention outputs, recompute cheap pointwise ops.
    return CheckpointPolicy.MUST_SAVE if op in SAVE else CheckpointPolicy.PREFER_RECOMPUTE

context_fn = functools.partial(create_selective_checkpoint_contexts, policy)

def block_forward(block, x):
    return checkpoint(block, x, use_reentrant=False, context_fn=context_fn)

This selective-checkpoint API arrived as a prototype around PyTorch 2.5 and the op names depend on your attention backend, so list the ops your model actually dispatches (the profiler or TorchDispatchMode will show them) before writing the set.

Turning it on in PyTorch, Hugging Face and Megatron

Plain PyTorch wraps a callable with torch.utils.checkpoint.checkpoint. Pass use_reentrant=False explicitly: recent releases warn when it is omitted, and the non-reentrant implementation supports keyword arguments, nested checkpoints and inputs that do not require gradients. The reentrant variant, with only a warning, produces no gradients for parameters inside a block when none of the block's tensor inputs require grad, which bites frozen-embedding fine-tunes.

import torch
from torch.utils.checkpoint import checkpoint

class Model(torch.nn.Module):
    def __init__(self, blocks, ckpt_every=1):
        super().__init__()
        self.blocks = torch.nn.ModuleList(blocks)
        self.ckpt_every = ckpt_every          # 1 = all blocks, 2 = every other, 0 = none

    def forward(self, x):
        for i, blk in enumerate(self.blocks):
            if self.training and self.ckpt_every and i % self.ckpt_every == 0:
                x = checkpoint(blk, x, use_reentrant=False)   # rng state preserved by default
            else:
                x = blk(x)
        return x

Hugging Face models expose a switch, model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False}), and the Trainer has a gradient_checkpointing argument. The KV cache is useless during training and conflicts with recomputation, so set model.config.use_cache = False. For FSDP, wrap the transformer block class with apply_activation_checkpointing so the checkpoint boundary and the FSDP unit boundary coincide (see FSDP in depth). Megatron-LM uses flags: --recompute-activations for selective recomputation, or --recompute-granularity full --recompute-method uniform --recompute-num-layers 1 for full-block recomputation; --recompute-method block recomputes only the first N layers of each pipeline stage, which is how you checkpoint "just enough".

Reading utilisation after recompute

Recomputation inflates the FLOPs the GPU executes without inflating the FLOPs the model needs. The Megatron team separates the two: model FLOPs utilisation (MFU) counts only the 6ND training arithmetic, while hardware FLOPs utilisation (HFU) includes the recompute. Turn on full recompute and HFU can rise while MFU and tokens per second fall. Always compare configurations on tokens per second per GPU or MFU, never on HFU or raw kernel throughput.

A related trap is profiling. Recomputed forward kernels appear inside the backward region of a trace, so a step looks backward-heavy. That is expected. What you want to see is that memory headroom became a larger micro-batch or longer sequence, and that the larger batch recovered most of the lost throughput through better kernel efficiency and fewer gradient-accumulation steps.

Failure modes

  • Non-deterministic recompute. If a block uses randomness (dropout) and RNG state is not restored, the recomputed forward differs from the original and gradients are wrong without any error. Keep preserve_rng_state=True (the default); custom CUDA ops with their own RNG need the same care.
  • Side effects inside the block. Anything stateful in the checkpointed function runs twice: running-statistic updates in BatchNorm, counters, logging, MoE load-balancing accumulators. Move side effects out of the wrapped region or guard them.
  • No memory saved. Checkpointing a block whose output is also stored elsewhere (a residual list, a hook, a debugging cache) keeps the tensors alive anyway. Measure with torch.cuda.max_memory_allocated() before and after; do not assume.
  • Frozen inputs with reentrant mode. LoRA and adapter training with frozen embeddings plus use_reentrant=True yields zero gradients for adapters. Use non-reentrant, or call enable_input_require_grads() in Hugging Face models.
  • Overlap with communication. Under FSDP, recomputing a block needs its parameters gathered again; mismatched wrap and checkpoint boundaries cause extra all-gathers. Align them.

Trade-offs

OptionActivation memoryExtra computeUse when
NoneFull0%It fits with the batch you want
Selective / op policyMost removedA few percentDefault for large transformers
Partial (N of L blocks)TunableProportional to N/LYou need a little more room
Full blockMinimal~33% FLOPsLong context, big micro-batch, last resort before offload
CPU offload of activationsMinimal on GPUPCIe trafficWhen compute is the bottleneck but PCIe is idle

Compare against the alternatives too. Tensor and sequence parallelism also shrink per-GPU activations but add communication and need fast links; a smaller micro-batch plus more gradient accumulation costs nothing in FLOPs but hurts kernel efficiency. Recomputation wins when you are already at the parallelism degree your interconnect supports.

What to do next

  1. Compute your activation term with sbh(34 + 5as/h) per layer (drop the second term if you use fused attention) and compare it with HBM left after weights and optimizer state.
  2. Profile one step with torch.cuda.memory._record_memory_history() to confirm the peak is activations, as described in LLM memory profiling.
  3. Enable checkpointing with use_reentrant=False on whole blocks; measure peak memory and tokens per second.
  4. Reduce the number of checkpointed blocks until memory is just under the limit with a safety margin.
  5. Try an op-level selective policy that saves matmul outputs; keep it if tokens per second improve.
  6. Spend reclaimed memory on a larger micro-batch, then report MFU, not HFU.
  7. Add a test that compares gradients with and without checkpointing on a small model with dropout enabled, so an RNG or side-effect bug fails CI instead of a run.
Key takeaway: Activation checkpointing trades roughly one extra forward pass for keeping only block inputs alive. Estimate activations as sbh(34 + 5as/h) per layer, checkpoint only as many blocks as needed, prefer selective policies that save matmul outputs, pass use_reentrant=False, keep RNG state preserved, and judge the result on tokens per second and MFU rather than hardware utilisation.