Summary
- GQA (Grouped Query Attention) divides the Query heads into G groups, and each group shares a single K,V head.
- As a compromise between MQA, where all heads share the same values, and MHA, where every head has its own independent values, the GQA-8 model shows quality close to MHA with speed close to MQA.
- Uptraining: The K, V projections of an already-trained MHA checkpoint are merged with mean pooling, and then the model is converted into an MQA/GQA model with additional training of only 5% of the original pre-training compute.
1. Problem & Motivation
As shown in Fast Transformer Decoding: One Write-Head is All You Need, MQA reduced the KV cache to h1 by having all heads share a single K, V. However, the MQA approach has several problems.
- Quality degradation: Since there is only one K, V head, model performance can drop.
- Training instability: Training becomes unstable, especially on tasks with long inputs.
- Separate training cost: Most large publicly released models are MHA, so retraining an MQA model from scratch for inference is costly.
To address these problems, the paper proposes two methods.
- Uptraining, which converts existing MHA checkpoints into MQA/GQA with little compute
- Grouped-Query Attention, which interpolates between MHA and MQA
2. Method
2.1. Uptraining
Uptraining produces a Multi-Query model from a Multi-head model, and the process consists of the following two steps.
- Converting Checkpoint: The per-head Key and Value projection matrices are mean pooled into a single Key and Value head.
- Additional pre-training: The converted model is trained further with the same recipe for a proportion α of the original pre-training compute so that it adapts to the new structure.
2.2. Grouped Query Attention
The H Query heads are divided into G groups, and each group shares one K, V head. A model built this way is called GQA-G.
- GQA-1 = MQA (1 K, V head)
- GQA-H = MHA (H K, V heads)
- A G in between gives an intermediate model with higher quality than MQA and faster speed than MHA.
When converting an MHA checkpoint to GQA, mean pooling is done only among the heads within each group.
So why is GQA an effective trade-off for large models?
- GQA can increase the number of groups in proportion to model size. MQA reduces to a single K, V head regardless of the original number of heads, so the representational capacity of K, V decreases as the model gets larger, whereas GQA can keep K, V heads in proportion to the number of heads.
- The larger the model, the smaller the share of the KV cache bottleneck. KV cache size is proportional to dmodel, but compute and parameters are proportional to dmodel2. So as the model grows, the relative impact of KV cache memory bandwidth decreases, and there is no need to use only a single K, V.
- Sharding waste is reduced. When a model is split across multiple GPUs/TPUs, MQA's single K, V head is replicated on every partition. GQA can set the number of groups to match the number of devices, which reduces this waste.
3. Experiments
3.1. Setup
- Base model: T5.1.1 Large, XXL (encoder-decoder). MHA-Large and MHA-XXL use the public checkpoints as they are.
- Uptraining: MQA-XXL and GQA-8-XXL are trained further for α=0.05 of the original pre-training steps.
- Tasks
- Summarization: CNN/Daily Mail, arXiv, PubMed, MediaSum, Multi-News (ROUGE-1)
- Translation: WMT 2014 En-De (BLEU)
- Question answering: TriviaQA (F1)
3.2. Results
The results of comparing MQA-, MHA-, and GQA-based models across multiple tasks are as follows.
- Uptrained MQA-XXL has both higher quality (46.6 vs 46.0) and faster speed (0.24 vs 0.37) than MHA-Large. In other words, uptraining a large model into MQA is a better trade-off than using a smaller MHA model. However, compared with MHA-XXL, its quality is 0.6 lower.
- GQA-8-XXL achieves almost the same quality as MHA-XXL (47.1 vs 47.2) while its speed is close to MQA (0.28 vs 0.24). It is about 5.4x faster than MHA-XXL (1.51).
- In the graph above, GQA-XXL sits at the upper left, i.e., in the fast and high-quality position.
3.3. Ablations
Since the ablations were run on a subset of the summarization tasks, the absolute scores differ from the main results. Performance is compared across checkpoint conversion methods, uptraining steps, and the number of groups.
3.3.1. Checkpoint conversion
Three methods for merging the K, V heads into one were compared.
- Mean: Average the K, V projections of all heads.
- First: Take only the K, V projection of the first head.
- Random: Initialize the K, V projections randomly from scratch.
The results are in the order Mean > First > Random. The more information from the pre-trained model is preserved, the better the performance.
3.3.2. Uptraining steps
- GQA achieves decent performance with conversion alone, without uptraining (α=0). In contrast, MQA's performance drops sharply right after conversion, so uptraining is essential.
- Both models improve substantially up to α=0.05, and the gains shrink at α=0.1. So the paper uses α=0.05 as the default.
3.3.3. Number of groups
- Increasing the number of groups from 1 (MQA) to 8 does not change inference time much
- Beyond that, inference time increases rapidly, and at 64 it becomes the same as MHA.
- The number of groups was set to 8 as the balance point between quality and speed.
4. Conclusion
GQA is a method that interpolates between MQA and MHA. It achieves quality close to MHA at speed close to MQA. On top of that, with uptraining, an existing MHA checkpoint can be converted into a GQA model with only 5% of the original pre-training compute.
However, there are also several limitations. The uncertainty of the evaluation metrics used to score each task was raised as an issue, and since the uptrained model was not compared with a model trained from scratch, no performance comparison with a from-scratch model was made. All tasks were run only with an encoder-decoder model (T5), and the effect on decoder-only models was not verified either. Decoder-only models, the mainstream of today's LLMs, have only self-attention and no cross-attention, so the paper only expected that the effect of GQA would be even larger.
5. My Take
1. It connects to the hypothesis I made in the MQA post.
In the MQA review, I hypothesized that "in attention, the Query is the most important, and it is enough for K, V to be shared." The GQA results also point in this direction. Even when the K, V heads are reduced from 64 to 8, quality stays almost the same. However, since quality drops with MQA (1 head), it also shows that K, V still need at least a minimum level of diversity.
2. The fact that mean pooling works best may mean that the K, V of different heads are quite similar.
One might expect that averaging completely different projections would produce meaningless values, but in practice Mean performed better than First and Random. This shows that there is a lot of redundancy among the heads' K, V projections. This redundancy seems to be why GQA performs reasonably well even at α=0.
References
- Ainslie, J. et al. (2023). GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. EMNLP 2023. arXiv:2305.13245
- Shazeer, N. (2019). Fast Transformer Decoding: One Write-Head is All You Need. arXiv:1911.02150
This post was translated from the original Korean version with the help of AI, so some errors may remain.