Skip to main content

Grouped-Query Attention: Head Sharing and KV Cache

Trace how grouped-query attention keeps independent query heads while sharing fewer key/value heads, projections, and compact KV-cache rows during LLM decoding.

Grouped-query attention (GQA) keeps the full set of query heads but uses fewer key and value heads. Several query heads read the same K/V head, while each query still computes its own scores, softmax weights, and output. The architectural change narrows the K/V projections and—most importantly for autoregressive serving—stores fewer rows in the KV cache.

This places GQA on one exact axis. If the number of K/V heads equals the number of query heads, the layer is ordinary multi-head attention (MHA). If there is one K/V head, it is multi-query attention (MQA). Intermediate counts are GQA.

Follow one query from projection to decode

Change the K/V-head count to move between MHA, GQA, and MQA. Click a query lane to trace its group, then play the sequence through scoring, value mixing, compact cache storage, and the next-token append.

The head mapping

Let Hq be the number of query heads and Hkv the number of K/V heads. In the common evenly grouped layout, Hq is divisible by Hkv, and query head h reads K/V head

g(h) = \left\lfloor h HkvHq \right\rfloor

Each K/V head therefore serves Hq/Hkv query heads. With eight queries and two K/V heads, Q0…Q3 map to KV0, while Q4…Q7 map to KV1.

For query head h, attention remains

Oh = \operatorname{softmax}\!(Qh Kg(h)^\top√(d) + M)Vg(h)

where M is the causal or other attention mask. Queries in one group share K and V tensors. They do not share query projections, attention-weight rows, or output vectors.

What changes in the projections

For hidden states X ∈ ℝ^(B×T×D), a GQA layer commonly uses these learned maps:

  • WQ ∈ ℝ^(D×Hq·d) for all query heads,
  • WK ∈ ℝ^(D×Hkv·d) for the smaller key set,
  • WV ∈ ℝ^(D×Hkv·d) for the smaller value set,
  • WO ∈ ℝ^(Hq·d×D) after concatenating query-head outputs.

The reshaped tensors are

  • Q: [B, Hq, T, d],
  • K: [B, Hkv, T, d],
  • V: [B, Hkv, T, d].

Reducing Hkv narrows the K and V projection outputs, parameters, and projection work. Q and output projection widths remain tied to Hq. At one decode step, every query head still scores the available context, so the number of query-key score pairs is Hq × T; it does not fall to Hkv × T.

Why the KV cache shrinks

During autoregressive generation, each layer retains past K and V vectors. For batch size B, Nl layers, cached length T, Hkv K/V heads, head dimension d, and b bytes per element, the raw tensor payload is

\text{cache bytes} = 2 Nl B T Hkv d b

The leading 2 accounts for both K and V. Relative to MHA with the same Hq, the exact payload fraction is

\text{GQA cache}\text{MHA cache} = HkvHq

This is a tensor-payload calculation, not a process-memory forecast. Allocator rounding, page tables, cache blocks, quantization metadata, padding, and framework bookkeeping can change observed device memory.

Concrete model configurations

ModelQuery headsK/V headsQ per K/VRelevant context mechanism
Llama 2 70B64884,096-token context
Llama 3 8B32848,192-token context
Mistral 7B v0.132844,096-token sliding window

For Llama 2 70B at 4,096 cached positions, 80 layers, head dimension 128, one sequence, and bf16, the raw GQA payload is 1.25 GiB. A hypothetical 64-K/V-head MHA cache with the same dimensions would be 10 GiB. These values exclude runtime overhead and any cache quantization.

Shape-faithful teaching implementation

The following code makes the group mapping visible. index_select creates query-aligned K/V tensors for clarity; it is not the storage strategy to copy into a production cache.

code
import math import torch def grouped_attention(q, k, v, mask=None): # q: [batch, Hq, query_tokens, head_dim] # k/v: [batch, Hkv, key_tokens, head_dim] hq, hkv = q.shape[1], k.shape[1] assert hq % hkv == 0 assert k.shape == v.shape group = torch.arange(hq, device=q.device) * hkv // hq k_for_q = k.index_select(1, group) # teaching expansion only v_for_q = v.index_select(1, group) scores = torch.einsum('bhtd,bhsd->bhts', q, k_for_q) scores = scores / math.sqrt(q.shape[-1]) if mask is not None: scores = scores.masked_fill(~mask, float('-inf')) weights = scores.softmax(dim=-1) return torch.einsum('bhts,bhsd->bhtd', weights, v_for_q)

A production implementation should keep the cache compact as [B, Hkv, T, d] per layer and use a kernel with native grouped-query broadcasting when available. Materializing and storing Hq repeated K/V copies forfeits the memory benefit. Rotary position embeddings, causal masking, FlashAttention-style kernels, paged caches, and cache quantization are orthogonal concerns and can be combined with GQA.

Converting an MHA checkpoint

GQA is an architectural weight-shape choice, not a safe inference-only switch. The original GQA work describes uptraining an MHA checkpoint: preserve query and output projections, mean-pool K/V heads inside each target group, and continue training. Its experiments use roughly 5% of the original pretraining compute for uptraining. That result motivates conversion, but it does not guarantee quality recovery for every model, dataset, group count, or downstream task.

Training GQA from scratch is also possible. In either route, choose Hkv using evaluation and deployment constraints rather than a universal “best” group count.

Engineering boundaries

Divisibility and partitioning

The common contiguous mapping requires Hq % Hkv == 0. Powers of two are convenient, not mathematically required. Tensor-parallel sharding may impose additional local divisibility constraints.

Cache layout

Store K/V under the K/V-head dimension, not the query-head dimension. Frameworks vary in axis order, so inspect the actual contract rather than assuming one universal [layer, batch, head, sequence, dimension] layout.

Logical broadcast versus physical repeat

A query-aligned view is useful for explaining the math. Whether expansion is zero-copy, fused, or materialized depends on the tensor strides and kernel. Profile the actual backend; do not label an expand + reshape sequence universally free.

Quality is empirical

Cache fraction and tensor shapes are exact; a generic quality label is not. Sharing can affect model quality, and the result depends on training, scale, data, evaluation, and Hkv. Report task metrics rather than an invented quality curve.

Performance is more than cache bytes

GQA can reduce K/V projection work, cache capacity, and cache bandwidth. End-to-end latency and throughput still depend on sequence lengths, batching, memory layout, kernel support, tensor parallelism, quantization, and hardware. The attention score/mix work across query heads remains.

Primary sources

If you found this explanation helpful, consider sharing it with others.

Mastodon