Model and Inference Engineering · Staff
Activation checkpointing saved GPU memory. Why did FSDP training slow down?
Take a few minutes to form your approach. Then open a worked answer and compare the decisions.
Reveal a worked answer
That memory saving has a bill. Instead of retaining every intermediate activation from forward, checkpointing stores selected boundaries and recomputes work during backward. FSDP also shards parameters and gathers the needed weights around computation. The combination can add recompute, extra parameter materialization or communication pressure, depending on wrapping, resharding and checkpoint boundaries. PyTorch's FSDP and activation-checkpointing discussion describes the extra backward computation, while FSDP documentation explains when unsharded parameters are freed and gathered again.
“We fit the model now” and “we train more tokens per second” are different wins. If the freed memory allows a larger microbatch, longer sequence or a model that previously could not run, recomputation may be worth it. If the batch remains unchanged, the same job can simply do more forward work per optimizer step. Broad checkpoint boundaries may also reduce overlap between all-gathers and useful compute.
I would profile one step at a fixed effective token count with and without checkpointing. Break down forward, recompute, backward, all-gather wait and optimizer time. Track peak allocated and reserved GPU memory. Then try the microbatch or sequence length the memory saving was intended to enable and compare useful tokens per second at equal quality settings. A profiler trace can tell us whether the issue is extra FLOPs, repeated communication, or a poor FSDP wrapping choice. GPU utilization alone cannot make that distinction.
The boundary matters. Checkpoint every tiny operation and the overhead grows. Checkpoint too little and memory stays high. Align checkpoint regions with sensible transformer blocks and understand whether the FSDP implementation reshares parameters after forward, because that determines what needs to be gathered again. Do not assume every checkpointed recomputation forces exactly one extra all-gather in every implementation.
If someone says the run is now slower so checkpointing failed, I would ask what constraint we were solving. Training a larger model or doubling useful tokens per step may justify a slower individual step. State the objective and compare the full training recipe, not the isolated memory number.
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 →