[Paper review] Fast Transformer Decoding: One Write-Head is All You Need

Jumyung Song·2026년 9월 25일

Paper review

목록 보기
2/6
post-thumbnail

Summary

  • MQA (Multi-Query Attention): Queries remain different for each head, while Keys and Values are shared across all heads, accelerating the incremental decoding step.
  • Incremental inference in Transformers is slow not because of the amount of computation, but because of memory bandwidth: K and V must be read from memory repeatedly every time a token is generated.
  • The K, V cache size is reduced to 1h\frac{1}{h} with respect to the number of heads hh, making decoder inference about 12x faster with only a very small drop in quality.


1. Problem & Motivation

Transformers exchange information within and across sequences through attention layers. During training, the entire sequence can be processed in parallel at once, which makes it fast, but this kind of parallelization cannot be applied during incremental inference.

  • During incremental inference, only one token is generated at each step, so parallelization is not possible.
  • For the attention operation, the K, V of all previous positions must be loaded from memory at every step. In other words, the amount of data read is large relative to the amount of computation, so memory bandwidth becomes the bottleneck.

This paper resolves the bottleneck by using MQA to reduce the size of the K, V cache that has to be read.

2. Background: Neural Attention

The paper expresses every operation in einsum notation, a generalized contraction between tensors. Thanks to this notation, the shape of each tensor and the flow of computation are clearly visible.

2.1. Dot-Product Attention

Attention for a single query.

def DotProductAttention(q, K, V):
  """
  q: [k]      K: [m, k]      V: [m, v]
  Returns y: [v]
  """
  logits = tf.einsum("k,mk->m", q, K)
  weights = tf.softmax(logits)
  return tf.einsum("m,mv->v", weights, V)

2.2. Multi-head Attention

hh different attention layers (heads) are used in parallel.

  • Query: created by projecting the current input xx.
  • Key, Value: created by projecting the entire sequence MM to be referenced.
def MultiheadAttention(x, M, P_q, P_k, P_v, P_o):
  """
  x: [d]    M: [m, d]
  P_q, P_k: [h, d, k]    P_v, P_o: [h, d, v]
  """
  q = tf.einsum("d,hdk->hk", x, P_q)
  K = tf.einsum("md,hdk->hmk", M, P_k)
  V = tf.einsum("md,hdv->hmv", M, P_v)
  logits = tf.einsum("hk,hmk->hm", q, K)
  weights = tf.softmax(logits)
  o = tf.einsum("hm,hmv->hv", weights, V)
  y = tf.einsum("hv,hdv->d", o, P_o)
  return y

2.3. Multi-head Attention (Batched)

In actual MHA computation, to use the GPU as efficiently as possible, queries are generated at nn different positions at once so that nn tokens are processed together, and bb sequences that do not interact with each other are processed as a batch. In other words, nbnb tokens are processed in a single step.

def MultiheadAttentionBatched(X, M, mask, P_q, P_k, P_v, P_o):
  """
  X: [b, n, d]    M: [b, m, d]    mask: [b, h, n, m]
  """
  Q = tf.einsum("bnd,hdk->bhnk", X, P_q)
  K = tf.einsum("bmd,hdk->bhmk", M, P_k)
  V = tf.einsum("bmd,hdv->bhmv", M, P_v)
  logits = tf.einsum("bhnk,bhmk->bhnm", Q, K)
  weights = tf.softmax(logits + mask)
  O = tf.einsum("bhnm,bhmv->bhnv", weights, V)
  Y = tf.einsum("bhnv,hdv->bnd", O, P_o)
  return Y

Performance Analysis

