Model and Inference Engineering · Staff
Activation checkpointing fits the model. Why do its dropout gradients no longer match?
The question
Interview question
To fit a larger training batch, a team checkpoints some activations and recomputes them during backward. The forward loss is the same as before for a seeded microbatch. Gradients differ from the non-checkpointed reference and training becomes noisy. The checkpointed region contains dropout. What should recomputation replay?
Take a few minutes to form your approach. Then open a worked answer and compare the decisions.
Reveal a worked answer
Activation checkpointing trades stored activations for a second forward computation during backward. The recomputed values need to represent the same stochastic forward whose loss is being differentiated. If dropout draws a new mask during recomputation, backward uses a different function than the one that produced the visible loss. A seed set once at the start of the step is not enough if the RNG stream has advanced before recomputation. PyTorch's checkpointing documentation describes preserving RNG state by default to make checkpointed dropout behave like the non-checkpointed path, with caveats around devices and moving tensors.
I would compare one microbatch in a small deterministic setup, capturing dropout masks or the RNG state at the checkpoint boundary, forward activations, gradients and the parameters after a single update. Disable unrelated sources of nondeterminism for this diagnostic. Then reproduce it with the real distributed configuration. A rank-local RNG seed, model-parallel dropout stream and device transfer can complicate state restoration. The exact framework setting and checkpoint implementation matter, so inspect them rather than assuming all checkpointing libraries preserve RNG the same way.
The fix may be to preserve the relevant RNG state, use a stateless counter-based mask keyed to the same layer, token and step, or move stochastic operations outside the recomputed region where that is valid. If the team intentionally accepts different masks to save overhead, that is a changed training computation, not a free memory optimization. Compare convergence and quality under controlled runs. Measure the throughput cost of saving RNG state as well as the memory saving from checkpointing.
What if dropout is disabled? Then this particular explanation goes away. You still need to check custom ops with side effects, changing global state, non-deterministic kernels and tensor movement. Every training rank has a different shard. Why are its random augmentations identical? covers identical augmentations across ranks because seeds were reused. This question is the opposite temporal problem: the same sample's backward replay gets a different stochastic forward from the one whose loss it is differentiating.
Continue reading
Related questions
Read beyond the question
Explore more model and inference engineering
Follow another question in this area, or search the complete Question Library.
Browse this area →Browse Question Library →