Model and Inference Engineering · Principal
Ranks process different token counts. Why does averaging their losses change the objective?
The question
Interview question
Two training ranks each compute a mean loss over non-padding tokens. Rank A has 100 valid tokens and rank B has 900. Distributed training averages the two gradients equally. The intended objective was mean loss over all 1,000 tokens. Are those gradients the same?
Take a few minutes to form your approach. Then open a worked answer and compare the decisions.
Reveal a worked answer
No. Let the sum of token gradients on A be G_A and on B be G_B. Each local mean gives G_A/100 or G_B/900. Equal-rank averaging produces (G_A/100 + G_B/900)/2. The global token mean is (G_A + G_B)/1000. A token on the short rank receives nine times the weight of a token on the long rank in the equal-rank average. If rank losses happen to have similar gradients, the error may be hard to notice, but the objective changed. PyTorch's DistributedDataParallel documentation describes reduction of gradients across ranks. The framework cannot infer whether your loss was normalized by examples, sequences, valid tokens or some task-specific unit.
I would log the valid-token count and unnormalized loss sum by rank for every optimizer update. Decide the actual objective before fixing the math. For token-mean cross-entropy, each rank can backpropagate its local token-loss sum scaled by the global valid-token count, with the scaling adjusted for whether DDP itself averages by world size. In a simple two-rank DDP averaging setup, scaling the local sum by 2 / global_valid_tokens before backward gives the desired global gradient after averaging. With gradient accumulation, count valid tokens across all microbatches in the update, not just the current one. Other distributed frameworks may sum rather than average or apply their own loss scaling, so verify the actual reduction path.
Do not automatically assume token mean is right for every task. If the goal is equal weight per conversation or per document, a long sequence should not necessarily dominate by token count. State and implement that weighting deliberately. The danger is reporting “global batch size unchanged” while an incidental padding distribution makes the effective weighting move. Check packing boundaries, ignored labels, multimodal token definitions and any auxiliary MoE loss separately. They may require distinct denominators.
For a test, construct two tiny ranks with different valid-token counts and known gradients, then compare the distributed update to a single-process computation of the intended objective over their combined data. Repeat across a checkpoint resume and a changed number of ranks. We doubled gradient accumulation after losing GPUs. Why did the training recipe change? concerns accumulation changing the recipe, and The gradient norm is clipped on every microbatch. Why is the accumulated norm still too large? concerns when clipping applies to the accumulated gradient. This question is about what that gradient represents before clipping or optimization begins.
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 →