An annotated companion · AI Primer

Multi-query attention, annotated

About this page. This is a companion, not a copy. It follows the paper section by section, quotes at most a sentence or two per section (clearly marked and attributed), and explains everything in its own words. The paper is distributed on arXiv under arXiv's non-exclusive licence, so its tables appear here only as a few selected rows or as redrawn charts, always attributed. The paper has no figures; the diagrams here are drawn from its code and its analysis. The equations are reproduced because mathematics is not copyrightable. Every section heading links to the original.

How to read this page

  • Any dotted word explains itself on hover, focus or tap, and so does every symbol in every equation.
  • The diagram in §3 is live: hover or tap a part to see what it holds and how big it is. The chart after §3.1 has a slider for the batch size.

You will get the most from this page after the attention companion (which introduced the multi-head attention this paper modifies) and the KV-cache section of the inference lesson. Every idea climbs the same ladder: everyday picture, tiny example, diagram, the math, why it matters today.

Abstract · original

“We propose a variant called multi-query attention, where the keys and values are shared across all of the different attention “heads”, greatly reducing the size of these tensors and hence the memory bandwidth requirements of incremental decoding.”Shazeer (2019), Abstract

Everyday picture

Eight researchers work in one library. Each has a different question, but in the old arrangement each also kept a private copy of the whole card catalogue, and every time anyone wanted to look something up, all eight catalogues had to be wheeled out of storage. The new arrangement keeps eight researchers with eight questions and one shared catalogue. The questions are just as varied; the trolley is eight times lighter.

What the paper claims

  • Training a Transformer is fast, but generating from it one token at a time is slow, because every step reloads the large stored keys and values from memory.
  • Multi-query attention keeps many query heads but only one key head and one value head, so there is far less to reload.
  • On English-to-German translation, the decoder ran about 12 times faster per token, and quality fell only slightly (and less than any other way of shrinking the keys and values that the paper tried).

Why it matters today

Sharing keys and values between heads is now standard in large language models. The in-between version, grouped-query attention, shares one key/value head among a group of query heads; Mistral 7B uses 8 for its 32 query heads. Both descend from this paper.

1 Introduction · original

Everyday picture

A chef who cooks one dish at a time and must fetch every ingredient from a cellar for each dish spends most of the evening on the stairs, not at the stove. The chef is not slow at cooking; the stairs are the bottleneck. A GPU generating text is that chef: its arithmetic is fast, but for each new token it must fetch a lot of stored data from its main memory, and the fetching sets the pace. The paper calls such generation incremental inference.

Tiny example

To write a 128-token sentence, a model runs 128 separate steps, and step 100 cannot start until step 99 has chosen its token. In training, by contrast, all 128 positions of a known sentence are processed in one pass. The work in the two cases is about the same; what differs is how often the same stored numbers have to be reloaded.

Why it matters

The paper's plan is to measure the problem precisely (§2), change the architecture so the problem shrinks (§3), and check that the model is still good (§4). The inference lesson separates the two phases of generation, the parallel prefill and the one-token-at-a-time decode, for the same reason.

2 Background: neural attention · original

The paper explains attention through short TensorFlow functions written with einsum, a notation that names every dimension of every tensor with a letter. It is worth learning the letters, because the paper's whole idea is one letter:

The dimension letters used throughout the paper
LetterMeaningIn the paper's model
dwidth of every token's vector1,024
hnumber of attention heads8
k, vwidth of one head's keys and values128
mnumber of positions being attended to (the memory)up to 128
nnumber of query positionsup to 128
bbatch size: independent sequences processed together1,024 in the speed test

2.1 Dot-product attention · original

Everyday picture

You walk into a room with a question. Every person there wears a name badge (a key) and holds a note (a value). You compare your question with each badge, decide how much each person is worth listening to, and leave with a blend of their notes, mostly from the people whose badges matched best. That is attention for a single query.

Tiny example

