Model and Inference Engineering · Staff
The all-reduce was launched. Why did one rank step before it finished?
The question
Interview question
A custom distributed trainer starts asynchronous gradient all-reduce, then calls `optimizer.step()` while communication is still in flight. The intent was to overlap network work with computation. Training sometimes diverges, especially under network jitter. Every rank did call all-reduce, so what is wrong?
Take a few minutes to form your approach. Then open a worked answer and compare the decisions.
Reveal a worked answer
Launching a collective is not the same as having a complete global gradient available for the update. The optimizer must consume the reduced gradient after the work has completed and any required stream synchronization has made it visible. PyTorch's distributed documentation distinguishes asynchronous work handles and waiting for completion. A manual trainer that reads or mutates the gradient buffer while the collective may still write it has a race. Timing can make the bug intermittent and hard to catch with a small local test.
I would trace one parameter through backward, gradient buffer production, collective enqueue, completion, clipping, unscaling and optimizer step. Verify that every rank uses the same reduction group and scaling convention. Test with two ranks whose local gradients are deliberately different, and delay one rank's communication to force the race. Compare the actual update with a synchronous reference. Do not treat identical loss after the first forward as evidence because the bad operation affects the next weights. Check whether optimizer.step() zeros or reuses the buffer before all-reduce completes.
Overlap is still possible. Launch reductions for buckets whose gradients are ready while backward computes other buckets, then make the optimizer wait on all dependencies it needs. DistributedDataParallel already coordinates gradient reduction, so a custom overlap path should be justified and tested against its semantics. A separate CUDA stream can require an explicit event or stream wait even if a host-side work handle returned. The exact synchronization rule depends on backend and framework. “I called wait() somewhere” is not enough if it happened after the optimizer consumed the data.
The interviewer may ask how this relates to Ranks process different token counts. Why does averaging their losses change the objective?, where different token counts make simple rank-loss averaging wrong. That answer can have a correctly synchronized gradient for the wrong objective. Here the desired objective can be correct, but the reduction result is not ready at the moment of update. The fix is a happens-before edge between complete reduction and the optimizer's read, not a new learning-rate schedule.
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 →