Model and Inference Engineering · Staff
The vocabulary is sharded across GPUs. Why is each rank's softmax the wrong loss?
The question
Interview question
You split the language model's final projection by vocabulary. Each rank holds one slice of the logits. Someone computes cross entropy on each slice and averages the losses. Training runs, but perplexity goes in a strange direction. What went wrong?
Take a few minutes to form your approach. Then open a worked answer and compare the decisions.
Reveal a worked answer
For target token y, the loss is log(sum_v exp(z_v)) - z_y. That sum is over every vocabulary token. A local softmax asks a different question: how likely is a token among the tokens on this GPU? Averaging those local answers cannot reconstruct the global probability. On most ranks, the target token is not even in the local vocabulary.
You can keep the logits sharded without gathering a huge [batch, sequence, vocab] tensor. For each token position, find the maximum logit across ranks. Subtract it before exponentiating. Sum each rank's exponentials, then sum those partial sums across the tensor-parallel group. The rank that owns y supplies its target logit, and the others contribute zero. The result is m + log(global_sum_exp) - z_y. Megatron's vocabulary-parallel cross entropy follows this distributed max, sum and target selection pattern. PyTorch's loss-parallel tutorial explains why the full logits need not be gathered.
One more trap: the rank that lacks the target cannot simply ignore the token. Its logits still appear in the denominator and must receive gradients. Masking the non-owning ranks out of the loss drops real competitors. On the other side, padding the vocabulary to divide evenly among ranks must not make padded token IDs into valid competitors. Keep target IDs in global coordinates and test the ignore-label mask before doing a local index lookup.
I would prove equivalence on a tiny vocabulary split over two ranks. Give the target to each rank in turn, include a padded vocabulary row and an ignored label, then compare per-token loss and local-logit gradients with an unsharded reference. Only after that would I benchmark the collective cost and precision. If the interviewer suggests gathering all logits as the safe version, it is a good correctness oracle for a small test. It is usually a poor production plan at large vocabulary and sequence sizes. Ranks process different token counts. Why does averaging their losses change the objective? is about weighting different numbers of tokens across data ranks. This is about normalizing one token over a vocabulary split across tensor ranks.
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 →