인지야공/딥러닝 기초 정리/7번째 글
Attention 은 어떻게 계산되는가 — Q·K·V 와 병렬 계산의 원리
실행:
python 03_attention_qkv.py모든 수치는 실행 결과이며, PyTorchF.scaled_dot_product_attention/nn.MultiheadAttention과 대조 검증(오차 5.96e-08)했습니다.
한 줄 공식
Attention(Q, K, V) = softmax( Q Kᵀ / √d_k ) V
1. Q, K, V가 도대체 뭔가
도서관 비유가 가장 정확하다.
| 기호 | 이름 | 도서관 비유 | 역할 |
|---|---|---|---|
| Q | Query (질의) | 내가 사서에게 하는 질문 | “나는 지금 어떤 정보가 필요하지?” |
| K | Key (색인) | 각 책의 등에 붙은 라벨 | “나는 이런 내용의 단어다” |
| V | Value (내용) | 책의 실제 내용 | “실제로 전달할 정보” |
절차는 딱 3단계:
- 내 질문(Q) 을 모든 책 라벨(K) 과 비교해서 점수를 매긴다 → 내적
- 점수를 확률로 바꾼다 (합 = 1) → softmax
- 그 확률로 책 내용(V) 을 섞는다 → 가중평균
왜 K와 V를 나누나? “찾는 기준”과 “가져올 내용”이 달라도 되기 때문이다. 예: “그것”이라는 단어는 K로는 “지시대명사”라고 광고하지만, V로는 “실제 가리키는 대상의 의미”를 전달할 수 있다.
실제 어디서 나오나
Q, K, V는 하늘에서 떨어지는 게 아니라, 같은 입력 x에 서로 다른 가중치를 곱해 만든다.
Q = x @ W_q # (N, d_model) @ (d_model, d_k)
K = x @ W_k
V = x @ W_v
- 셀프 어텐션: Q, K, V가 전부 같은 문장에서 나옴 (문장이 자기 자신을 봄)
- 크로스 어텐션: Q는 디코더에서, K·V는 인코더 출력에서 (번역할 때 원문 참조)
2. 손계산 — 토큰 3개로 끝까지
토큰: ["나는", "학교에", "간다"], d_k = 2
Q = [[1, 0], K = [[1, 0], V = [[10, 0],
[0, 1], [0, 1], [ 0, 10],
[1, 1]] [1, 1]] [ 5, 5]]
(1) 점수 = Q Kᵀ
i행 j열 = “i번째 단어가 j번째 단어를 볼 점수”
K1 K2 K3
Q1 → [ 1 0 1 ] Q1·K1 = 1×1+0×0 = 1
Q2 → [ 0 1 1 ] Q1·K2 = 1×0+0×1 = 0
Q3 → [ 1 1 2 ] Q3·K3 = 1×1+1×1 = 2
내적이 크다 = 두 벡터가 같은 방향 = 비슷하다 = 서로 봐야 한다
(2) √d_k 로 나누기
√2 = 1.4142
[[0.7071, 0.0000, 0.7071],
[0.0000, 0.7071, 0.7071],
[0.7071, 0.7071, 1.4142]]
(3) softmax (행 방향, dim=-1)
[[0.4011, 0.1978, 0.4011], ← 합 = 1
[0.1978, 0.4011, 0.4011], ← 합 = 1
[0.2483, 0.2483, 0.5035]] ← 합 = 1
“나는”은 자기 자신(40%)과 “간다”(40%)를 주로 보고, “학교에”는 20%만 본다는 뜻.
(4) V의 가중평균
"나는" = 0.401×[10,0] + 0.198×[0,10] + 0.401×[5,5] = [6.017, 3.983]
"학교에" = 0.198×[10,0] + 0.401×[0,10] + 0.401×[5,5] = [3.983, 6.017]
"간다" = 0.248×[10,0] + 0.248×[0,10] + 0.503×[5,5] = [5.000, 5.000]
✔ PyTorch F.scaled_dot_product_attention 결과와 오차 2.38e-07 (동일)
3. 왜 하필 √d_k 로 나누는가
이유: softmax 포화(saturation)를 막기 위해.
Q, K 원소가 평균 0, 분산 1이면 내적 q·k 의 분산은 d_k에 비례한다.
따라서 표준편차는 √d_k. 실제로 측정한 값:
| d_k | 내적 표준편차 | 이론값 √d_k | √d_k로 나눈 뒤 |
|---|---|---|---|
| 4 | 2.025 | 2.000 | 1.013 |
| 64 | 8.121 | 8.000 | 1.015 |
| 512 | 22.444 | 22.627 | 0.992 |
나누지 않으면 무슨 일이 생기나:
점수 [10, 2, 1] 을 그대로 softmax → [0.9995, 0.0003, 0.0001] 거의 한 곳에 몰림
같은 점수를 √64로 나눈 뒤 softmax → [0.591, 0.217, 0.192] 골고루 분포
softmax가 0/1로 포화되면 그 지점의 기울기가 거의 0 → 역전파 신호가 안 흐름 → 학습 정지. √d_k로 나누면 차원이 4든 512든 항상 비슷한 크기의 점수 분포가 유지된다.
4. 병렬 계산이 되는 원리 ★ 핵심
왜 RNN/LSTM은 병렬이 안 되나
h₁ = f(x₁, h₀)
h₂ = f(x₂, h₁) ← h₁ 이 나와야만 계산 가능
h₃ = f(x₃, h₂) ← h₂ 를 기다려야 함
시간축에 순차 의존성이 있다. 문장 길이가 100이면 100번을 줄 서서 계산해야 한다. GPU에 코어가 1만 개 있어도 소용없다.
왜 Attention은 병렬이 되나
out₁ = softmax(q₁Kᵀ/√d)V
out₂ = softmax(q₂Kᵀ/√d)V ← out₁ 이 전혀 필요 없다!
out₃ = softmax(q₃Kᵀ/√d)V ← 셋 다 동시에 계산 가능
모든 출력이 오직 (Q, K, V) 전체에만 의존하고, 서로를 참조하지 않는다. 그래서 for 루프를 통째로 하나의 행렬곱으로 접을 수 있다.
# 느린 버전: 토큰 하나씩 (개념용)
for i in range(N):
for j in range(N):
score[j] = dot(Q[i], K[j]) / sqrt(D)
out[i] = softmax(score) @ V
# 실제로 쓰는 버전: 딱 3줄
scores = Q @ K.transpose(-2, -1) / math.sqrt(Q.size(-1)) # (B, N, N)
attn = F.softmax(scores, dim=-1)
out = attn @ V # (B, N, D)
실측 (N=32, D=64, CPU):
| 방식 | 시간 | 결과 차이 |
|---|---|---|
| for 루프 | 9.73 ms | — |
| 행렬곱 | 0.21 ms | 4.84e-07 (같은 계산) |
47배 빠르다. 그것도 겨우 32토큰에서. N이 커지면 격차는 더 벌어진다.
왜 행렬로 묶으면 빠른가
Q @ Kᵀ는(N,D) × (D,N)단일 GEMM(행렬곱 커널). GPU의 수천 개 코어가 나눠서 동시 처리.- for 루프는 코어 1개가
N×N번 순차 실행 + 매번 파이썬 인터프리터 오버헤드. - 학습 시에는 디코더도 병렬이다. 정답 문장을 이미 알고 있으므로(teacher forcing) 전체를 한 번에 넣고 마스크로 미래만 가린다.
⚠️ 단, 추론(생성)은 순차다. 한 글자를 만들어야 다음 글자를 만들 수 있다. 그래서 이미 계산한 K, V를 저장해두는 KV 캐시로 보완한다. “Transformer가 빠르다”는 말은 학습이 빠르다는 뜻이다.
5. 마스크 — 미래를 못 보게 가리기
Causal Mask (디코더용)
i번째 토큰이 i보다 뒤를 못 보게 막는다. 안 막으면 정답을 컨닝한다.
mask = torch.triu(torch.ones(N, N, dtype=torch.bool), diagonal=1)
scores = scores.masked_fill(mask, float("-inf"))
왜 -inf 인가? exp(-inf) = 0 이므로 softmax 확률이 정확히 0이 된다.
0을 곱하는 게 아니라 softmax 전에 -inf를 넣는 게 핵심 (0을 곱하면 다른 확률의 합이 1이 안 됨).
실행 결과:
마스크 없음 Causal 마스크 적용
[0.5447 0.2022 0.1807 0.0724] [1.0000 0.0000 0.0000 0.0000]
[0.1293 0.0724 0.3665 0.4318] [0.6411 0.3589 0.0000 0.0000]
[0.3523 0.0342 0.3032 0.3103] [0.5108 0.0496 0.4396 0.0000]
[0.0453 0.7613 0.0712 0.1222] [0.0453 0.7613 0.0712 0.1222]
첫 토큰은 자기 자신만 보므로 확률이 1.0. 아래로 갈수록 볼 수 있는 범위가 넓어진다.
Padding Mask
길이가 다른 문장을 배치로 묶으면 빈칸(PAD)이 생긴다. 그 자리를 무시해야 한다.
[0.7293 0.2707 0.0000 0.0000] ← 뒤 2칸이 패딩. 확률 0
6. Multi-Head Attention — 여러 관점으로 동시에
d_model=512를 head 8개로 쪼개면 head당 d_k = 64.
왜 쪼개나? head마다 다른 종류의 관계를 학습하기 위해서다.
- head 1: 문법 관계(주어-동사)
- head 2: 지시 대상(“그것” → 실제 명사)
- head 3: 인접 단어
- …
핵심: 쪼개도 총 연산량은 거의 같다. 8 × 64 = 512 이므로.
그리고 reshape 하나로 모든 head를 한 번의 배치 행렬곱으로 동시 처리한다.
# (B,N,D) -> (B,N,h,d_k) -> (B,h,N,d_k) head 축을 배치처럼 취급
q = self.w_q(x).view(B, N, h, d_k).transpose(1, 2)
k = self.w_k(x).view(B, N, h, d_k).transpose(1, 2)
v = self.w_v(x).view(B, N, h, d_k).transpose(1, 2)
scores = q @ k.transpose(-2, -1) / math.sqrt(d_k) # (B, h, N, N) ★모든 head 동시
attn = F.softmax(scores, dim=-1)
ctx = attn @ v # (B, h, N, d_k)
# head 다시 이어붙이기
ctx = ctx.transpose(1, 2).contiguous().view(B, N, D)
out = self.w_o(ctx) # 마지막 선형변환
transpose(1,2) 로 head를 배치 차원 쪽으로 보내는 것이 병렬화의 트릭이다.
(B, h, N, d_k) 에서 앞 두 축이 배치처럼 취급되어 B×h 개의 행렬곱이 한 번에 돈다.
✔ 직접 만든 구현과 nn.MultiheadAttention 의 최대 오차: 5.96e-08 (동일)
마지막 W_o 가 필요한 이유: head들을 그냥 이어붙이면 각자 따로 논다.
W_o 가 head들의 결과를 섞어서 하나의 표현으로 통합한다.
7. 계산 복잡도 — Transformer의 약점
| 연산 | 복잡도 |
|---|---|
Q @ Kᵀ | O(N² · D) |
softmax @ V | O(N² · D) |
| attention 행렬 메모리 | O(N²) |
N(토큰 수)이 2배가 되면 계산과 메모리가 4배가 된다. 문장 4096토큰이면 attention 행렬만 4096² = 1600만 개.
그래서 나온 해법들:
- FlashAttention: N² 행렬을 통째로 메모리에 안 올리고 타일 단위로 계산 (수학적으로 동일, 메모리 O(N))
- Sliding Window / Sparse Attention: 가까운 토큰만 보기
- Linear Attention: softmax를 근사해 O(N)으로
8. 한 장 요약
1. Q = xW_q, K = xW_k, V = xW_v 같은 입력에서 3가지 역할 뽑기
2. scores = Q Kᵀ 누가 누구를 볼지 점수
3. scores /= √d_k softmax 포화 방지
4. scores.masked_fill(mask, -inf) 미래/패딩 가리기
5. attn = softmax(scores, dim=-1) 확률로 변환 (행 합 = 1)
6. out = attn @ V 내용 가중평균
7. out = out @ W_o (멀티헤드면) 통합
병렬화 이유 = 출력들 사이에 의존성이 없어서 for 루프를 GEMM 하나로 접을 수 있음
병렬화 방법 = head 축을 배치 축으로 보내는 reshape + transpose
대가 = O(N²) 계산/메모리