Model and Inference Engineering · Staff
Every attention head is correct. Why is the tensor-parallel output wrong?
Take a few minutes to form your approach. Then open a worked answer and compare the decisions.
Reveal a worked answer
You split attention heads across GPUs. Each rank's attention probabilities and head output match a single-GPU reference. Yet the block output does not. The first question I would ask is what happened at the output projection, W_o.
The heads produce pieces of one vector, h = [h_1, h_2, ...]. If W_o is partitioned along its input dimension to match those pieces, each rank computes a partial result, y_i = W_{o,i} h_i. Every y_i already has the output dimension. The answer is y = Σ_i y_i, followed by the projection bias once, if there is one. Concatenating those partial results, or treating one rank's result as the complete attention output, is wrong even though each head was calculated perfectly. PyTorch's tensor-parallel tutorial uses column-wise Q/K/V projections and a row-wise output projection for this reason.
That sum is a communication boundary. In ordinary tensor parallelism, an all-reduce can leave the complete projected result on every rank so it can be added to the replicated residual. In a sequence-parallel layout, reduce-scatter can deliver sequence shards instead. I would inspect the layout expected by the residual addition, rather than prescribing one collective without looking at the next operator.
To locate the bug, run one short sequence with dropout disabled and compare at four points: projected Q/K/V, per-rank head output, each rank's partial W_o product, and the result after communication. Verify the head order and weight shard axis as well. If the sum of the local projected products matches the reference but the actual block does not, inspect the collective, bias placement and residual layout. If the sum itself is wrong, look at the weight partition or head ordering.
An interviewer might say, “Why not gather all heads and apply the original projection?” That can be a correct reference implementation, but it moves the full head activation to every rank and loses the point of the row-parallel projection. Use it as a debugging oracle, then restore the partitioned computation with the required reduction. The real invariant is that every output feature receives contributions from all the input head partitions.
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 →