인지야공

인지야공/딥러닝 기초 정리/7번째 글

Attention 은 어떻게 계산되는가 — Q·K·V 와 병렬 계산의 원리

실행: python 03_attention_qkv.py 모든 수치는 실행 결과이며, PyTorch F.scaled_dot_product_attention / nn.MultiheadAttention 과 대조 검증(오차 5.96e-08)했습니다.

한 줄 공식

Attention(Q, K, V) = softmax( Q Kᵀ / √d_k ) V

1. Q, K, V가 도대체 뭔가

도서관 비유가 가장 정확하다.

기호이름도서관 비유역할
QQuery (질의)내가 사서에게 하는 질문“나는 지금 어떤 정보가 필요하지?”
KKey (색인)각 책의 등에 붙은 라벨“나는 이런 내용의 단어다”
VValue (내용)책의 실제 내용“실제로 전달할 정보”

절차는 딱 3단계:

  1. 내 질문(Q) 을 모든 책 라벨(K) 과 비교해서 점수를 매긴다 → 내적
  2. 점수를 확률로 바꾼다 (합 = 1) → softmax
  3. 그 확률로 책 내용(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로 나눈 뒤
42.0252.0001.013
648.1218.0001.015
51222.44422.6270.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 ms4.84e-07 (같은 계산)

47배 빠르다. 그것도 겨우 32토큰에서. N이 커지면 격차는 더 벌어진다.

왜 행렬로 묶으면 빠른가

  1. Q @ Kᵀ 는 (N,D) × (D,N) 단일 GEMM(행렬곱 커널). GPU의 수천 개 코어가 나눠서 동시 처리.
  2. for 루프는 코어 1개가 N×N 번 순차 실행 + 매번 파이썬 인터프리터 오버헤드.
  3. 학습 시에는 디코더도 병렬이다. 정답 문장을 이미 알고 있으므로(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 @ VO(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²) 계산/메모리

표시는 이 브라우저에만 남는다. 서버로 가는 것은 없다.