Transformers — Compute Attention By Hand
Keep the Query Heads, Shrink Only the Key/Value Heads
In one line
There is only one knob that separates MHA, MQA and GQA — how many key/value heads to have. The number of query heads is not touched. So what shrinks is the memory the KV cache takes up, and the number of multiplications in attention stays almost the same.
Why this was needed
Generation moves forward one token at a time. Every time it produces a new token, the model looks again at the keys and values of every position produced so far. So instead of erasing those keys and values and rebuilding them, it holds on to them. That is the KV cache.
The problem is that this cache grows as it is when the context gets longer, when there are more layers, and when more requests are handled at once. The model weights are fixed regardless of context length, but the cache is not. The cache eats the room left after loading all the weights, so from some point the cache decides the number of concurrent requests.
On top of that, the computation done to produce one token is small. To do that small computation, you have to read the whole cache in from memory. This re-reading is exactly what the original MQA paper named as the bottleneck of incremental decoding. The title is the answer itself — you need only one head to use.
That is where the MQA decision comes from. Keep just one key/value head. The query heads stay as they are, whether eight or sixty-four, and all of them look at the same K·V. The cache shrinks to one over the number of query heads.
But there is something you lose. In a structure built so that each head looks at a different place, you have tied into one the side that decides "what to put out". The original GQA paper makes room in between. It divides the query heads into several groups and keeps one key/value head per group.
Three names are three spots on one axis
Say there are h query heads and g groups. Then the number of key/value heads is g.
| Number of groups g | Name | Key/value heads | Query heads looking at one K·V |
|---|---|---|---|
| g = h | MHA | h | 1 |
| 1 < g < h | GQA | g | h/g |
| g = 1 | MQA | 1 | h |
You should not memorize the three lines separately. It is a difference in how far you turned a single knob. There is only one condition — h must be divisible by g. Otherwise the group sizes become uneven, and some groups bundle three heads and others two. If you break this condition in a real configuration file, it usually gets stopped at the point of loading the model.
What shrinks and what does not
This is the heart of this topic, and where people go wrong most. If you write out the formulas, you see it at once.
The number of real numbers that go into the cache is as follows.
KV 캐시 원소 수 = 2 × 층 수 × (키·값 헤드 수) × 길이 × 헤드 차원
The 2 at the front is the two sets, K and V. The number of query heads is not in this formula. It is because queries are not kept in the cache. So if you reduce the key/value heads from eight to one, the cache becomes exactly one eighth.
Now let us count the multiplications. This is for passing once through one layer with length n, h query heads and head dimension d.
점수 Q·Kᵀ h × n × n × d ← g 가 없다
가중합 × V h × n × n × d ← g 가 없다
K 투영 n × d_model × (g × d)
V 투영 n × d_model × (g × d)
There is no g in the core of attention. The reason is clear when you look at the implementation. When using the folded key/value heads, you repeat and unfold them as many times as the number of query heads and run as usual. Only the number of distinct values decreased; the number of times attention runs is decided by the number of query heads. The terms that shrink are only the two K·V projections, and those are not a big share of the whole.
So it comes out like this. What GQA reduces is memory. It gets faster because that memory is read less, not because fewer multiplications are done. If you cannot make this distinction, you cannot answer the question "I switched to GQA, so why is prefill the same?"
How to merge the groups
Say you are converting an already-trained multi-head model to GQA. You have to make the four key/value heads in a group into one, but how do you make it?
You can keep just one and throw away the rest, you can average them, or you can retrain from scratch at the reduced size. The path the original GQA paper chose is to average the key/value heads in a group to use as the initial values, and then train a little more from there. It is important that it does not stop at averaging — the average is a starting point, not the answer.
The next lab also uses the average. But it does not train. So the number "the output changed by this much" that comes out of this lab does not mean quality gets worse by that much; it is a value that shows the fact that if you tie only the keys and values while keeping the same weights, the answer changes. You will write this distinction in the report as well.
What it looks like in the field
First, the configuration file has the number of heads written on two lines. In today's model configurations, the number of query heads and the number of key/value heads are separate. If the two values are equal it is MHA, if the key/value side is 1 it is MQA, and in between it is GQA. There is no need to look for the three names; you can just look at the ratio of the two numbers.
Second, the question "it's GQA, so why isn't it faster" comes up. When processing many positions at once, as in prefill, computation is the bottleneck, and g hardly affects that. The gain comes from the stretch where you hold a long context and draw tokens one at a time, and from raising the number of concurrent requests.
Third, the upper limit on the number of concurrent requests is set by the cache. The remaining memory divided by the size of one set of cache is roughly the upper limit. If you reduce the key/value heads to four, that one set gets smaller, so the upper limit rises. What grows at that point is throughput, not the speed of a single request.
Fourth, the number of key/value heads becomes an obstacle in tensor parallelism. When you split the heads across several devices, if there are fewer key/value heads than devices, there is not enough to split, and several devices end up holding the same key/value. If you lowered g to shrink the cache but it is replicated on every device, it does not shrink as much as expected.
Fifth, changing the number of heads makes it a different model. If you put weights made by averaging the groups into service as they are, the output changes. The original paper also does additional training after converting. It is not "just change the configuration".
What really matters in practice
- What shrinks is memory. The number of multiplications stays almost the same. You can decide where to use it only if you know where the gain comes from.
- h must be divisible by g. If the group sizes are not even, the implementation does not hold.
- The two ends are on the same axis. If g=h it is MHA, and if g=1 it is MQA. If you write the code so that one set of code can run all three, comparison becomes easy.
- Converting and training are different jobs. The average is only an initial value.
What you will do in the next lab
You grow /root/work/tf-gqa/gqa.py one step at a time. You start by grouping the query heads into contiguous chunks, fold the key/value heads within each group by averaging, repeat them and unfold them as many times as the number of query heads, and then, through single-head attention, run the three methods with the same weights. If you set g to h, it must not differ from MHA in a single place, and if you reduce g, the output must differ. You measure that difference in numbers.
Then comes the point of this lab. You count separately the number of KV cache elements and the number of multiplications and place them side by side. You will confirm, with numbers you counted yourself instead of someone else's explanation, that the cache shrinks in direct proportion to g while the multiplications of the core of attention never shrink even once. At the end you leave those two tables and the output difference as a report.
You do not measure time. This Pod has no GPU and the CPU is shared with other work, so "it got faster" is not something you can measure here. The system Python has no numpy either (it exists only inside /opt/onnx-lab). The standard library alone is enough, and every judgment is integer counting, such as number of elements and number of multiplications, and numerical comparison with a tolerance. The grader actually imports your module, pokes at the functions with a different number of heads and number of groups every time, and checks them against values it computes separately.