One query q = (1, 0) and three positions with keys (1, 0), (0, 1), (1, 1) and one-number values 1, 2, 3. The scores are the dot products 1, 0, 1. Softmax turns them into shares 0.422, 0.155, 0.422, and the output is 0.422 × 1 + 0.155 × 2 + 0.422 × 3 = 2.0.

In words: “score the query against every key, turn the scores into shares that add up to 1, and hand back the values blended by those shares.”

With the numbers: q Kᵀ = (1, 0, 1); softmax gives (0.422, 0.155, 0.422); y = 0.422 × 1 + 0.155 × 2 + 0.422 × 3 = 2.0. The paper leaves out the usual division by √k and notes that it can be folded into the query projection.

In Python:

import math
q = [1.0, 0.0]
K = [[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]]
V = [1.0, 2.0, 3.0]
# logits = einsum("k,mk->m", q, K): one dot product per position
logits = [sum(q_i * k_i for q_i, k_i in zip(q, K_m)) for K_m in K]
logits  # → [1.0, 0.0, 1.0]
# weights = softmax(logits)
exps = [math.exp(s) for s in logits]
weights = [e / sum(exps) for e in exps]
[round(w, 3) for w in weights]  # → [0.422, 0.155, 0.422]
# y = einsum("m,mv->v", weights, V)
round(sum(w * v for w, v in zip(weights, V)), 3)  # → 2.0

Why it matters

Notice what the output needs: every key and every value. However short the query, all of K and V must be read. That fact is the whole story of this paper. The scaled_dot_product_attention function in the attention lesson is the same computation for many queries at once.

2.2 Multi-head attention · original

Everyday picture

Instead of one person asking one question, send h people into the room, each with their own question and their own way of reading the badges and notes. One cares who did what, another which word a pronoun refers to. Afterwards their findings are translated back into the room's common language and added together. That is multi-head attention, from the Transformer paper.

Tiny example

In the paper's model, d = 1,024 and h = 8 heads of width k = 128. Each head has four learned matrices of size 1,024 × 128: one each to make its query, keys, values, and to project its output back. Stacked over the 8 heads, each of the four projection tensors Pq, Pk, Pv, Po holds 8 × 1,024 × 128 = 1,048,576 numbers, which is exactly 1,024², because 8 × 128 = 1,024.

In words: “for each head, make a query from the input and keys and values from the memory, each with that head's own matrices; run attention; project the result back to width d; and add up the heads.” This is the paper's einsum code written as one line; the paper writes the head index and the head count with the same letter h, and this page uses i for the index.

With the numbers: each P holds h × d × k = 8 × 1,024 × 128 = 1,048,576 numbers, so one attention layer has 4 × 1,048,576 = 4,194,304 of them. The keys this layer computes for m positions form a tensor of h × m × k numbers: 8 × 128 × 128 = 131,072 for a 128-token memory.

In Python:

d, h, k, m = 1024, 8, 128, 128
# P_q, P_k, P_v and P_o each have shape [h, d, k]
P = h * d * k
P  # → 1048576
# the same as d², because h · k = d
P == d * d  # → True
# four projections in one attention layer
4 * P  # → 4194304
# K = einsum("md,hdk->hmk", M, P_k): one key per head per position
h * m * k  # → 131072

Why it matters

Look at the shape of K: it has an h in it. Every head makes its own keys and values from the same memory. Hold on to that letter; §3 deletes it. The lesson's MultiHeadAttention builds this layer from scratch.

2.3 Batched: many queries at once · original

Everyday picture

If a hundred people have questions for the same room, you open the doors once and let them all in together. In training, a model knows every position of the target sentence in advance, so it computes the queries for all n positions at once, and processes b unrelated sentences side by side as a batch. A mask of −∞ on the forbidden scores stops each position from reading the ones after it.

Tiny example

The training batches in the paper hold 128 sentences of 256 target tokens: 32,768 query positions per step. All of them use the same projection matrices, which are loaded from memory once per step and then used 32,768 times.

Why it matters

“Loaded once, used many times” is the property that makes training efficient on modern hardware. The next section counts it.

2.3.1 Why batched attention is fast · original

Everyday picture

