All posts

Grouped query attention

/AI/4 min read

Every attention head keeping its own keys and values makes the cache enormous. Grouped query attention has heads share them in small groups, cutting the memory by the size of the group.

Attention normally runs several heads side by side, and each head has its own queries, keys and values. Grouped query attention keeps the separate queries but has heads share keys and values in groups.

That one change is aimed at a single number: how much memory a model needs while it is generating.

Why the cache grows

A model generates one token at a time, and each new token attends to everything before it. Rather than recompute the keys and values of all the earlier tokens at every step, they are worked out once and kept. That store is the KV cache.

It grows with the sequence, and it is multiplied by the number of heads, because every head has its own keys and values to keep.

Take a model with 8 heads, a head dimension of 128, and 32 layers, holding 8,000 tokens in fp16. Every token adds a key and a value per head per layer:

2×32 layers×8 heads×128×2 bytes=128 KiB per token2 \times 32 \text{ layers} \times 8 \text{ heads} \times 128 \times 2\text{ bytes} = 128 \text{ KiB per token}

Over 8,000 tokens that is just under a gigabyte, for one conversation. Serving many at once, the cache — not the weights — is what fills the GPU.

Two ways to shrink it

The heads are the multiplier, so the saving has to come from there.

The blunt version is to keep one set of keys and values for the whole layer and let every head share it. Each head still has its own queries, so it still asks its own question, but every head now looks at the same evidence. This is multi-query attention, and with 8 heads it cuts the cache eight-fold.

It also costs quality. What made the heads different was partly that they built different keys and values — different views of what each token offers. Collapse those into one and the heads lose much of their independence.

Grouped query attention is the version in between. Split the heads into groups, and give each group its own keys and values.

MHA 8 kv sets 1,000 MiB GQA 2 kv sets 250 MiB MQA 1 kv set 125 MiB
Same eight query heads throughout. Only the number of key-value sets below them changes.

The queries stay separate in all three. That matters, because the query is what decides what a head goes looking for — it is where most of a head's individuality lives. Sharing the keys and values costs less than sharing the queries would.

It is one dial, not three designs

Grouped query attention is not a third scheme sitting between two others. It is the general case, and the other two are its endpoints.

Set the number of groups equal to the number of heads and every head gets its own keys and values — that is ordinary multi-head attention. Set it to 1 and every head shares one set — that is multi-query attention. Anything between is what gets called GQA.

In code the dial usually appears as two numbers, num_attention_heads and num_key_value_heads, and the group size is just one divided by the other. A model with 32 query heads and 8 key-value heads runs 8 groups of 4.

What it saves

Same model as before — 8 heads, head dimension 128, 32 layers, 8,000 tokens, fp16:

kv setsper tokenat 8,000 tokensvs MHA
multi-head8128 KiB1,000 MiB—
grouped, 2 groups232 KiB250 MiB4× smaller
multi-query116 KiB125 MiB8× smaller

The saving is exactly the group size. Four heads per group, four times less cache. Nothing subtler is going on.

What that buys is not really speed at the arithmetic — it is room. A smaller cache means longer contexts fit, and more conversations fit on one GPU at the same time. Generation is limited by moving memory around rather than by multiplying numbers, so a cache a quarter of the size is also read a quarter as fast.

Converting a model that already exists

A model trained with full multi-head attention can be moved to grouped attention without starting again.

Take the heads destined for one group and average their key projection matrices into a single one. Do the same for the values. The model now has one key-value set per group, built out of what its heads already learned, and a short period of further training recovers most of what the averaging blurred.

That is cheap next to the original training run, which is a large part of why the technique spread quickly — existing models could adopt it rather than waiting for the next generation.

The short version

  • The KV cache is multiplied by the number of attention heads, and it is what fills memory during generation.
  • Multi-query attention gives every head one shared set of keys and values — smallest cache, some quality lost.
  • Grouped query attention splits heads into groups and gives each group its own set.
  • Queries always stay separate, because that is where most of a head's individuality sits.
  • Groups equal to heads is multi-head attention; one group is multi-query; it is a single dial.
  • The cache shrinks by exactly the group size, which buys longer contexts and more concurrent users.
  • An existing multi-head model can be converted by averaging key and value matrices within each group and briefly training on.