콘텐츠로 이동

Lab 1. Transformer와 Self-Attention

실습 실행: Colab 노트북 링크를 추가한다.

학습 목표

  • Self-Attention이 문장 안의 토큰 관계를 계산하는 방식을 이해한다
  • Query, Key, Value의 역할을 구분한다
  • Scaled Dot-Product Attention을 PyTorch로 직접 구현한다

핵심 개념

Scaled Dot-Product Attention

각 토큰은 Query로 다른 토큰의 Key와 유사도를 계산하고, 그 가중치로 Value를 합산한다.

\[ \text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V \]
구성 요소 역할
Query (\(Q\)) 지금 토큰이 찾고 싶은 정보
Key (\(K\)) 각 토큰이 가진 정보의 색인
Value (\(V\)) 실제로 전달되는 정보
\(\sqrt{d_k}\) 내적 값이 커져 softmax가 한쪽으로 쏠리는 것을 막는 스케일링

환경 설정

import torch
import torch.nn as nn
import torch.nn.functional as F

torch.manual_seed(42)

Self-Attention 구현

토큰 4개, 임베딩 차원 8인 입력으로 Attention을 계산한다.

batch, n_tokens, d_model = 1, 4, 8
x = torch.randn(batch, n_tokens, d_model)

W_q = nn.Linear(d_model, d_model, bias=False)
W_k = nn.Linear(d_model, d_model, bias=False)
W_v = nn.Linear(d_model, d_model, bias=False)

q, k, v = W_q(x), W_k(x), W_v(x)

scores  = q @ k.transpose(-2, -1) / (d_model ** 0.5)   # (1, 4, 4)
weights = F.softmax(scores, dim=-1)                    # 각 행의 합 = 1
out     = weights @ v                                  # (1, 4, 8)

print(weights.shape, out.shape)
print(weights[0].sum(dim=-1))   # 모든 값이 1

PyTorch 내장 함수와 비교

out_ref = F.scaled_dot_product_attention(q, k, v)
print(torch.allclose(out, out_ref, atol=1e-6))   # True

해석

weights의 \(i\)행은 \(i\)번째 토큰이 다른 토큰들에 얼마나 주목하는지를 나타내는 확률 분포이다. 출력 out의 각 토큰 벡터는 문장 전체의 Value를 이 가중치로 섞은 결과이므로, 문맥 정보를 담게 된다.

과제 1: 스케일링 제거

/ (d_model ** 0.5)를 제거하고 d_model을 512로 늘려 실행하라. weights의 각 행이 어떻게 바뀌는가?

과제 2: 인과 마스크

GPT처럼 앞쪽 토큰만 보도록 마스크를 적용하라. 결과가 F.scaled_dot_product_attention(q, k, v, is_causal=True)와 같은지 확인하라.

mask   = torch.triu(torch.ones(n_tokens, n_tokens, dtype=torch.bool), diagonal=1)
scores = scores.masked_fill(mask, float("-inf"))