Distributed training must preserve the intended average
Calculate the difference between rank means and a sample-weighted mean.
Distributed data parallel training combines gradients across processes. The intended objective usually averages loss over a global set of examples or tokens. That objective is not always the same as averaging each rank's local mean. Unequal batch sizes or variable valid-token counts can change the weighting.
Write the mathematical denominator before configuring the collective. If every rank has the same number of equally weighted examples, averaging local mean gradients matches the global mean. If one rank has one valid example and another has three, equal rank weighting gives the small rank too much influence. Variable-length sequence losses often encounter this through padding masks.
Framework behavior matters. DDP synchronizes gradients according to its contract, while the loss reduction happens in user code. Gradient accumulation adds another denominator: microbatches can have unequal sizes, and the last accumulation window may be smaller. Do not divide by a fixed accumulation count without checking the actual effective sample or token count.
Also verify data partitioning. Each rank should receive the intended shard, and epoch changes should produce the intended shuffle. Duplicate samples or dropped tails may be acceptable under an explicit protocol, but they change the dataset exposure. A distributed run that trains faster on fewer effective examples is not an equal-work speed comparison.
Worked example
A teaching setup has two ranks. Rank A has one example with gradient 4. Rank B has three examples with gradients 0, 0, and 0. The desired global example mean is 4/4 = 1. The local means are 4 and 0. Averaging rank means produces 2, which is wrong for the stated objective.
One valid approach computes local gradient sums and normalizes by the global example count under the framework's reduction semantics. If the framework averages gradients across two ranks, compensate for that averaging in the loss scaling. The exact implementation must be derived from the chosen API, not copied from this arithmetic without checking its contract.
Exercise and solution
Rank A has two valid tokens with loss-gradient sum 6. Rank B has four valid tokens with sum 6. Calculate the global token mean and the average of local means.
The global mean is 12/6 = 2. Local means are 3 and 1.5, whose average is 2.25. Award one point each for these values, identifying the denominator mismatch, and proposing a global valid-token count. Include a single-device reference test using the same six tokens. This test checks objective equivalence before performance comparisons.
Lab artifact: derive compensation for DDP averaging
Assume R ranks, a framework that averages parameter gradients across ranks, and a desired mean over N globally valid tokens. Let local rank r compute differentiable loss sum L_r. Set its local scalar to R × L_r / N before backward. The framework average then gives (1/R) × sum_r [R × gradient(L_r) / N], which equals the desired global gradient sum divided by N.
This sketch omits gradient reset, mixed-precision handling, and accumulation details to expose the denominator. The count is data, not a differentiable model output. A rank with zero valid tokens still needs to participate in the required collective and backward protocol with a compatible graph; skipping backward on just that rank can hang or violate synchronization expectations. Use the framework's supported handling rather than copying the sketch as production code.
If the framework sums gradients instead of averaging, the world-size factor is wrong. If a communication hook changes the reduction contract, derive the factor again. DDP's documented behavior is the authority for a specific implementation. The algebra here is an original derivation under the stated averaging assumption.
A second failure case: unequal accumulation windows
Suppose two microbatches contain one and three valid tokens. Their loss-gradient sums are four and six. The global token mean over the effective batch is ten divided by four, or 2.5. Averaging microbatch means gives (4 + 2) / 2 = 3. The same weighting defect can arise on one device without distributed training.
Unit
Valid count
Gradient sum
Local mean
Microbatch A
1
4
4
Microbatch B
3
6
2
Effective batch
4
10
2.5
A final short accumulation window creates another trap. Dividing every loss by the planned number of four microbatches when the final update contains only two changes its scale unless the protocol deliberately accounts for it. Derive the loss normalization from actual desired sample or token weight and document how incomplete windows are treated.
Exercise: verify three ranks against one device
Rank A has two valid tokens with gradient sum four. Rank B has one with gradient sum five. Rank C has three with gradient sum three. The desired global mean is twelve divided by six, or two. The local means are two, five, and one; their unweighted rank average is eight thirds, about 2.667.
With averaged gradients across three ranks, each local sum should be multiplied by 3/6 before backward. Their local gradients become two, 2.5, and 1.5; averaging yields two. Award one point for the global count, one for the desired mean, one for the incorrect rank mean, and two for the compensation derivation and single-device reference check.
Misconceptions to correct
“DDP automatically chooses the scientifically correct loss denominator” fails because user code defines local reduction and token masks. “Equal device count means equal training exposure” fails when sampler padding duplicates examples or tails are dropped. Inspect sample identities and counts in addition to the gradient formula.
Distributed equivalence tests should use intentionally uneven counts, then an equal-count sanity case, and a zero-valid-token boundary if supported. Compare the same initial parameters and effective batch with the single-device reference. A passing throughput benchmark without this reference can reward a faster implementation of the wrong objective.
The final packet should name the gradient collective semantics, loss sum definition, global denominator, accumulation boundary, and sampler behavior. These details let a reviewer reproduce both the arithmetic and the workload. They are also the most useful interview questions when a distributed run “converges differently” without an obvious exception.
Interview probe
Original practice: Why can adding GPUs change loss even when the nominal global batch size is unchanged? A strong answer checks per-rank reduction, valid-token counts, accumulation, and sampling. Follow up with uneven final batches. A weak answer assumes the framework guarantees the user's loss denominator.
A rank has zero valid tokens while others have valid work. What is required?
ASkip every collective and backward on that rank.BPretend it has one valid token to avoid zero division.CFollow a supported synchronized protocol with compatible graph and global count.DTerminate other ranks after their optimizer step.
A sampler duplicates tail examples to balance ranks. What should the performance report do?
AAssume equal devices imply equal unique-example exposure.BIgnore duplicates if throughput improved.CTreat repeated examples as new unique data.DRecord actual sampling and counts when defining equal work.
Can you derive the global loss gradient under the actual collective semantics? Rate confidence from 1 to 5 and verify it against an uneven-count single-device fixture.
Not yetGetting thereConfident
Wrap-up
Derive the global objective and verify the framework's reduction. Compare distributed output with a small single-device reference.