
Transformer는 attention layer로 sequence 내부 및 sequence 간 정보를 주고받는다. 학습할 때는 전체 sequence를 한 번에 병렬로 처리할 수 있어 빠르지만, incremental inference(추론) 과정에서는 이와 같은 병렬화를 적용할 수 없다.
해당 논문에서는 MQA 방식을 통하여 읽어오는 K, V cache의 크기를 줄이는 방식으로 bottleneck을 해결한다.
해당 논문은 모든 연산을 tensor 사이 일반화된 contraction인 einsum notation으로 표현한다. 이 표기법 덕분에 각 텐서의 shape과 연산 흐름이 명확하게 보인다.
Query 하나에 대한 attention이다.
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)
개의 서로 다른 attention layer(head)를 병렬로 사용한다.
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
실제 MHA 연산시에는 GPU를 최대한 효율적으로 활용하기 위하여 개의 서로 다른 position에서 query를 한번에 생성하여 개의 token을 한번에 처리하고 서로 상호작용하지 않는 개의 sequence를 batch로 처리한다. 즉, 한 step에서 개의 token을 처리하는 것이다.
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 표현을 위하여 다음의 가정이 적용되었다.
위의 조건을 바탕으로 분석한 performance는 다음과 같다.
는 낮을수록 좋다. 적은 데이터를 읽어서 많은 연산을 한다는 뜻이기 때문이다. 학습 때는 와 이 충분히 커서 이 비율이 작고, 따라서 GPU/TPU를 효율적으로 쓸 수 있다.
추론 시에는 한 step에 token 하나만 처리한다. 이때 매번 에 를 곱해서 K, V를 새로 만드는 것이 아니라, 이전 position까지의 K, V(prev_K, prev_V)에 현재 token의 값만 이어 붙인다.
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
항은 batch를 키우면 줄일 수 있지만, 항은 구조적인 문제다. 이 항을 줄이는 것이 MQA의 목표이다.
서로 다른 head들이 하나의 Key, Value 세트를 공유한다.
구현 관점에서는 정말 단순하다. Multi-head 코드에서 K, V와 관련된 텐서의 차원을 지우면 된다.
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] # h 차원이 사라짐
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는 여전히 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) # 모든 head가 같은 K를 봄
weights = tf.softmax(logits)
o = tf.einsum("bhm,bmv->bhv", weights, new_V) # 모든 head가 같은 V를 봄
y = tf.einsum("bhv,hdv->bd", o, P_o)
return y, new_K, new_V
P_q: [h, d, k] 그대로 유지P_k, P_v: [h, d, k] → [d, k]prev_K, prev_V: [b, h, m, k] → [b, m, k] → 크기가 배문제였던 항이 K, V를 공유함으로써 , 즉 배로 줄어든다.
| Operations | Memory access | Ratio | |
|---|---|---|---|
| MHA (Batched) | |||
| MHA (Incremental) | |||
| MQA (Incremental) |
Quality 평가에 앞서 K, V projection이 줄어든 만큼 파라미터 수가 줄어들기 때문에, baseline과 전체 파라미터 수를 맞추기 위해 FFN의 hidden dimension 를 키웠다.
K, V를 head끼리 공유한 효과는 inference step에서 확실하게 드러난다. (Time per token)
| Encoder | Decoder | |
|---|---|---|
| Baseline (MHA) | 1.7μs | 46μs |
| MQA | 1.5μs | 3.8μs |
MQA는 결국 K, V를 head별로 통일하는 방식이다. 서로 다른 head는 token 간의 서로 다른 상관관계를 학습하기 때문에, K, V를 head별로 통일하면 성능이 훨씬 많이 떨어질 거라고 예상했다. 그런데 실제로는 baseline과 매우 근소한 차이만 났다. 이를 보면 attention에서는 결국 현재 token이 요구하는 정보인 Query가 가장 영향력이 큰 것으로 보인다. head마다 "무엇을 물어볼지"만 다르면, "어디서 찾을지(K)"와 "무엇을 가져올지(V)"는 공유해도 충분히 다양한 패턴을 만들 수 있다는 뜻이다.
head 수를 줄이거나 K, V의 dimension을 줄였을 때 성능이 낮아지는 것은 K, V 자체가 담는 정보가 줄어들기 때문으로 보인다. 이 경우 FFN dimension을 늘려서 파라미터 수를 맞춰도 token 간 정보 공유는 attention에서만 일어나기 때문에 그 손실을 메우지 못한다. 반면 MQA는 K, V의 dimension()은 유지하고 head 간 중복만 제거하기 때문에 정보 손실이 적다.
논문에서 사용한 WMT14 번역, Billion-Word LM 등은 지금 기준으로 다소 오래되고 규모가 작은 task다. 그래서 dataset 규모가 크거나 더 어려운 task에서도 MQA가 좋은 성능을 보일지는 의문이 남는다. 다만 실제로 PaLM, Falcon 등 대형 모델에서 MQA가 채택되었고, LLaMA-2 70B는 MQA를 일반화한 GQA를 사용하는 것을 보면 대규모에서도 충분히 쓸 만한 trade-off로 받아들여진 것 같다.