The following assumptions are applied to express performance intuitively.

  • m=nm = n (# of key/value = # of query)
  • k=v=dhk = v = \frac{d}{h} (combining all heads gives dimension dd)

The performance analyzed under these conditions is as follows.

  • Operations: Θ(bnd2)\Theta(bnd^2)
  • Memory access: Θ(bnd+bhn2+d2)\Theta(bnd + bhn^2 + d^2)
    • bndbnd: X,M,Q,K,V,OX, M, Q, K, V, O
    • bhn2bhn^2: logits & weights
    • d2d^2: projection matrices
  • Ratio (memory / operations): Θ(1k+1bn)\Theta\left(\frac{1}{k} + \frac{1}{bn}\right)

The lower MemoryOperations\frac{\text{Memory}}{\text{Operations}}, the better, because it means a lot of computation is done while reading little data. During training, kk and bnbn are large enough that this ratio is small, so the GPU/TPU can be used efficiently.

2.4. Multi-head Attention (Incremental)

During inference, only one token is processed per step. Instead of creating K, V anew each time by multiplying MM by PK,PVP_K, P_V, only the value of the current token is appended to the K, V up to the previous position (prev_K, prev_V).

def MultiheadSelfAttentionIncremental(x, prev_K, prev_V, P_q, P_k, P_v, P_o):
  """
  x: [b, d]
  prev_K: [b, h, m, k]    prev_V: [b, h, m, v]
  Returns y: [b, d], new_K: [b, h, m+1, k], new_V: [b, h, m+1, v]
  """
  q = tf.einsum("bd,hdk->bhk", x, P_q)
  new_K = tf.concat([prev_K, tf.expand_dims(tf.einsum("bd,hdk->bhk", x, P_k), axis=2)], axis=2)
  new_V = tf.concat([prev_V, tf.expand_dims(tf.einsum("bd,hdv->bhv", x, P_v), axis=2)], axis=2)
  logits = tf.einsum("bhk,bhmk->bhm", q, new_K)
  weights = tf.softmax(logits)
  o = tf.einsum("bhm,bhmv->bhv", weights, new_V)
  y = tf.einsum("bhv,hdv->bd", o, P_o)
  return y, new_K, new_V

Performance Analysis (over all nn steps)

  • Operations: Θ(bnd2)\Theta(bnd^2) → same as batched
  • Memory access: Θ(bn2d+nd2)\Theta(bn^2d + nd^2)
    • bn2dbn^2d: the entire K, V is re-read at every step
    • nd2nd^2: the projection matrices are re-read at every step
  • Ratio: Θ(nd+1b)\Theta\left(\frac{n}{d} + \frac{1}{b}\right)
    - nd\frac{n}{d}: as the sequence length nn approaches dd, the cost of reading K, V becomes dominant.
    - 1b\frac{1}{b}: when the batch is small, the cost of reading the projection matrices becomes dominant.

The 1b\frac{1}{b} term can be reduced by increasing the batch size, but the nd\frac{n}{d} term is a structural problem. Reducing this nd\frac{n}{d} term is the goal of MQA.

3. Multi Query Attention (MQA)

Key Idea

Different heads share a single set of Keys and Values.

From an implementation standpoint, it is really simple. Just remove the hh dimension from the K, V-related tensors in the multi-head code.

def MultiquerySelfAttentionIncremental(x, prev_K, prev_V, P_q, P_k, P_v, P_o):
  """
  x: [b, d]
  prev_K: [b, m, k]    prev_V: [b, m, v]        # the h dimension is removed
  P_q: [h, d, k]    P_k: [d, k]    P_v: [d, v]    P_o: [h, d, v]
  Returns y: [b, d], new_K: [b, m+1, k], new_V: [b, m+1, v]
  """
  q = tf.einsum("bd,hdk->bhk", x, P_q)          # Query is still per head
  new_K = tf.concat([prev_K, tf.expand_dims(tf.einsum("bd,dk->bk", x, P_k), axis=1)], axis=1)
  new_V = tf.concat([prev_V, tf.expand_dims(tf.einsum("bd,dv->bv", x, P_v), axis=1)], axis=1)
  logits = tf.einsum("bhk,bmk->bhm", q, new_K)  # all heads see the same K
  weights = tf.softmax(logits)
  o = tf.einsum("bhm,bmv->bhv", weights, new_V) # all heads see the same V
  y = tf.einsum("bhv,hdv->bd", o, P_o)
  return y, new_K, new_V
  • Query projection P_q: [h, d, k] kept as is
  • Key/Value projection P_k, P_v: [h, d, k] → [d, k]
  • KV cache prev_K, prev_V: [b, h, m, k] → [b, m, k] → size reduced by a factor of 1h\frac{1}{h}

Performance Analysis (Incremental)

  • Operations: Θ(bnd2)\Theta(bnd^2) → same
  • Memory access: Θ(bnd+bn2k+nd2)\Theta(bnd + bn^2k + nd^2)
    • bndbnd: x,q,o,yx, q, o, y (xx is accessed nn times)
    • bn2kbn^2k: K, V → ∑m=1nbmk=Θ(bn2k)\sum_{m=1}^{n} bmk = \Theta(bn^2k)
    • nd2nd^2: projection → hdk×n=hd(dh)n=nd2hdk \times n = hd\left(\frac{d}{h}\right)n = nd^2
  • Ratio: Θ(1d+ndh+1b)\Theta\left(\frac{1}{d} + \frac{n}{dh} + \frac{1}{b}\right)

By sharing K, V, the problematic nd\frac{n}{d} term is reduced to ndh\frac{n}{dh}, i.e., by a factor of 1h\frac{1}{h}.

OperationsMemory accessRatio
MHA (Batched)bnd2bnd^2bnd+bhn2+d2bnd + bhn^2 + d^21k+1bn\frac{1}{k} + \frac{1}{bn}
MHA (Incremental)bnd2bnd^2bn2d+nd2bn^2d + nd^2nd+1b\frac{n}{d} + \frac{1}{b}
MQA (Incremental)bnd2bnd^2bnd+bn2k+nd2bnd + bn^2k + nd^21d+ndh+1b\frac{1}{d} + \frac{n}{dh} + \frac{1}{b}

4. Experiments & Results

4.1. Quality

Before evaluating quality: since reducing the K, V projections also reduces the number of parameters, the FFN hidden dimension dffd_{ff} was increased to match the total parameter count of the baseline.

  • MQA shows only a very slight drop in performance compared to the baseline (multi-head), and is even slightly higher in beam-4 test BLEU.
  • In contrast, other ways of reducing the K, V size (reducing the number of heads hh, or reducing dk,dvd_k, d_v) perform clearly worse than MQA.

4.2. Speed

The effect of sharing K, V across heads shows up clearly in the inference step. (Time per token)

EncoderDecoder
Baseline (MHA)1.7μs46μs
MQA1.5μs3.8μs
  • The encoder is processed in parallel just like in training, so there is almost no difference.
  • The decoder becomes about 12x faster. In incremental decoding, where memory bandwidth was the bottleneck, the effect of reducing the KV cache to 1h\frac{1}{h} shows up directly.

5. My Take

5.1. The Query seems to be the core of attention.

MQA ultimately unifies K, V across heads. Since different heads learn different correlations between tokens, I expected that unifying K, V across heads would degrade performance much more. In practice, however, the difference from the baseline was very small. This suggests that in attention, the Query, which is the information the current token asks for, has the greatest influence. It means that as long as each head differs in "what to ask," "where to look (K)" and "what to retrieve (V)" can be shared while still producing sufficiently diverse patterns.

5.2. Even for the same KV size reduction, results differ by method.

The performance drop when reducing the number of heads or the K, V dimension seems to come from the reduced information contained in K, V themselves. In that case, even if the FFN dimension is increased to match the parameter count, the loss cannot be recovered, because information sharing between tokens only happens in attention. MQA, on the other hand, keeps the K, V dimension (dk=128d_k = 128) and only removes the redundancy across heads, so there is less information loss.

5.3. Questions about the experimental tasks

The tasks used in the paper, such as WMT14 translation and Billion-Word LM, are somewhat old and small-scale by today's standards. So it remains a question whether MQA would perform well on tasks with larger datasets or harder tasks. That said, given that MQA was actually adopted in large models such as PaLM and Falcon, and that LLaMA-2 70B uses GQA, a generalization of MQA, it seems to have been accepted as a worthwhile trade-off at large scale as well.

References

  • 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.

0개의 댓글