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.