Checkpointing asks which layers; selective asks which tensors

Classic gradient checkpointing partitions a network into segments, keeps only the segment boundaries live, and replays each segment’s forward pass during backward. The knob is where to cut, the answer is the sqrt(L) rule, and the price is roughly one extra forward pass — about 33% of a training step’s FLOPs.

That framing hides an assumption: that an activation is an activation. It is not. Inside one transformer layer the tensors autograd keeps differ by orders of magnitude in both bytes occupied and arithmetic needed to reproduce them. A LayerNorm output and a down-projection output can be the same shape, [s·b, h], and cost the same memory — while rebuilding one takes a handful of FLOPs per element and the other takes thousands.

Selective recomputation drops the segment abstraction and decides per tensor. That needs one number ranking tensors by how good a deal each is. That number is the whole subject.

Advertisement

The ratio: FLOPs paid per byte freed

Define, for a saved tensor T:

ρ(T) = (FLOPs needed to rebuild T) / (bytes freed by dropping T)

units: FLOPs per byte

Low ρ is a bargain: you hand back a lot of memory for very little arithmetic. High ρ is a bad trade. Because every candidate is measured in the same units, ρ gives a total ordering, and the optimal policy under a memory budget is simply: sort ascending by ρ, drop tensors from the front until you fit, stop.

Two things make this more than bookkeeping. First, ρ has a closed form for the operators a transformer is built from — you compute it symbolically rather than profile it. Second, the values are not spread smoothly; they cluster into two tight groups separated by a factor of a hundred or more. When a ranking has a gap that wide you do not need an optimizer, you need a threshold. (ρ is a close cousin of arithmetic intensity, which counts all traffic in and out of an op; ρ counts only the bytes you stop storing.)

Advertisement

Why ρ is the producing op’s FLOPs per output byte

Rebuilding T means re-running the op that produced it, so ρ is that op’s FLOP count divided by the bytes of its output. Do this for a matmul. A projection X·W with X: [N, k] and W: [k, d] costs 2Nkd FLOPs and emits N·d elements, or 2Nd bytes in bf16:

ρ_matmul = 2Nkd / (2Nd) = k       FLOPs per byte

  QKV / attn-out projection:  k = h    = 768
  FFN up-projection:          k = h    = 768
  FFN down-projection:        k = 4h   = 3072

The N and d cancel: a matmul’s ρ is just its reduction depth. Now do an elementwise op. GELU on [N, 4h] costs on the order of 8 FLOPs per element and frees 2 bytes per element, so ρ ≈ 4. RMSNorm or LayerNorm, two passes plus an affine, lands near ρ ≈ 4 as well. Dropout and softmax sit in the same single-digit neighbourhood.

So 768 versus 4 — a factor of roughly 190. That gap, derived rather than asserted, is the selective-recompute argument.

Ranking the transformer’s tensors

Applying the formula across one pre-norm layer at h = 768, a = 12 heads, head dim d_h = 64:

Saved tensorProducing opρ (FLOPs/byte)Verdict
Dropout output / maskelementwise~1recompute
Softmax probabilitiesrow softmax~3recompute
Norm outputRMSNorm / LayerNorm~4recompute
GELU / SwiGLU outputelementwise~4recompute
Attention scores QK^Tbatched matmuld_h = 64recompute
Q, K, V projectionsmatmulh = 768store
FFN up-projection outmatmulh = 768store
Attention output PVbatched matmuls = 2048store
FFN down-projection outmatmul4h = 3072store

The interesting row is QK^T. It is a matmul, yet its ρ is only d_h = 64, because attention reduces over the head dimension, not the model dimension — while its output is the layer’s largest tensor, scaling as s². Biggest bytes, shallowest reduction: the best trade in the network, and one no layer-granularity policy can express.

Chains and anchors: you free a span, not a tensor

A subtlety that per-tensor ranking hides: you cannot drop an elementwise output in isolation and expect to rebuild it for 4 FLOPs per byte, because rebuilding it requires its input to still be live. If that input was also dropped, the true cost is the whole chain back to the nearest surviving tensor.

So the real unit of decision is a span between two anchors. Pick tensors to keep (the anchors), and everything strictly between two anchors is free to discard; the cost is the span’s total forward FLOPs and the saving is every intermediate inside it. Evaluate ρ over the span:

anchors: Q, K, V (stored)
span:    S = QK^T/√d_h  →  P = softmax(S)  →  P’ = dropout(P)

cost  ≈ 2bas²d_h  +  ~5bas²  +  ~2bas²      (matmul + softmax + mask)
freed ≈ 3 · 2bas²                             (three s² tensors, bf16)

ρ_span ≈ (2·64 + 7) / 6 ≈ 22 FLOPs per byte

Still an order of magnitude under any projection matmul. Anchoring on Q, K and V is what makes the span cheap: they are small, linear in s, and they cut the chain exactly where the quadratic tensors begin.

Worked example: a 12-layer CPU-trainable SLM

Take L = 12, h = 768, a = 12, s = 2048, b = 1, bf16. Two unit sizes:

linear unit  sbh   = 2048 · 1 · 768   = 1.57M elems  →  3.15 MB
quad   unit  bas²  = 1 · 12 · 2048²  = 50.3M elems  →  100.7 MB

per layer:  ~17 linear-unit tensors  →   53.5 MB
            3 quadratic tensors      →  302.0 MB
whole stack (×12):                      4.27 GB   (642 MB + 3.62 GB)

The s² tensors are 85% of activation memory. Now the bill for the attention span, per layer:

span FLOPs ≈ 2 · 50.3M · 64 + 7 · 50.3M ≈ 6.8 GFLOP
layer fwd  ≈ 24sbh² + 4bas²d_h ≈ 29.0 + 12.9 = 41.9 GFLOP
step (fwd+bwd) ≈ 3 × fwd ≈ 125.7 GFLOP per layer

overhead = 6.8 / 125.7 ≈ 5.4%

3.6 GB removed for a 5% slowdown. Full checkpointing removes the last 642 MB too — and charges 33%.

The frontier: greedy on ρ is steep, then flat

Sort every candidate span by ρ and plot cumulative memory freed against cumulative FLOPs paid. Because you spend the cheapest bytes first, the curve is concave: enormous early returns, then a long flat tail. For the example above, three tiers:

TierρMemory freedStep overhead
Norms, activations, dropout~1–4~0.4 GB< 0.2%
Attention s² span~22~3.6 GB~5%
Projection outputs768–3072~0.25 GB~28%

The first two tiers deliver about 94% of the achievable saving for about 5% of a step. The third tier costs five times as much compute as the first two combined and buys a rounding error. That is not a close call, and it is why the correct default is not “checkpoint everything” but “checkpoint down to ρ ≈ 50 and stop.”

Boundaries move with shape: as s grows the quadratic tier swells while its ρ stays pinned at d_h.