skip to content

Memory Budget

What actually fills the device during a training step, and the two levers that buy room back: accumulating gradients over micro-batches and recomputing activations. Interviewers love this arithmetic.

on this pageshow

explore

questions

14

What does activation checkpointing trade away to cut a training step's activation memory?

level: juniorimportance: must knowfreq 58%

answer

  1. one resource bought with another
  2. keep the boundaries, drop the middle
  3. rebuild what was dropped in backward
  4. about one extra forward per step

basics

~20 s

Activation checkpointing trades compute for memory. Instead of keeping every layer's forward outputs alive until the backward pass, it keeps only a few segment boundaries and recomputes the rest on demand, costing roughly one extra forward pass per step.

solid answer

~50 s

By default every intermediate tensor produced in the forward pass has to stay resident until the backward pass reaches it, so peak activation memory grows with depth and with batch size, sequence length or resolution. Activation checkpointing cuts the network into segments, keeps only the tensor entering each segment, and discards the segment's interior. When the backward pass arrives, that segment is run forward again from its saved input to rebuild the interior values, the local gradients are formed, and the rebuilt values are freed immediately. Only one segment's interior is live at a time. The price is roughly one extra forward pass over the checkpointed region per step — about a third more compute, since a backward pass costs roughly twice a forward one. The gradients are mathematically unchanged as long as recomputation is deterministic.

go deeper

for a junior

Be ready to say in one breath what is traded for what: memory is saved, extra compute is spent, and the model's mathematics is unchanged. Also be ready to distinguish it from saving model checkpoints to disk.

for a middle

Explain the mechanics — segment boundaries retained, interiors discarded and rebuilt during the backward pass — and produce the cost estimate of roughly one extra forward pass, about a third more compute per step.

for a senior

Show you know when it is the wrong lever: a compute-bound step with headroom to spare, or a model that already fails at batch size one. Be ready to say how you measured the memory profile before turning it on.

for a principal

Own the throughput argument: recompute costs a fixed fraction of step time, and the case for it is that the batch size it unlocks earns more efficiency back than it spends. Frame it against buying more devices or shrinking the model.

## What actually fills memory during a forward pass When a network runs forward, each layer produces an output tensor, and reverse-mode automatic differentiation needs many of those intermediate values a second time to form local gradients on the way back. The backward pass walks the graph in reverse, so the values produced earliest are needed last. By default they all sit in device memory from the moment they are produced until the backward pass consumes them. Peak activation memory therefore grows with depth and with everything that multiplies the size of a single intermediate: batch size, sequence length, image resolution, channel width. This is the term that usually pushes a training job over the device limit, because unlike the weights it scales with the batch. ## The idea Activation checkpointing — also called gradient checkpointing or rematerialization — refuses to hold that whole stack. The network is cut into segments. Only the tensor entering each segment is retained; everything computed inside a segment is released as soon as the segment's output has been produced. When the backward pass reaches a segment, the segment is executed forward a second time from its retained input, the interior values are reconstructed, the gradients for that segment are computed from them, and the reconstructions are dropped again before moving to the next segment. At any instant only the small set of boundary tensors plus one segment's interior is alive. The technique is trading a resource that is elastic (compute time) for one that is a hard ceiling (device memory). Running slower is a decision; running out of memory is a crash. ## Counting the cost A good accounting rule: if a forward pass costs F, the backward pass costs roughly 2F, because at most layers it forms both a gradient with respect to the layer's input and a gradient with respect to its parameters. An ordinary step is therefore about 3F. Checkpointing the whole network adds one more forward, giving about 4F — roughly a third more compute per step. This is why the technique is usually quoted as costing about thirty percent of step time when everything is checkpointed; checkpointing only part of the network costs proportionally less. Wall-clock time often moves by less than the compute ratio, for a practical reason: teams rarely enable checkpointing and keep the same batch size. The freed memory is spent on a larger batch, which usually improves device efficiency, so some of the recompute cost is earned back as better utilization. ## What it does not help Checkpointing touches the activation term only. Parameters, their gradient buffers and optimizer state are untouched, because none of them is a recomputable intermediate — they persist across the whole step by definition. If a model cannot fit with a batch of one and no activations at all, checkpointing cannot rescue it; that situation needs a different lever such as sharding or reducing the model itself. It also does not reduce the number of parameters, change the learning rate that is appropriate, or alter what the optimizer does. ## Correctness Recomputation reproduces the same values from the same inputs, so the gradients handed to the optimizer are the ones an un-checkpointed run would produce. The exception is any layer whose forward pass is not a pure function of its input — a layer that draws random values, or one that mutates internal state as a side effect of running forward. Those need explicit care so the second forward reproduces the first; otherwise the recomputed values do not match the ones that produced the loss. ## A naming trap worth pre-empting Activation checkpointing has nothing to do with saving model checkpoints to disk for fault tolerance or for resuming a run. The word is shared, the mechanism is not: one is about which intermediate tensors stay in device memory inside a single step, the other is about periodically persisting parameters. Interviewers ask this deliberately, because a candidate who has only read the flag name often conflates the two. ## When to reach for it Use it when activations dominate the memory profile and the model is deep enough that a segment interior is much smaller than the whole stack — long-sequence encoders, high-resolution vision models, anything where the useful batch size is being squeezed to one or two. Skip it when the step is already compute-bound and memory is comfortable, because you would be paying a third of your throughput for room you do not need.

  • Does activation checkpointing change the gradients the optimizer receives?
    Not if recomputation is deterministic — the rebuilt intermediates are the same values the first forward produced, so the gradients match an un-checkpointed run. The exception is a layer whose forward pass is not a pure function of its input, such as one drawing random values or updating internal state, which must be made to reproduce itself on the second forward.
  • Why is the overhead roughly a third rather than double the step time?
    Because the extra work is one forward pass, not a whole step. A backward pass costs about twice a forward pass, so an ordinary step is about three forward-equivalents; adding one recompute makes four. That is a ratio of about 4/3, and it shrinks further if only part of the network is checkpointed.
  • Can activation checkpointing help a model that will not fit even at batch size one?
    Usually not. At batch size one the activation term is already small and the remaining pressure comes from parameters, gradients and optimizer state, none of which checkpointing removes. That case calls for a different lever — sharding state across devices, a smaller model, or lower-precision storage of the persistent tensors.

