Model and Inference Engineering · Staff
The GQA cache has the right shape. Why do answers change after tensor parallel reshaping?
The question
Interview question
A model has 32 query heads and 8 key/value heads. Four query heads share one KV head. The team changes tensor parallel degree and rewrites a kernel that expands K and V for attention. Shapes match, throughput improves, but token probabilities shift. Why can a cache with exactly eight KV heads still be wrong?
Take a few minutes to form your approach. Then open a worked answer and compare the decisions.
Reveal a worked answer
The number of stored heads is only half the contract. Each query head must use the right KV head from the trained model. In a common contiguous grouping, query heads 0 to 3 use KV head 0, heads 4 to 7 use head 1, and so on. A wrong reshape or interleave may assign head 1 to queries 1, 9, 17 and 25 instead. Everything has a legal shape, and every head can compute attention, but it is a different model. The GQA paper defines grouping query heads around shared K and V heads. A specific checkpoint's head order, conversion and tensor parallel partitioning still need to be verified against its implementation.
I would take one short sequence through the reference attention path and the new path, capture Q, K and V after projection and rotary position handling, and compare each query head's chosen KV index. Next compare per-head attention output before the final output projection, then logits. This localizes a mapping error without a generation text diff. Test prefill and one cached decode step because a kernel can be right for one and wrong for the other. Include a nonuniform K/V fixture so swapping two KV heads cannot accidentally pass. A repeated identical head fixture would hide the bug.
Tensor parallelism makes the mapping less obvious. Query and KV heads may be sharded or replicated differently when the number of KV heads is smaller than the number of ranks. The code must preserve the checkpoint's global head order after gathering, replication and any fused QKV projection reorder. An arithmetic rule like query_head // 4 is correct only for that checkpoint layout and head convention. Write the reference mapping down as an invariant and validate it across supported parallel degrees.
Someone may say loss or aggregate benchmark quality looks almost unchanged. A head mix-up can have uneven effects by prompt and position, so inspect logit parity on a targeted fixture and slice downstream evaluation. Grouped query attention cuts KV memory. Why might prefill still be expensive? asks why GQA saves decode KV bandwidth but not necessarily prefill compute. This question asks whether the same GQA model survived a serving layout change.
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 →