
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.
This paper resolves the bottleneck by using MQA to reduce the size of the K, V cache that has to be read.
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.
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)
different attention layers (heads) are used in parallel.
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
In actual MHA computation, to use the GPU as efficiently as possible, queries are generated at different positions at once so that tokens are processed together, and sequences that do not interact with each other are processed as a batch. In other words, 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
The following assumptions are applied to express performance intuitively.
The performance analyzed under these conditions is as follows.
The lower , the better, because it means a lot of computation is done while reading little data. During training, and are large enough that this ratio is small, so the GPU/TPU can be used efficiently.
During inference, only one token is processed per step. Instead of creating K, V anew each time by multiplying by , 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
The term can be reduced by increasing the batch size, but the term is a structural problem. Reducing this term is the goal of MQA.
Different heads share a single set of Keys and Values.
From an implementation standpoint, it is really simple. Just remove the 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
P_q: [h, d, k] kept as isP_k, P_v: [h, d, k] → [d, k]prev_K, prev_V: [b, h, m, k] → [b, m, k] → size reduced by a factor of By sharing K, V, the problematic term is reduced to , i.e., by a factor of .
| Operations | Memory access | Ratio | |
|---|---|---|---|
| MHA (Batched) | |||
| MHA (Incremental) | |||
| MQA (Incremental) |
Before evaluating quality: since reducing the K, V projections also reduces the number of parameters, the FFN hidden dimension was increased to match the total parameter count of the baseline.
The effect of sharing K, V across heads shows up clearly in the inference step. (Time per token)
| Encoder | Decoder | |
|---|---|---|
| Baseline (MHA) | 1.7μs | 46μs |
| MQA | 1.5μs | 3.8μs |
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.
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 () and only removes the redundancy across heads, so there is less information loss.
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.
This post was translated from the original Korean version with the help of AI, so some errors may remain.