Why can activation checkpointing around a stochastic layer corrupt a training run?
answer
- recompute assumes a pure function
- the second pass draws different values
- gradients no longer match the loss
- snapshot and restore, not a global seed
basics
~20 sThe 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.
solid answer
~50 sCheckpointing assumes a segment's forward pass is a pure function of its input, so replaying it reproduces the same values. A layer that draws random values breaks that assumption: the first forward uses one draw, the recompute during backward pulls the next values from the generator stream and produces different intermediates. The gradients are then computed against activations the loss never saw — effectively gradients of a different sampled network, which injects noise, breaks the match with an un-checkpointed baseline and makes runs hard to reproduce. The remedy is to snapshot the generator state on entry to each segment and restore it before the recompute, so the second pass draws identical values. The same discipline applies to any forward pass with side effects — anything that mutates internal state while running forward would apply that mutation twice per step.
go deeper
Know that recomputing a segment assumes replaying it gives the same values, and that a layer drawing random numbers breaks that assumption unless the generator state is saved and restored.
Explain the mechanism: the generator advances between the two forwards, so the rebuilt intermediates differ from the ones that produced the loss and the gradients no longer correspond to it.
Describe the symptom as a silent quality regression against an un-checkpointed baseline, name the state capture-and-restore fix, and give the diagnostic run that proves it.
Generalise it to a review rule: any forward-pass side effect executes twice under recomputation. Decide whether the team enforces purity in checkpointed regions or accepts case-by-case suppression.
## The assumption checkpointing rests on Activation checkpointing discards a segment's interior and rebuilds it later by running the segment forward a second time from its retained input. That is only sound if the segment is a pure function: same input, same parameters, same output. Every deterministic layer satisfies this. Anything that consumes randomness or mutates state as a side effect of running forward does not. ## What goes wrong with a stochastic layer A layer that samples random values during training draws from a generator whose state advances with every draw. On the first forward pass through the segment it draws one set of values. That set shapes the intermediates, which shape the loss the optimiser is about to differentiate. Later, when the backward pass reaches the segment and triggers the recompute, the generator has moved on — every other draw in the step has consumed from the same stream — so the second pass draws a different set. The rebuilt intermediates are the activations of a *different* sampled network. The gradients then formed are not the gradients of the loss that was actually computed. They are a mismatched pair: a loss from one sample, derivatives from another. Nothing raises an error, because every tensor has a valid shape and a finite value. ## How it shows up The symptoms are quiet and easy to misattribute: - The checkpointed run no longer matches an un-checkpointed run with the same configuration, even though checkpointing is supposed to be mathematically neutral. - Training is noisier than the baseline, and convergence is slower or reaches a worse plateau, because every step adds a mismatch term that does not average to a useful direction. - Effects scale with how much randomness the segment contains: a model with stochastic behaviour in every block is hit far harder than one with a single stochastic layer near the head. - Small models often survive it, which is what makes it dangerous — the bug is discovered only after it is scaled up. A candidate who says the run will crash has not thought it through. Silent quality loss is the whole problem. ## The fix Capture and restore the generator state. On entry to a checkpointed segment, record the state of every generator that segment will consume from. Before recomputing that segment, restore the recorded state so the second pass draws exactly the values the first one drew. The recompute then reproduces the original intermediates bit for bit, and the gradients match the loss again. This is per-segment bookkeeping, not a global seed. Setting one seed at the start of training does not help at all: both forwards still consume from the same advancing stream, just from different positions in it. Nor is disabling stochastic behaviour during the recompute a fix — that reproduces a different network again, this time a deterministic one, and biases the gradients in the opposite direction. The requirement is exact reproduction of the first pass, and only saving and restoring the state achieves it. Note what this bookkeeping costs: a little state per segment and one restore per recompute. It is negligible against the recompute itself, which is why correct implementations simply always do it. ## The wider class: side effects in the forward pass Randomness is the loudest member of a family. Anything a layer *does* while running forward, beyond producing its output, will happen twice under checkpointing: - A layer that maintains a running statistic updated on each forward pass will update it twice per step, so the statistic moves at double the intended rate and drifts from the value an un-checkpointed run would hold. - A counter, a cache, a queue written during the forward pass gets two writes per step. - Any per-sample statistic computed inside the segment must be reproducible from the segment input alone; if it is, the recompute regenerates it exactly and nothing is lost, which is the good case. The general rule for reviewing a checkpointed model: for each segment, ask what the forward pass does other than return a value. If the answer is anything at all, that side effect now happens twice, and either it must be made idempotent, suppressed on the recompute pass, or the boundary must be redrawn so the offending layer sits outside the checkpointed region. ## Interview framing This is a favourite senior question because it separates people who have read the flag description from people who have debugged a training run. The tell of the second group is that they describe the symptom as a silent quality regression against a baseline, name generator state as the mechanism, and then generalise to forward-pass side effects without being prompted.
- Would setting a single global seed at the start of training fix this?No. Both forwards still read from the same advancing generator stream, only at different positions, so they draw different values. A fixed seed makes the whole run repeatable end to end but does nothing to make a segment's recompute match its original forward. Only per-segment state capture and restore does that.
- What happens to a layer that updates a running statistic during the forward pass?It updates twice per step under checkpointing, once on the original forward and once on the recompute, so the statistic advances at double the intended rate and diverges from what an un-checkpointed run would hold. Either suppress the update on the recompute pass or move the layer outside the checkpointed segment.
- How would you detect this bug on a model already in training?Run a short comparison: identical configuration and data order, one run with checkpointing and one without, and compare per-step losses. A correct implementation matches closely; a re-randomising one diverges immediately even though neither run errors. A single-step gradient comparison is the sharpest version of the same check.
saying these in an interview costs you the question
- Says the run will crash rather than silently degrade
- Believes one global seed makes both forwards draw identically
- Suggests disabling the stochastic layer during recompute
- Assumes checkpointing is always mathematically neutral
- Ignores forward-pass side effects that now run twice