지난 글에서 Scaled Dot-Product Attention을 수식부터 코드까지 따라갔다. 거기서 본 계산은 한 번에 하나의 가중치 분포를 만든다. 「나는 사과를 먹었다」에서 「먹었다」라는 토큰의 표현을 만들 때 주어와의 관계, 목적어와의 관계, 시제를 알려 주는 어미와의 관계를 동시에 담아야 하는데 분포가 하나면 그 셋이 한 벌의 가중치를 놓고 다툰다. Multi-Head Attention은 같은 계산을 h개의 낮은 차원 부분 공간에서 병렬로 돌려 이 다툼을 없애는 구조다.
여기서 헤드란 자기만의 질의·키·값 투영 행렬을 갖고 독립적으로 어텐션을 한 번 수행하는 단위를 말한다. 헤드가 쓰는 차원 d_k는 d_model을 h로 나눈 값이고, 이 나눗셈이 이 글에서 다루는 거의 모든 계산의 출발점이다. 이 글은 그 나눗셈에서 시작해 네 가지를 차례로 따라간다. 왜 나눠야 하는가, 코드에서 어떻게 나누는가, 나눠 놓은 헤드가 실제로 무엇을 하는가, 그리고 그 나눔이 메모리와 속도에 어떤 값을 매기는가다.
단일 헤드의 한계
분포 하나의 제약
어텐션 가중치는 소프트맥스를 지난 값이라 한 행의 합이 1이다. 합이 정해진 예산을 나눠 쓰는 것이므로, 「먹었다」가 「나는」에 0.5를 주면 나머지 전부에 0.5밖에 못 준다. 서로 다른 성격의 관계 셋을 한 분포로 표현하려면 셋을 평균한 애매한 가중치가 나오고, 그 평균은 어느 관계도 제대로 담지 못한다.
헤드를 여럿 두면 예산이 헤드마다 따로 생긴다. 한 헤드가 주어에 몰아 주는 동안 다른 헤드는 목적어에 몰아 줄 수 있고, 둘의 결과는 뒤에서 합쳐진다. 표현력을 늘렸다기보다 예산을 쪼개 경쟁을 없앤 것에 가깝다. 한 헤드의 차원을 여덟 배로 키우는 것과 헤드를 여덟 개 두는 것을 견줘 보면 차이가 분명해진다. 앞쪽은 여전히 분포가 하나라 예산 다툼이 그대로 남고, 뒤쪽은 차원 총합이 같은데도 분포가 여덟 개 생긴다. 실제로 원논문의 절제 실험에서 헤드 하나짜리 설정은 헤드 여덟짜리보다 번역 점수가 낮았고, 파라미터 수는 두 설정이 같았다.
차원을 나누는 계산
헤드를 여덟 개 두면서 각 헤드가 d_model 전체를 쓰면 연산량이 여덟 배가 된다. 원논문은 그렇게 하지 않고 d_k를 d_model / h로 잡았다. 512와 8이면 각 헤드는 64차원을 쓴다.
이 선택 덕분에 총 연산량이 단일 헤드와 거의 같아진다. 질의·키·값 투영은 헤드마다 d_model × d_k인데 h개를 합치면 d_model × d_model이라 h와 무관하고, 어텐션 점수 계산도 헤드마다 T × T × d_k인데 h를 곱하면 T × T × d_model이 되어 역시 h가 사라진다. 헤드 수는 공짜로 늘릴 수 있는 값이며, 늘릴 때 잃는 것은 연산량이 아니라 헤드당 차원이다. 「공짜」에는 단서가 하나 붙는다. 점수 텐서의 원소 수는 B × h × T × T이므로 헤드 수에 정비례해 늘어난다. 곱셈의 총량은 그대로인데 중간에 들고 있어야 하는 값은 늘어나는 것이고, 이 어긋남이 뒤의 「비용의 두 축」에서 다시 나온다.
W_O가 하는 일
헤드 h개의 출력을 이어 붙이면 길이 d_model인 벡터가 되지만 그 벡터는 아직 여덟 조각이 나란히 놓인 상태다. 앞의 64칸은 첫 헤드만, 다음 64칸은 둘째 헤드만 만든 값이라 조각 사이에 어떤 섞임도 없다. 출력 투영 행렬 W_O가 이 조각들을 선형 결합해 하나의 표현으로 만든다.
W_O를 빼고 이어 붙이기만 하면 각 헤드가 출력 벡터의 정해진 구간만 담당하는 구조가 되어, 헤드가 서로 보완하지 못하고 다음 레이어도 조각 경계에 맞춰 학습해야 한다. 이어 붙이기는 자리를 만드는 일이고 섞는 일은 W_O 몫이다. 실제 구현에서 이 행렬은 d_model × d_model 짜리 선형 계층 하나로 들어가 있어 눈에 잘 띄지 않는다. 코드에서 W_Q·W_K·W_V와 나란히 놓인 네 번째 선형 계층이 그것이며, 파라미터의 4분의 1을 여기가 쓴다.
헤드 분할 구현
view와 transpose
구현에서 헤드를 나누는 일은 행렬을 h개 따로 만드는 것이 아니라 하나의 큰 텐서를 다르게 보는 것이다. (B, T, d_model) 텐서에 view(B, T, h, d_k)를 걸면 마지막 축만 둘로 쪼개지고, transpose(1, 2)로 축 순서를 바꾸면 (B, h, T, d_k)가 된다. 두 연산 모두 메모리를 복사하지 않고 원소를 읽는 규칙만 바꾼다.
문제는 transpose 뒤의 텐서가 메모리에서 연속이 아니라는 점이다. 헤드를 다시 합칠 때 view를 부르면 연속 텐서를 요구하다 오류가 나므로, 그 직전에 contiguous()로 한 번 복사해야 한다. 앞의 분할에서는 필요 없고 뒤의 병합에서만 필요하다는 비대칭이 여기서 나온다. reshape를 쓰면 필요한 경우에만 알아서 복사하므로 오류는 안 나지만, 어디서 복사가 일어나는지 코드에 안 보이게 된다. 메모리 사용을 따져야 하는 자리에서는 contiguous와 view를 나눠 적는 편이 읽기에 낫다.
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, num_heads=8, dropout=0.1):
super().__init__()
assert d_model % num_heads == 0
self.h = num_heads
self.dk = d_model // num_heads
self.W_Q = nn.Linear(d_model, d_model)
self.W_K = nn.Linear(d_model, d_model)
self.W_V = nn.Linear(d_model, d_model)
self.W_O = nn.Linear(d_model, d_model)
self.drop = nn.Dropout(dropout)
def split_heads(self, x, B, T):
return x.view(B, T, self.h, self.dk).transpose(1, 2)
def forward(self, Q, K, V, mask=None):
B = Q.size(0)
q = self.split_heads(self.W_Q(Q), B, Q.size(1))
k = self.split_heads(self.W_K(K), B, K.size(1))
v = self.split_heads(self.W_V(V), B, V.size(1))
scores = q @ k.transpose(-2, -1) / math.sqrt(self.dk)
if mask is not None:
scores = scores.masked_fill(mask, float('-inf'))
alpha = self.drop(F.softmax(scores, dim=-1))
ctx = (alpha @ v).transpose(1, 2).contiguous()
ctx = ctx.view(B, -1, self.h * self.dk)
return self.W_O(ctx), alpha
마스크의 모양
점수 텐서는 (B, h, T, T)인데 마스크를 그 모양으로 만들 필요는 없다. 가릴 자리는 헤드마다 같기 때문이다. (B, 1, T, T)로 두면 브로드캐스트가 헤드 축을 알아서 늘려 주고, 메모리도 h분의 1만 쓴다.
마스크에서 자주 나는 실수는 참과 거짓의 방향이다. masked_fill은 참인 자리를 채우므로 마스크는 「가릴 곳이 참」이어야 한다. 패딩 마스크를 만들 때 유효 토큰을 참으로 표시해 두고 그대로 넘기면 볼 곳만 가리는 정반대 동작이 되는데, 오류가 나지 않고 손실만 이상하게 높은 채로 학습이 진행되므로 찾기 어렵다. 마스크가 제대로 걸렸는지는 학습 전에 한 번 확인할 수 있다. 짧은 배치를 하나 통과시켜 가중치 텐서를 꺼내고, 가려야 할 자리의 값이 정확히 0인지 보면 된다. 소프트맥스에 음의 무한대를 넣었으므로 0이 아니라 아주 작은 수가 나온다면 마스크가 덜 걸린 것이다. 인과 마스크와 패딩 마스크를 함께 걸어야 하는 자리에서는 둘을 논리합으로 합친 뒤 한 번에 넣는다. 따로 두 번 적용하면 두 번째 masked_fill이 이미 음의 무한대가 들어간 자리를 다시 덮으면서, 한 행이 통째로 가려지는 경우에 소프트맥스가 NaN을 내놓는다. 패딩 토큰이 질의 쪽에 있을 때 실제로 생기는 상황이다.
단계별 모양
배치 32, 길이 128, d_model 512, 헤드 8로 두고 한 번 통과시키면 모양이 이렇게 움직인다.
| 단계 | 모양 |
|---|---|
| 입력 | (32, 128, 512) |
W_Q 통과 후 |
(32, 128, 512) |
view로 쪼갬 |
(32, 128, 8, 64) |
transpose |
(32, 8, 128, 64) |
| 점수 | (32, 8, 128, 128) |
| 가중합 | (32, 8, 128, 64) |
transpose + view |
(32, 128, 512) |
W_O 통과 후 |
(32, 128, 512) |
들어간 모양과 나온 모양이 같고, 중간에 한 번만 다른 모양이 나타난다. 그 한 줄이 점수 텐서이고 길이가 제곱으로 들어간 유일한 자리다. 디버깅할 때 이 표를 옆에 두고 실제 텐서의 shape를 찍어 대조하면 대부분의 오류가 한 줄 안에서 잡힌다. 특히 셀프 어텐션이 아닌 크로스 어텐션에서는 질의의 길이와 키의 길이가 달라 점수 텐서가 정사각이 아니게 되는데, 이 표를 정사각으로 외워 두면 그 자리에서 모양이 안 맞는 이유를 못 찾는다.
헤드의 역할 분담
가중치 행렬 읽기
학습된 모델에서 헤드가 무엇을 하는지는 가중치 행렬을 꺼내 보면 대강 분류할 수 있다. 문장 여러 개를 통과시켜 헤드마다 (T, T) 행렬을 모으고, 대각선 근처에 몰린 비율, 바로 앞뒤 토큰에 몰린 비율, 문장 첫 토큰에 몰린 비율을 재는 방식이다.
이렇게 세어 보면 인접 토큰에 붙는 헤드, 구문상 걸리는 토큰을 찾는 헤드, 문장 첫 토큰에 대부분의 가중치를 주는 헤드가 나뉜다. 마지막 유형은 볼 곳이 마땅치 않을 때 가중치를 버리는 자리로 첫 토큰을 쓰는 것이라 해석되며, 소프트맥스 합이 1이라는 제약 때문에 어딘가에는 주어야 해서 생기는 현상이다. 이 특화는 지시해서 생긴 것이 아니라 언어 모델링 목표만으로 나타난다. 다만 이런 분류를 해석으로 받아들일 때는 조심할 필요가 있다. 가중치가 어디에 몰렸는지는 그 헤드가 무엇을 계산하는지의 일부일 뿐이고, 실제 출력에 얼마나 기여했는지는 값 벡터의 크기까지 함께 봐야 알 수 있다. 가중치가 고르게 퍼진 헤드가 사실상 평균을 내고 있을 수도 있다.
절제 실험
헤드가 실제로 필요한지 확인하는 방법은 하나씩 0으로 만들고 성능을 다시 재는 것이다. 이렇게 해 보면 지워도 점수가 거의 안 떨어지는 헤드가 상당수 나오고, 반대로 하나만 지워도 특정 작업이 크게 무너지는 헤드가 소수 나온다. 학습을 마친 모델에서 많은 헤드를 잘라 내도 품질이 유지된다는 보고가 여럿이지만, 같은 모델을 처음부터 적은 헤드로 학습하면 같은 품질에 닿지 못한다.
학습에는 여러 헤드가 필요하고 추론에는 덜 필요하다는 뜻이다. 헤드를 잘라 내는 가지치기가 추론 최적화 기법으로 다뤄지고 학습 설정으로는 다뤄지지 않는 이유가 여기에 있다. 학습 중에는 여러 헤드가 서로 다른 후보를 시도하는 탐색 장치처럼 쓰이고, 끝나고 나면 그중 일부만 실제 계산을 떠받친다고 보는 해석이 이 관찰과 맞는다.
d_k가 작아질 때
헤드를 늘리면 d_k가 줄어든다. d_model이 512일 때 h를 32로 올리면 d_k는 16이다. 16차원 공간에서 질의와 키의 내적으로 관계를 구분해야 하는데, 차원이 낮으면 서로 다른 토큰의 키 벡터가 비슷해지기 쉬워 가중치가 평평해진다. 평평한 분포는 사실상 평균을 내는 것이라 어떤 관계도 고르지 않은 것과 같고, 그런 헤드가 여럿 생기면 헤드를 늘린 만큼의 이득이 사라진다.
헤드당 차원을 64에서 128 사이에 두는 관행은 이 선에서 나왔다. 헤드 수를 정할 때 실제로 고르는 값은 헤드 수가 아니라 헤드당 차원이고, 모델 폭이 정해지면 헤드 수는 거기서 따라 나온다. 반대 방향으로 밀어붙인 사례도 있다. 헤드당 차원을 256 이상으로 크게 잡고 헤드 수를 줄이면 분포 수가 적어져 관계를 나눠 담기 어려워진다. 위아래 양쪽에 벽이 있고 그 사이가 64에서 128이라는 좁은 띠다.
비용의 두 축
파라미터와 활성값
헤드 수를 바꿀 때 무엇이 변하고 무엇이 그대로인지를 가르는 것이 이 구조를 다루는 핵심이다. 파라미터는 투영 행렬 넷뿐이고 각각 d_model × d_model이라 h가 등장하지 않는다. 반면 학습 중에 들고 있어야 하는 활성값에는 점수 텐서가 있고 이것은 B × h × T × T라 h에 정비례한다.
d_model, num_heads = 512, 8
params_per_mha = 4 * d_model * d_model
print(f"MHA 파라미터: {params_per_mha:,}") # 1,048,576
mha = nn.MultiheadAttention(embed_dim=512, num_heads=8, batch_first=True)
print(sum(p.numel() for p in mha.parameters())) # 편향 포함 1,050,624
헤드를 8에서 16으로 올려도 위 숫자는 그대로이고, 점수 텐서만 두 배가 된다. 「헤드를 늘렸더니 메모리가 터졌다」는 보고가 파라미터 수를 세어 보고 납득이 안 되는 이유다.
길이가 끌어올리는 것
배치 1, 헤드 8, 16비트로 두고 점수 텐서 크기를 길이만 바꿔 계산하면 차이가 분명해진다. 레이어 하나 기준이다.
| 길이 | 점수 텐서 원소 수 | 크기 |
|---|---|---|
| 512 | 209만 | 4MB |
| 2,048 | 3,355만 | 64MB |
| 8,192 | 5억 3,687만 | 1,024MB |
길이를 네 배로 늘리면 열여섯 배가 된다. 여기에 레이어 수와 배치를 곱해야 하므로, 32층 모델에 배치 8이면 길이 8,192에서 점수 텐서만 수백 기가바이트가 된다. 문맥 창을 늘리는 일이 설정 하나 고치는 일이 아닌 이유다. 추론에서는 사정이 조금 낫다. 새 토큰 하나를 만들 때 필요한 점수는 (1, T) 한 줄이라 제곱이 아니라 길이에 비례하고, 대신 앞에서 본 KV 캐시가 길이에 비례해 쌓인다. 학습에서 무거운 것은 점수 텐서이고 추론에서 무거운 것은 캐시라는 구분을 해 두면 최적화 기법이 어느 쪽을 겨냥한 것인지 바로 읽힌다.
행렬을 만들지 않는 길
FlashAttention은 이 곱셈의 전제를 깬다. 점수 행렬 전체를 메모리에 두고 소프트맥스를 걸어 값과 곱하는 대신, 키와 값을 작은 조각으로 나눠 조각마다 부분 결과를 구하고 소프트맥스의 정규화 항을 이어 가며 갱신한다. 최종 결과는 같은데 (T, T) 행렬이 한 번도 통째로 만들어지지 않는다.
메모리가 길이의 제곱이 아니라 길이에 비례하게 되고, 느린 메모리를 오가는 횟수가 줄어 속도까지 붙는다. 연산량 자체는 그대로라는 점은 짚어 둘 만하다. 줄어든 것은 메모리와 메모리 이동이지 곱셈의 수가 아니다. 그런데도 실제 학습 속도가 몇 배 붙는 것은 지금의 GPU에서 병목이 곱셈이 아니라 메모리 대역폭이기 때문이다. 연산 장치는 놀고 있고 데이터가 오기를 기다리는 시간이 대부분인 상황에서, 오가는 양을 줄이는 것이 곱셈을 줄이는 것보다 크게 먹힌다.
KV 공유 갈래
캐시 크기 식
생성할 때는 이미 계산한 키와 값을 저장해 두고 새 토큰의 질의만 그것들과 맞춘다. 저장해 두는 양이 KV 캐시이고 크기는 배치 × 레이어 × 2 × 헤드 수 × 길이 × d_k × 바이트 수다. 헤드 수가 곱해져 있다는 점이 중요하다. 파라미터와 달리 캐시는 헤드 수에 정비례한다.
키와 값의 헤드 수를 g로 줄이면 캐시도 g에 비례해 준다. 질의 헤드는 h 그대로 두고 키와 값만 줄이는 것이 이 갈래의 전부다. 질의를 줄이지 않는 이유도 같은 식에서 나온다. 질의는 새 토큰 하나에 대해서만 계산하고 버리므로 저장할 것이 없고, 캐시에 쌓이는 것은 키와 값뿐이다.
| 방식 | 질의 헤드 | 키·값 헤드 | 캐시 |
|---|---|---|---|
| MHA | h | h | 기준 |
| GQA | h | g | h/g 배 감소 |
| MQA | h | 1 | h배 감소 |
g를 고르는 자리
g를 1까지 내리면 캐시는 최소가 되지만 품질이 떨어진다. 실무에서 g는 대개 4에서 8 사이에 놓이는데, 캐시 절감이 g에 반비례해 가파르게 떨어지는 데 비해 품질 하락은 g가 작아질 때 급해지기 때문이다. 헤드 32개를 g 8로 줄이면 캐시가 4분의 1이 되고 품질 손실은 재기 어려울 만큼 작다는 보고가 많아, 이 근방이 기본값이 됐다. 이미 학습한 MHA 모델을 GQA로 바꾸는 방법도 쓰인다. 같은 그룹에 들어갈 키·값 헤드들의 가중치를 평균해 초기값으로 삼고 원래 학습량의 몇 퍼센트만 추가 학습하는 방식인데, 처음부터 GQA로 학습한 것에 가까운 품질에 닿는다는 보고가 있다.
MQA가 무너지는 자리
키와 값을 하나로 합치면 모든 질의 헤드가 같은 키 공간을 보게 된다. 여러 관계를 동시에 봐야 하는 작업, 특히 긴 입력에서 특정 사실을 집어내야 하는 작업에서 손실이 눈에 띈다. 요약처럼 문장 전체의 분위기를 반영하는 작업은 덜 민감하고, 긴 문서에서 정확한 한 줄을 찾아 인용해야 하는 작업은 민감하다. 품질을 재는 지표가 평균 점수 하나뿐이면 이 차이가 안 보인다는 점도 함께 짚어 둘 만하다. 전체 벤치마크 점수는 0.2점 떨어지는데 긴 문맥 검색 정확도는 크게 떨어지는 식이라, 캐시를 줄이는 변경을 넣을 때는 대상 작업에 맞는 지표를 따로 세워 두고 재야 한다. 자세한 비교는 Multi-Query와 Grouped-Query Attention에서 다룬다.
연습 문제
d_model이 4096이고 헤드당 차원을 128로 두려 한다. 헤드 수와 투영 행렬 넷의 파라미터 수를 구하라.헤드 수는 4096 / 128로 32다. 파라미터는 4 × 4096 × 4096으로 6,710만 8,864개이며 헤드 수와 무관하다.같은 모델에서 헤드를 32에서 64로 늘렸다. 파라미터 수와 점수 텐서 크기는 각각 어떻게 되는가.
파라미터는 그대로다. 헤드당 차원이 64로 줄 뿐 투영 행렬의 크기는 변하지 않는다. 점수 텐서는 헤드 축이 두 배가 되므로 두 배로 늘어난다.레이어 32개, 헤드 32개, 헤드당 차원 128, 길이 4,096, 배치 1, 16비트 기준으로 KV 캐시 크기를 구하라.
1 × 32 × 2 × 32 × 4096 × 128 × 2바이트로 약 68억 7,200만 바이트, 곧 6.4GB다. 같은 조건에서 키와 값 헤드를 8로 줄이면 1.6GB가 된다.W_O를 항등 행렬로 고정하면 무엇이 깨지는가.헤드 출력이 섞이지 않는다. 출력 벡터의 각 구간을 헤드 하나가 독점하게 되어 헤드끼리 보완할 수 없고, 다음 레이어가 그 구간 경계에 맞춰 학습해야 한다.
헤드를 나누는 구조를 잡았으니 다음 글에서는 어텐션이 순서를 전혀 모른다는 문제로 넘어간다.
읽어주셔서 감사합니다. 😊