A factory's speed can be limited by its machines or by its delivery trucks. What matters is the ratio: how many parts arrive per unit of machining. Modern chips, the paper notes, can do about a hundred times more arithmetic than their memory can feed them, so the ratio of numbers fetched to operations done must be small, or the machines sit idle. This ratio is the inverse of what the primer calls arithmetic intensity.

Tiny example

With the training batch above (b = 128, n = 256) and the paper's widths (d = 1,024, h = 8), one attention layer does about b·n·d² = 34 billion operations and touches about 102 million numbers: the activations (b·n·d = 34 million), the attention weights (b·h·n² = 67 million) and the projections (d² = 1 million). That is about 3 numbers fetched per thousand operations.

In words: “the work grows with the batch, the length and the width squared; the data touched is the activations, the attention weights and the weights of the projections; divide one by the other and what is left is one over the head width plus one over the number of query positions, both small.”

With the numbers: the exact ratio is (33,554,432 + 67,108,864 + 1,048,576) / 34,359,738,368 = 0.0030, under the bound 1/k + 1/(bn) = 1/128 + 1/32,768 = 0.0078. The paper assumes k = d/h and n ≤ d, which both hold here (256 ≤ 1,024).

In Python:

b, n, d, h = 128, 256, 1024, 8
k = d // h
# activations: X, M, Q, K, V, O and Y, each about b·n·d numbers
b * n * d  # → 33554432
# the logits and the weights: b·h·n²
b * h * n * n  # → 67108864
# the four projection tensors: d² each
d * d  # → 1048576
arithmetic = b * n * d * d
memory = b * n * d + b * h * n * n + d * d
round(memory / arithmetic, 4)  # → 0.003
# the paper's bound, 1/k + 1/(bn)
round(1 / k + 1 / (b * n), 4)  # → 0.0078

Why it matters

A ratio far below 1 means the chip can be kept busy: the data arrives faster than it is consumed. This is the compute-bound regime, and it is why training and prefill run near a chip's peak. The arithmetic_intensity function in the inference lesson computes the same idea for a whole model's weights.

2.4 Incremental: one token at a time · original

Everyday picture

When the model writes, the token it chooses at one position becomes the input at the next, so the positions can no longer be processed together. The model keeps a notebook of every earlier position's keys and values (today called the KV cache) and, for each new token, writes one more line in it and then reads the whole notebook.

Tiny example

The paper's function MultiheadSelfAttentionIncremental receives the new token's vector x, of shape [b, d], and the cached prev_K and prev_V, of shape [b, h, m, k]. It computes one query per head, appends one new key and value per head to the cache (making m + 1 positions), runs attention, and returns the output together with the longer cache. Every call reads the entire cache.

A note for careful readers of the paper's code: the incremental listings use M where they mean the new input x, and O where they mean the lower-case o computed a line earlier. The multi-query version in §3 also keeps axis=2 when appending the new key, which was right for the four-dimensional multi-head cache but not for the three-dimensional shared one, whose position axis is 1. They are typos; the shapes in the docstrings make the intent clear.

Why it matters

This is exactly how every chat model generates today. The TinyDecoder in the inference lesson generates with and without such a cache and counts the work saved.

2.4.1 Why incremental attention is slow · original

“When n ≈ d or b ≈ 1, the ratio is close to 1, causing memory bandwidth to be a major performance bottleneck on modern computing hardware.”Shazeer (2019), §2.4.1

Everyday picture

Back to the chef. With a big batch, each trip to the cellar for the recipe book (the weights) serves a thousand dishes, so that trip is cheap. But each dish also has its own notebook (its KV cache), and that notebook must be carried up for its dish alone, every single step. Batching does nothing for it.

Tiny example

One decoder self-attention layer, batch b = 1,024, at step m = 128: the cached keys and values are 2 × b × h × m × k = 2 × 1,024 × 8 × 128 × 128 = 268 million numbers, while the four projection tensors are 4 × 1,024² = 4.2 million. The step reads 64 times more cache than weights, and it does only about one multiply-add with each cached number it reads.