Like not photographing every step of a recipe: you keep a photo at the start of each stage, and if you need to inspect a middle step you just cook that stage again from its saved starting point.

saying these in an interview costs you the question

  • Confuses it with saving model checkpoints to disk
  • Claims it reduces parameter or optimizer-state memory
  • Says the memory saving is free of compute cost
  • Believes it changes the gradients or the loss by design
  • Thinks the activations are written to disk rather than recomputed

context

open as a page

Why does halving the training batch size cut device memory when the weights are unchanged?

level: juniorimportance: must knowfreq 78%

basics

~20 s

Weights, gradients and optimizer state are sized by the parameter count, not by batch size. Activations, the per-layer outputs held for the backward pass, exist once per sample, so halving the batch roughly halves that share.

open as a page

In gradient accumulation, why is each micro-batch's loss divided by the number of micro-batches?

level: middleimportance: must knowfreq 68%

basics

~20 s

Because each micro-batch's loss is usually already averaged over its own samples, so adding N of them gives N times the true full-batch gradient. The extra 1/N factor restores the mean over all samples and keeps the step size honest.

open as a page

Your training job reports 22 GB reserved but 9 GB allocated, then fails to allocate — why?

level: middleimportance: must knowfreq 60%

basics

~20 s

Reserved is what the caching allocator holds from the driver; allocated is what live tensors use. The gap is cached blocks of the wrong sizes, and a new tensor needs one contiguous block none of them can supply.

open as a page

How much device memory does an Adam-trained 1.5B-parameter model's state need before any activations?

level: middleimportance: must knowfreq 66%

basics

~20 s

About 24 GB. Adam training in 32-bit floats costs 16 bytes per parameter - weight, gradient, and two moment estimates at 4 bytes each - so 1.5 billion parameters need roughly 24 GB before any activation exists.

open as a page

What does gradient accumulation let you do when only a few samples fit in device memory?

level: juniorimportance: should knowfreq 55%

basics

~20 s

Gradient accumulation runs several small micro-batches, adding their gradients into one buffer, and updates the weights only after the whole group. You get the update of a large batch while only one micro-batch's activations sit in memory at a time.

open as a page

How do you choose activation-checkpoint segment boundaries in a deep network?

level: middleimportance: should knowfreq 42%

basics

~20 s

Place boundaries so retained memory balances: with n uniform layers in k segments it scales like k + n/k, smallest near k equal to the square root of n. When layers differ, rank candidates by bytes freed per recomputed operation.

open as a page

Why can activation checkpointing around a stochastic layer corrupt a training run?

level: seniorimportance: should knowfreq 34%

basics

~20 s

The recomputed forward draws fresh random numbers, so the rebuilt activations differ from the ones that produced the loss and the gradients belong to a different sampled network. Fix it by capturing and restoring the generator state per segment.

open as a page

Why must global-norm gradient clipping run after the last micro-batch of an accumulation group?

level: seniorimportance: should knowfreq 45%

basics

~20 s

Clipping inside the loop rescales each micro-batch gradient separately, so the sum of clipped pieces is not the clipped full-batch gradient. Clip once after the final backward pass and before the step, so the threshold sees the real update.

open as a page

A training run OOMs at step 8,000 with no code change — leak or peak, and how do you tell?

level: seniorimportance: should knowfreq 52%

basics

~20 s

Log memory per step. Steady growth with a constant increment means retention, usually a logged metric still attached to its graph. Flat memory with one tall spike means a peak: the longest sample, or the evaluation pass.

open as a page

How do you estimate whether a 3D segmentation run on 256-cubed volumes fits in 40 GB?

level: seniorimportance: should knowfreq 45%

basics

~20 s

Budget two piles. Model state, parameter count times bytes per parameter, is often under a gigabyte here. Activations dominate: one 256-cubed feature map with 32 channels in half precision is about 1 GB, and several stay resident.

open as a page

Step time doubled after offloading stored activations to host memory — what do you check?

level: seniorimportance: nice to knowfreq 24%

basics

~20 s

Check whether the bytes moved each step exceed what the host link can carry in the time available. Offloading pays only when a round trip is faster than recomputing the values and overlaps with compute.

open as a page

Why do batch-statistics normalization layers not see the full effective batch under gradient accumulation?

level: seniorimportance: nice to knowfreq 32%

basics

~20 s

Each micro-batch is normalized with statistics computed from its own samples alone, because normalization happens in the forward pass, before anything is accumulated. Accumulation enlarges the update, not the sample set a forward pass can see.

open as a page

A training job OOMs the day before a deadline — in what order do you apply the memory levers?

level: principalimportance: nice to knowfreq 36%

basics

~20 s

Measure the true peak first, then go cheapest-to-riskiest: shrink the micro-batch, restore the effective batch with gradient accumulation, add activation recomputation, drop precision, then shard state across devices. Anything that changes what the model sees comes last.

open as a page