Model and Inference Engineering · Principal
The MHA checkpoint loads after converting to GQA. Why did answer quality fall?
Take a few minutes to form your approach. Then open a worked answer and compare the decisions.
Reveal a worked answer
First check what the conversion did to the key and value heads. In multi-head attention, each query head has its own learned K and V projections. In grouped-query attention, several query heads share a K/V head. If we had 32 query heads and moved to 8 K/V heads, the cache for those heads becomes roughly one quarter as large, holding sequence length, head dimension, layers and dtype fixed. But 32 learned K/V functions have just been compressed into 8. Loading tensors with the right shapes does not preserve the old function.
An arbitrary reshape is especially suspect. Even averaging the K/V projections within each intended group, as in the GQA uptraining paper, is an initialization for further training, not a promise that the converted model answers identically. Check exactly which query heads map to which K/V heads, the weight layout in the checkpoint and the serving kernel, and whether tensor parallel ranks agree on that mapping. A transpose error can be different from the expected quality loss caused by sharing.
I would separate those two with a small forward-pass comparison. Run the original model against a reference MHA implementation, then the converted checkpoint against a simple, independently checked GQA implementation. Compare per-layer K/V projections and logits on short prompts before touching cache behavior. If the reference GQA and production GQA agree, but both diverge from MHA, we probably have a real compression and adaptation problem. If the two GQA implementations disagree, debug the conversion or kernel first. Then uptrain and evaluate the converted model on long context, tool use and rare slices, not only average perplexity.
Would I automatically use the fewest K/V heads? No. The serving win depends on whether KV bandwidth or capacity was actually the limit. Fewer heads can save cache bytes while taking an unacceptable hit in quality. Choose the head count from a measured quality, memory and latency curve, and keep the old checkpoint as a rollback path. Grouped query attention cuts KV memory. Why might prefill still be expensive? asks why GQA does not remove prefill cost. Here the question is what happens to a trained model when its attention parameterization changes.
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 →