In words: “across n steps the work is the same as in training, b·n·d², but now every step rereads the whole cache (b·n·d numbers, n times over) and the weights (d², n times over); divide, and the ratio is the length over the width, plus one over the batch.”

With the numbers: the paper's speed test uses b = 1,024 and n = 128 with d = 1,024. The ratio is 128 / 1,024 + 1 / 1,024 = 0.125 + 0.001 = 0.126, over 40 times the 0.003 of training. With a single sequence, b = 1, it is 0.125 + 1 = 1.125: roughly one number fetched per operation.

In Python:

n, d, h, k = 128, 1024, 8, 128
def ratio(b):
    # memory / arithmetic for multi-head incremental decoding: n/d + 1/b
    return n / d + 1 / b
round(ratio(1024), 3)  # → 0.126
ratio(1)  # → 1.125
# one layer's cache at step m = 128, batch 1,024: 2 · b · h · m · k
2 * 1024 * h * 128 * k  # → 268435456
# against its four projection tensors
4 * d * d  # → 4194304

Why it matters

The paper names two ways to push the ratio down. The 1/b term is easy: use a bigger batch, if memory allows. The n/d term is the hard one, because it comes from rereading the cache. One fix is to attend to fewer positions (a local window, as in Mistral 7B's sliding-window attention). This paper's fix is orthogonal: keep every position, but make each position's keys and values smaller. The paper describes the cache's size as b·h·m·k = b·n²; under its own assumptions (h·k = d, m = n) that product is b·n·d per step, and the n² appears once all n steps are added up, as in the formula above.

3 Multi-query attention · original

“Multi-query attention is identical except that the different heads share a single set of keys and values.”Shazeer (2019), §3

Everyday picture

Back to the room of badges and notes. Multi-head attention gave each of the h questioners their own private way of reading every badge and every note, and so a private copy of all of them. Multi-query attention keeps the h different questions but gives everyone the same badges and the same notes. Different questions still pick out different people; the room just stops storing h copies of everything.

Tiny example

Two query heads, q1 = (1, 0) and q2 = (0, 1), share the three keys and values from §2.1. Head 1 scores (1, 0, 1) and returns 2.0, as before. Head 2 scores (0, 1, 1), gets shares (0.155, 0.422, 0.422), and returns 2.267. Two heads, two different answers, one set of keys and values: 3 × 2 = 6 key numbers stored instead of the 2 × 3 × 2 = 12 that two private key sets would need.

In Python:

import math
K = [[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]]
V = [1.0, 2.0, 3.0]
def head(q):
    # the same dot-product attention as §2.1, against the shared K and V
    s = [sum(a * b for a, b in zip(q, K_m)) for K_m in K]
    e = [math.exp(x) for x in s]
    return round(sum(e_j / sum(e) * v for e_j, v in zip(e, V)), 3)
head([1.0, 0.0])  # → 2.0
head([0.0, 1.0])  # → 2.267
# key numbers stored: shared, then private per head
len(K) * 2, 2 * len(K) * 2  # → (6, 12)

Diagram: one decoding step, before and after

Multi-head: h = 4, four caches K₁ V₁K₂ V₂K₃ V₃K₄ V₄ q₁q₂q₃q₄ new token x cache read per step: 4 × m × 2k numbers Multi-query: h = 4, one shared cache K V q₁q₂q₃q₄ new token x cache read per step: 1 × m × 2k numbers

Hover or tap a part. Start with the new token x in the top half, then compare the two halves' caches.

One step of incremental decoding, drawn from the paper's two code listings (§2.4 and §3). The paper itself has no figures.

Reading it: the top half is multi-head attention, the bottom half multi-query; read each from the bottom up. The new token's vector x is projected into four queries, one per head, in both designs. In the top half, each query reads its own cache: four tall boxes, each holding m positions (the horizontal lines) of keys and values of width k. In the bottom half, all four queries read the same single box. Every query still has its own line up to a cache, so the heads still ask four different questions, but the memory to be fetched at every step is four times smaller here, and h times smaller in general. The label under each half counts the numbers read from the cache per sequence per step.

The change in code

The paper's instruction is literally to delete a letter: remove h from every einsum wherever it indexes the keys, the values, Pk or Pv. Pk goes from shape [h, d, k] to [d, k]; the cache goes from [b, h, m, k] to [b, m, k]; the score line becomes einsum("bhk,bmk->bhm", q, K), so every head's query is dotted with the same keys.

Why it matters

In the primer's MultiHeadAttention, multi-query attention is one argument: n_kv_heads=1. MultiHeadAttention.kv_params counts the key and value weights that shrink, and the inference lesson's kv_cache_bytes_per_token shows the cache shrinking by the same factor.

3.1 The analysis again · original

Everyday picture

Same chef, same thousand dishes, but each dish's notebook is now h times thinner. The trip to the cellar for the notebook still happens every step, it just carries much less.

Tiny example

At b = 1,024, m = 128, the cache read per layer per step drops from 268 million numbers to 2 × 1,024 × 128 × 128 = 33.5 million: 8 times less, one factor of h.

In words: “the cache term now has a k where it had a d, because only one head's worth of keys and values is stored; dividing by the unchanged work, the troublesome n/d term shrinks by a factor of h.”

With the numbers: b = 1,024, n = 128, d = 1,024, h = 8: 1/1,024 + 128/8,192 + 1/1,024 = 0.001 + 0.0156 + 0.001 = 0.0176, against 0.126 for multi-head: 7.2 times less memory traffic per operation. With b = 1 it is 1.017 against 1.125: almost no gain, because the weights, not the cache, dominate a batch of one.

In Python:

n, d, h = 128, 1024, 8
k = d // h
def mha(b):
    # n/d + 1/b
    return n / d + 1 / b
def mqa(b):
    # 1/d + n/(dh) + 1/b
    return 1 / d + n / (d * h) + 1 / b
round(mqa(1024), 4)  # → 0.0176
round(mha(1024) / mqa(1024), 1)  # → 7.2
round(mqa(1), 3), mha(1)  # → (1.017, 1.125)
# one layer's cache read at step 128, batch 1,024: 2 · b · m · k
2 * 1024 * 128 * k  # → 33554432

Why it matters

The analysis predicts when the trick pays: large batches and long sequences, exactly the setting of a busy server. That prediction is what the chart below lets you explore, and what §4.3 tests.

Try it: batch size and length

Drag the slider to change the batch size b. The chart plots the two ratios from §2.4.1 and §3.1 against the sequence length n, with d = 1,024 and h = 8 as in the paper. Lower is better: it means fewer numbers fetched per operation.

Hover or tap the chart to read both ratios at one length.

Reading it: the x-axis is the sequence length n on a logarithmic scale, with a point at every doubling from 16 to 1,024 tokens; the y-axis is the memory-to-arithmetic ratio, constant factors dropped as in the paper. The solid line is multi-head attention and the dashed line multi-query. At the paper's batch of 1,024, the multi-head line climbs steadily with length (it is mostly n/d), while the multi-query line climbs eight times more slowly, so the gap widens as sequences get longer. Now drag the batch down to 1: both lines jump to just above 1 and nearly merge, because with one sequence the weights dominate and sharing keys and values saves almost nothing. Multi-query attention is a large-batch, long-sequence optimization.

4 Experiments · original

4.1 Setup · original

Everyday picture

A fair race needs equal cars. Deleting the h from Pk and Pv removes weights, so a multi-query model starts with fewer parameters than the baseline and would be expected to do a little worse for that reason alone. The paper gives the removed weights back to the feed-forward networks, widening them until the total count matches.

Tiny example

The translation baseline is an encoder-decoder Transformer with 6 layers, d = 1,024, feed-forward width dff = 4,096, h = 8 and dk = dv = 128: 211 million parameters. Every attention layer (encoder self-attention, decoder self-attention and encoder-decoder attention) becomes multi-query, and dff grows from 4,096 to 5,440. The paper states the new width; the accounting below, which is this page's own, shows why it is exactly right, assuming 6 encoder layers and 6 decoder layers as in the original Transformer.

In words: “the weights removed from the key and value projections of every attention layer equal the weights added to the two matrices of every feed-forward layer.”

With the numbers: 18 attention layers (6 + 6 + 6) each lose 2 × 7 × 1,024 × 128 = 1,835,008 weights, 33,030,144 in all. 12 feed-forward layers each gain 2 × 1,024 × 1,344 = 2,752,512, also 33,030,144. The language model (6 decoder-only layers, dff from 8,192 to 9,088) balances the same way: 11,010,048 on each side.

In Python:

d, h, k = 1024, 8, 128
def removed(L_attn):
    # each layer's P_k and P_v go from [h, d, k] to [d, k]
    return L_attn * 2 * (h - 1) * d * k
def added(L_ffn, d_ff_old, d_ff_new):
    # each feed-forward layer has a d × d_ff matrix in and a d_ff × d matrix out
    return L_ffn * 2 * d * (d_ff_new - d_ff_old)
# translation: 18 attention layers, 12 feed-forward layers
removed(18), added(12, 4096, 5440)  # → (33030144, 33030144)
# language model: 6 of each
removed(6), added(6, 8192, 9088)  # → (11010048, 11010048)

Two more families of models go into the race. Local versions restrict decoder self-attention to the current position and the previous 31, to show that windows and multi-query attention combine. And smaller multi-head models shrink the keys and values the obvious way, by using fewer heads (h = 1, 2, 4) or narrower ones (dk = dv = 16, 32, 64), again with wider feed-forward layers to keep the count equal. Training ran for 100,000 steps on a 32-core TPU v3 cluster, about 2 hours per model. The paper also trains decoder-only language models on the Billion-Word Language Modeling Benchmark (192 million parameters each).

Why it matters

Matching parameter counts is what makes the comparison honest: any quality difference is caused by how the weights are used, not by how many there are.

4.2 Quality · original

Everyday picture

If you make everyone share one catalogue, you worry the questions get worse. The test: measure translation quality with BLEU (higher is better) and the model's surprise at the reference translation with perplexity (lower is better), for every variant.

Selected rows of Table 1 (WMT14 English-German, development set), from Shazeer (2019), reproduced with attribution
Attentionhdk, dvdffln(PPL)BLEU
multi-head (baseline)81284,0961.42426.7
multi-query81285,4401.43926.5
multi-head, fewer heads2646,7841.48026.2
multi-head, one head11286,7841.51825.8

Tiny example

The table lists the natural logarithm of perplexity. Undoing it: e1.424 = 4.15 for the baseline and e1.439 = 4.22 for multi-query, a 1.5% rise in perplexity per subword token. The best of the “fewer or narrower heads” models reaches e1.480 = 4.39, a 5.8% rise.

In Python:

import math
base, mqa, fewer = 1.424, 1.439, 1.480
# perplexity is e to the logged value
round(math.exp(base), 2), round(math.exp(mqa), 2), round(math.exp(fewer), 2)  # → (4.15, 4.22, 4.39)
# rise over the baseline, in percent
round(100 * (math.exp(mqa - base) - 1), 1)  # → 1.5
round(100 * (math.exp(fewer - base) - 1), 1)  # → 5.8

On the test set with beam search (4 beams), multi-query scored 28.5 BLEU against the baseline's 28.4, the highest in the table; with greedy decoding it scored 27.5 against 27.7. On the Billion-Word benchmark, per-word perplexity was 29.9 for the baseline and 30.2 for multi-query, while every smaller multi-head model landed between 30.9 and 31.2.

Why it matters

The key comparison is not multi-query against the baseline but multi-query against the other ways of shrinking K and V. Shrinking the number or width of heads also shrinks the queries, and quality drops more. Keeping h separate queries while sharing the keys and values is the cheap cut.

4.3 Speed · original

Everyday picture

The paper times everything on one 8-core TPU v2, as a cost per token: the total time of a step divided by the tokens that step produced.

Tiny example

For the baseline, generating a batch of 1,024 sentences took 47 ms per decoder step, one token per sentence: 47 / 1,024 = 45.9 μs per token, which the paper rounds to 46. The multi-query decoder took 3.9 ms per step: 3.8 μs per token. The encoder, which reads all 128 source tokens of all 1,024 sentences in one parallel pass, took 222 ms, only 1.7 μs per token.

In words: “the cost of one token is the time of a step shared out over every token that step produced.”

With the numbers: baseline decoder 47 ms / (1,024 × 1) = 45.9 μs; multi-query decoder 3.9 ms / 1,024 = 3.8 μs; baseline encoder 222 ms / (1,024 × 128) = 1.69 μs; multi-query encoder 195 ms / (1,024 × 128) = 1.49 μs. The decoder speedup is 46 / 3.8 = 12.1 times. Training barely moves: 433 ms against 425 ms per step of 32,768 tokens, 13.2 against 13.0 μs per token.

In Python:

b, src = 1024, 128
def t_token(T_step_ms, n_step):
    # milliseconds per step, shared over b · n_step tokens, in microseconds
    return T_step_ms / (b * n_step) * 1000
round(t_token(47, 1), 1), round(t_token(3.9, 1), 1)  # → (45.9, 3.8)
round(t_token(222, src), 2), round(t_token(195, src), 2)  # → (1.69, 1.49)
# decoder speedup, from the paper's rounded values
round(46 / 3.8, 1)  # → 12.1
# training, per (input + target) token: 32,768 per step
round(433 / 32768 * 1000, 1), round(425 / 32768 * 1000, 1)  # → (13.2, 13.0)

Reading it: the bars redraw the decoder column of the paper's Table 2, in TPU-microseconds per output token; solid bars are multi-head models, striped bars multi-query. Each group has its own axis, because beam search costs several times more. Greedy decoding: 46 → 3.8, and with a 32-token local window, 23 → 3.3. Beam-4 search: 203 → 32, and 47 → 16 with the local window. In greedy decoding the local window alone halves the multi-head cost, multi-query alone cuts it by about 12, and together they cut it by 14: the two ideas attack different factors of the same cache (how many positions, and how big each one is).

The paper notes one caveat about the setup: to keep tensor shapes fixed, the cache was padded to its maximum length (128, or 32 for local attention), so every step cost the same; a growing cache would have saved time on the early steps.

Why it matters

The predicted gain (about 7 times less memory traffic per operation, §3.1) and the measured one (12 times faster decoding) point the same way; the analysis drops constant factors, so the two need not match exactly. Training speed is unchanged, which confirms that the saving comes entirely from the one-token-at-a-time setting.

5 Conclusion · original

Everyday picture

The whole paper is one move: find the number that is reloaded every step, and stop storing h copies of it.

What it concludes

Multi-query attention has much lower memory-bandwidth needs during incremental decoding, and the author expects it to make attention-based models practical where inference speed is critical. The experiments use models of about 200 million parameters and one design, a single shared key/value head; sharing one key/value head among each group of query heads was left to later work.

What happened next

DevelopmentWhat it changedBuild it
Grouped-query attention (Ainslie et al., 2023)A middle ground: a few key/value heads, each shared by a group of query heads, recovering most of multi-head quality at most of multi-query's savingMultiHeadAttention with 1 < n_kv_heads < n_heads
Mistral 7B (2023)32 query heads share 8 key/value heads, combined with a sliding window: the two cuts of §4.3, at scalesliding_window_decode
PagedAttention (2023)Once the cache is smaller, the next problem is packing many sequences' caches into memory without wasteKV-cache memory
Multi-head latent attention (DeepSeek-V2, 2024)Cache one small latent vector per token and rebuild every head's keys and values from itLatentKVAttention
FlashAttention (2022)A different memory problem: the n × n score matrix in training and prefill, solved by computing attention in tiles without changing the modelattention lesson

The efficient architectures lesson puts all of these cache-shrinking ideas on one chart, from multi-head attention's full cache down to a state-space model's fixed one.

Glossary

Every term with hover guidance on this page, in one place.