Model and Inference Engineering · Staff
More pipeline stages fit the model. Why did training tokens per second fall?
The question
Interview question
A team moves from four to eight pipeline stages to fit a larger model without increasing per-GPU memory. The global batch and sequence length stay fixed, and tokens per second falls. GPU graphs show idle gaps even though each stage runs quickly when it has work. An engineer proposes another four stages because the layers will be smaller. What is missing from that argument?
Take a few minutes to form your approach. Then open a worked answer and compare the decisions.
Reveal a worked answer
The stages form a pipeline only while there are microbatches in flight. At the start, later stages wait for the first microbatch. At the end, earlier stages wait while the last backward passes finish. We flush before the next optimizer step to preserve the intended batch update. Splitting layers across more GPUs reduces work and memory per stage, but it also deepens the pipeline that has to fill and drain.
For a simple balanced flush schedule with p stages and m microbatches per optimizer step, the extra bubble time relative to the ideal busy time is (p − 1) / m. Equivalently, ignoring communication and imbalance, useful utilization is roughly m / (m + p − 1). With four microbatches and eight stages, that idealized utilization is 4 / 11, only about 36%. This is a simplified scheduling estimate, not a prediction of real GPU utilization. NVIDIA's Megatron explanation lays out the pipeline fill and drain and the 1F1B and interleaved schedules. The common mistake is to call (p − 1) / m a percentage of actual wall time even when it exceeds one. Its denominator there is the ideal compute time.
I would draw the measured schedule from traces, not infer it from a single throughput number. Count microbatches per step, forward and backward time per stage, point-to-point transfers, tensor parallel collectives, activation recomputation, and time waiting at the flush boundary. Find the slowest stage. If stage 5 contains an expensive embedding or MoE layer, balancing by layer count will not balance work. A scheduler can also expose extra idle time from uneven sequence lengths or communication across slower interconnects.
Adding microbatches can amortize the bubble, but ask what stays fixed. If the global batch remains constant, making each microbatch smaller can turn efficient matrix multiplies into tiny ones and change communication overhead. If we simply increase the number of same-sized microbatches, we change the global batch or accumulation schedule, and that can change optimization. Interleaved virtual stages can reduce bubble at the cost of more communication. Activation recomputation may recover memory to allow a better microbatch shape while adding compute. Sometimes a different tensor, sequence or data parallel split gives a better end-to-end balance. These are measurements, not a rule to always use fewer pipeline stages.
The interviewer may say, “But we only care about fitting the larger model.” Fair, fitting is a hard constraint. The decision then is the fastest valid configuration under that constraint, with an acceptable training recipe. I would benchmark several p and m pairs at the actual sequence lengths, report tokens per second per GPU and total step time, and verify the same loss semantics and checkpoint behavior. A configuration that saves memory but leaves most stages empty is usable only if no better feasible partition exists.
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 →