Model and Inference Engineering · Staff
Each tensor-parallel GPU runs RMSNorm. Why does the full model disagree?
Take a few minutes to form your approach. Then open a worked answer and compare the decisions.
Reveal a worked answer
First ask which dimension was sharded. RMSNorm for a token uses the root mean square over its normalized hidden features. In a model with hidden width d, the denominator is sqrt(eps + sum(x_i²)/d) across those d features. PyTorch's RMSNorm definition makes that reduction axis explicit.
If two GPUs each hold half the hidden vector and independently normalize their local half, they calculate two denominators. That is generally a different function. Try x = [1, 1, 10, 10] split in the middle. One GPU sees an RMS near 1, the other near 10. The full-vector RMS is near 7.1. Local normalization makes all four outputs roughly 1 before the learned scales. Full RMSNorm leaves the small coordinates near 0.14 and the large ones near 1.4. The difference is not rounding noise.
There are two sensible layouts. Keep the full hidden vector on a rank at the norm boundary and normalize locally, then shard for the next projections. Or keep hidden shards and reduce the local sum of squares across ranks, divide by the global d, and apply each rank's local affine-weight slice using the common denominator. That collective must complete before downstream projections consume the result. PyTorch's tensor-parallel tutorial uses sequence parallelism for RMSNorm by sharding the sequence dimension while retaining the hidden features needed for each token's normalization.
Compare the single-rank reference immediately before and after RMSNorm on a short fixed input. Record the actual tensor layout, the normalized shape, eps, accumulator precision and weight slice. A correct attention all-reduce later cannot repair a wrong norm that already changed Q, K and V.
This question catches a common shortcut: “The operation is per token, so each GPU can do it alone.” Per token says nothing about whether each GPU has all features of that token. Know the reduction dimension before deciding whether a local kernel is equivalent.
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 →