skip to content

Why does every Llama 3 size use grouped-query attention rather than multi-head attention?

level: middleimportance: must knowfreq 58%

answer

  1. Cost is in the cache, not the weights
  2. Query heads kept, key/value heads shared
  3. Eight KV heads at every size
  4. Decoding is bandwidth-bound
  5. Four-to-one on 8B, eight-to-one on 70B

basics

~20 s

Grouped-query attention lets several query heads share one key/value head, so the KV cache shrinks by that ratio. Llama 3 uses 8 key/value heads at every size, cutting per-token cache memory about fourfold on 8B and eightfold on 70B, which is what makes long contexts and large batches affordable.

solid answer

~60 s

During decoding, a transformer must keep the key and value vectors of every past token in memory — the KV cache — and re-read them for each new token. With plain multi-head attention, that cache scales with the full head count. **Grouped-query attention** keeps all the query heads but gives each *group* of them a single shared key/value head. Llama 3 8B has 32 attention heads and 8 KV heads (a 4:1 group ratio); the 70B has 64 attention heads and the same 8 KV heads (8:1). The result is a KV cache four to eight times smaller than the equivalent multi-head model, with negligible quality loss. That matters twice over: decoding is memory-bandwidth-bound, so a smaller cache means fewer bytes read per token and lower latency, and the memory you save is memory you can spend on longer contexts or bigger batches. Llama 2 only used GQA on its 70B; making it universal in Llama 3 is precisely what made a 128K window practical on an 8B model.

code

python · 6 lines
python
from transformers import AutoConfig

