How do you choose activation-checkpoint segment boundaries in a deep network?
answer
- two competing memory terms
- boundaries kept plus one live interior
- k + n/k, and where it is smallest
- for uneven layers, bytes freed per rerun operation
basics
~20 sPlace 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.
solid answer
~50 sFor a stack of `n` similar blocks split into `k` equal segments, you retain `k` boundary tensors plus one live interior of about `n/k` blocks, so retained memory scales like `k + n/k`. That is minimised at `k` near the square root of `n` — a 48-block encoder wants roughly seven segments of seven blocks. The curve is flat near the optimum, so a rounder grid such as every fourth block is fine in practice. The recompute bill is one extra forward either way; `k` only shifts where the memory sits. Once blocks stop being uniform, stop using the heuristic and rank segments by bytes saved per recomputed operation: checkpoint the parts that hold large intermediates and are cheap to rerun, and leave alone the parts whose outputs are tiny, where you pay a recompute to free almost nothing.
go deeper
Know that the network is cut into segments and only segment inputs are kept, and that roughly the square root of the depth is a sensible number of segments for a uniform stack.
Derive it: k boundaries plus one interior of n/k gives k + n/k retained units, minimised near k equal to the square root of n. Be able to apply that number to a concrete block count.
Show you profile first. Explain ranking candidate segments by bytes freed per recomputed operation on a non-uniform model, and why boundaries belong at narrow points in the graph.
Frame segment count as a memory-shape decision with a fixed recompute bill, and set the stopping rule from the batch size the training plan needs rather than minimising memory for its own sake.
## The uniform case, and where the square root comes from Take a network that is a stack of `n` near-identical blocks, each producing an intermediate of about the same size. Cut it into `k` equal segments of `n/k` blocks each. What stays resident? First, the tensor entering each segment: that is `k` tensors, alive for the whole step because each is the starting point for a later recompute. Second, while the backward pass is working on one segment, that segment's interior is reconstructed and live: about `n/k` tensors. Everything else is gone. So retained activation memory scales like `k + n/k` units instead of `n`. Minimising `k + n/k` over `k` gives `k = sqrt(n)`, and the minimum value is `2*sqrt(n)`. That is the classic sublinear-memory result: memory that grew linearly with depth now grows with the square root of depth, for the price of one extra forward pass. For a 48-block text encoder, `sqrt(48)` is about 6.9, so the heuristic says roughly seven segments of seven blocks, retaining about fourteen block-intermediates instead of forty-eight. ## Why the exact number matters less than people expect The function `k + n/k` is flat near its minimum. Checkpointing every fourth block of that same 48-block encoder gives twelve segments of four: twelve boundaries plus a four-block interior, about sixteen units against the optimum's fourteen. A few percent of the activation budget, in exchange for a boundary grid that lines up with the architecture and is easy to reason about. The more important invariant is that `k` does not change the recompute bill at all. Every checkpointed block is recomputed exactly once per step regardless of how the segments are drawn, so the time cost is fixed and `k` only decides how the retained memory is split between many small boundary tensors and one large live interior. Choosing `k` is a memory-shape decision, not a speed decision. ## When the layers are not uniform Real networks are not stacks of identical blocks. A text model has an embedding table lookup, a stack of transformer-style blocks, a pooling step and a head. A vision model has high-resolution early stages holding enormous feature maps and cheap low-resolution late stages. The square-root heuristic silently assumes every block holds the same bytes and costs the same to rerun, and neither is true here. The right ranking is bytes saved per unit of recomputed work. For each candidate segment, ask how many bytes of intermediates you stop retaining, and how many operations you must rerun to get them back. Sort by that ratio and checkpoint from the top until the memory target is met. This ranking produces two clear rules. Checkpoint the parts that hold large intermediates but are cheap to reproduce — elementwise activations, normalisation outputs, high-resolution feature maps early in a vision stack. Do not checkpoint parts whose retained output is already tiny: a pooling layer that collapses a long sequence to a single vector stores almost nothing, so dropping it frees almost nothing while still forcing everything before it in that segment to rerun. ## The failure mode this prevents The common bad configuration is a segment grid drawn where the code was easiest to edit rather than where the bytes are. A team wraps the embedding lookup and the pooling head — the two cheapest, smallest parts of the model — leaves the expensive attention-and-feedforward blocks resident, and reports that checkpointing did not help. It did exactly what it was asked to do; it was pointed at the wrong tensors. The diagnostic is always the same: measure where the activation bytes actually are before drawing any boundary. A second failure mode is drawing boundaries mid-block, so a segment's retained input is the widest tensor in the model rather than the narrowest. Good boundaries sit at natural narrow points in the graph — a block's residual input, a stage transition where resolution drops — because the boundary tensor is exactly what you have to keep. ## The practical procedure Profile the activation footprint per stage. Identify the narrow points in the graph as candidate boundaries. If the stack is uniform, start near the square-root count and round to something architecturally natural. If it is not uniform, rank by bytes freed per recomputed operation and checkpoint down that list only until the batch size you want fits — every extra segment past that is throughput you spent for nothing.
- Does using more checkpoint segments make the recompute cost higher?No. Every checkpointed block is recomputed exactly once per step whatever the segment count, so the time cost is fixed. The segment count only decides how retained memory is split between many boundary tensors and one live interior — which is why the choice is a memory-shape decision, not a throughput one.
- Where in the graph should a checkpoint boundary physically sit?At a narrow point, because the boundary tensor is the thing you keep. A block's residual input or a stage transition where spatial resolution drops are good boundaries; slicing through the widest intermediate inside a block is the worst place, since you retain the largest tensor and still rerun everything around it.
- How do you decide when to stop adding segments?Stop as soon as the batch size or sequence length you actually want fits with headroom. Recompute is paid per checkpointed block, so any segment added beyond the memory target buys nothing and costs throughput. Treat the memory target as the stopping rule, not the memory minimum.
saying these in an interview costs you the question
- Thinks more segments means proportionally more recompute time
- Applies the square-root rule to a network with very uneven layers
- Checkpoints the cheapest, smallest layers and reports no benefit
- Places boundaries at the widest tensor in a block
- Checkpoints everything by default instead of to a memory target