cfg = AutoConfig.from_pretrained("meta-llama/Llama-3.1-8B-Instruct")
print(cfg.num_attention_heads)   # query heads
print(cfg.num_key_value_heads)   # shared KV heads
print(cfg.num_attention_heads // cfg.num_key_value_heads)  # group size

go deeper

for a junior

Know that grouped-query attention means several query heads share one key/value head, and that its purpose is to shrink the key/value cache the model keeps while generating.

for a middle

Be able to compute the KV cache from layers, KV heads, head dimension and dtype, and to state Llama 3's numbers: 8 KV heads at every size, 32 query heads on 8B, 64 on 70B.

for a senior

Connect it to production behaviour — cache size sets your concurrency ceiling and your tokens-per-second at long context, and it is the reason a 128K window on an 8B model is servable at all.

for a principal

Own the memory budget across the fleet: how cache footprint per request bounds batch size and cost per token, when to trade context length for concurrency, and why the group ratio is a procurement-time property of the checkpoint, not a knob.

## The problem GQA solves Autoregressive decoding generates one token at a time. For each new token, attention needs the key and value vectors of every token that came before it. Recomputing them would be quadratic waste, so serving stacks cache them — the **KV cache**. The cache is per-request, grows linearly with sequence length, and lives in accelerator memory alongside the weights. Size it yourself. Bytes per token equals `2 (K and V) x layers x kv_heads x head_dim x bytes_per_element`. For Llama 3 8B — 32 layers, 8 KV heads, head dimension 128, bf16 at 2 bytes — that is `2 x 32 x 8 x 128 x 2 = 131,072` bytes, i.e. 128 KiB per token. A single 128K-token conversation therefore parks about 16 GiB of KV cache. Now redo it with 32 KV heads, as plain multi-head attention would require: 512 KiB per token, 64 GiB for that same conversation. On a single 80 GB accelerator that alone decides whether the deployment is possible. ## What grouped-query attention actually is Three points on a spectrum: - **Multi-head attention (MHA).** Every query head has its own key head and value head. Maximum expressiveness, maximum cache. - **Multi-query attention (MQA).** All query heads share exactly one key/value head. Minimum cache, and measurably worse quality, with training instability reported at scale. - **Grouped-query attention (GQA).** The middle: query heads are partitioned into `G` groups, and each group shares one key/value head. `G = num_heads` reduces to MHA; `G = 1` reduces to MQA. GQA keeps the *query* projections untouched — the model still has 32 or 64 distinct query heads attending in 32 or 64 different ways. Only the key and value projections are narrowed. Empirically that costs very little quality while recovering nearly all of MQA's memory savings, which is why it has become the default across essentially every open-weight family, not just Llama. ## The Llama-specific facts In a Hugging Face `config.json` the two fields to read are `num_attention_heads` and `num_key_value_heads`; their ratio is the group size. For Llama 3 and 3.1: - **8B**: 32 attention heads, 8 KV heads, 32 layers, head dim 128 → group size 4. - **70B**: 64 attention heads, 8 KV heads, 80 layers → group size 8. Note the KV head count stays at 8 as the model scales. That is deliberate: the larger model gets more query heads and more layers, but the per-layer KV footprint does not grow proportionally, so cache pressure scales more gently than parameter count would suggest. Llama 2 is the useful contrast. Its 7B and 13B used ordinary multi-head attention; only the 70B adopted GQA. Making GQA universal in Llama 3 was a serving decision as much as a modelling one — it is a precondition for the 128K context that arrived in 3.1, because without it the cache for a full window on the small model would dwarf the weights. ## Why the saving shows up as latency, not just memory Token-by-token decoding does very little arithmetic per byte moved: for each generated token you stream the weights and the entire KV cache through the accelerator's memory system. That makes decode **memory-bandwidth-bound**, not compute-bound. Shrinking the KV cache by 4x removes 4x of the bytes that must be read per step from the cache side, which shows up directly as higher tokens-per-second — especially at long contexts, where the cache dominates the weights in bytes moved. It also raises the batch ceiling. Concurrency in a modern server is limited by how many requests' caches fit in memory at once. Quarter-size caches mean roughly four times the concurrent sequences at the same memory budget, which is throughput and cost per token, not just a headroom nicety. ## Where the tradeoff bites GQA is not free. Sharing key/value heads reduces the diversity of what each group can attend over, and ablations show a small quality cost relative to full MHA — small enough that every major family accepted it. Models are trained with their group ratio from the start; you cannot convert an MHA checkpoint to GQA at serving time without an uptraining procedure, so the ratio is a fixed property of the weights you downloaded, readable from the config and nothing you tune. ## Answering well Lead with the mechanism (query heads grouped over shared KV heads), give the Llama numbers (8 KV heads at every size, 4:1 on 8B, 8:1 on 70B), then explain *why anyone cares* — the KV cache is the thing that scales with context and concurrency, decoding is bandwidth-bound, and the saving converts into longer windows and larger batches. Candidates who describe GQA purely as a parameter-count optimisation have missed the point; the weights barely shrink, the cache does.

  • Does grouped-query attention reduce the model's parameter count much?
    Barely. Only the key and value projection matrices shrink, and attention projections are a modest slice of a transformer's parameters next to the feed-forward blocks. The saving that matters is at runtime: the KV cache, which scales with sequence length and concurrent requests. Describing GQA as a weight-compression technique is the classic misread — it is a serving-memory and memory-bandwidth technique.
  • Could you take a multi-head checkpoint and just serve it with fewer KV heads to save memory?
    Not without retraining. The key and value projections are learned per head; dropping or averaging them changes what the model computes and degrades quality sharply. The published approach is uptraining — mean-pool the key/value heads within each group, then continue pretraining for a small fraction of the original budget. Llama models ship with their group ratio baked in, so at serving time it is a fixed property of the weights.
  • Why does the 70B keep 8 KV heads rather than scaling them with its 64 query heads?
    Because KV heads are the term that multiplies the cache. Holding them at 8 while query heads double to 64 means the per-layer cache footprint stays flat as capacity grows, so the bigger model's cache pressure rises only with its layer count, not with its head count. It is an explicit choice to keep long-context and high-batch serving viable on the larger model.

saying these in an interview costs you the question

  • Saying GQA mainly shrinks the model weights
  • Confusing GQA with multi-query attention's single KV head
  • Claiming Llama 2 used GQA at every size
  • Thinking the group ratio is a serving-time tunable
  • Ignoring that decoding is memory-bandwidth-bound

context