-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmodel.py
More file actions
308 lines (269 loc) · 17.6 KB
/
Copy pathmodel.py
File metadata and controls
308 lines (269 loc) · 17.6 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
"""kopus 모델 — 밑바닥부터 만드는 한국어 GPT.
구성: RMSNorm(pre-norm) · RoPE · GQA · SwiGLU · tied embedding.
bias 없음, dropout 없음 (110M을 0.5B 토큰으로 학습하는 이 예산에서는 과적합이
아니라 과소적합이 문제라 정규화 장치가 필요 없다).
옵티마이저 구성은 train.py 소관 — 이 파일은 순수하게 모델만 담는다.
"""
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
# 모델의 "형태"를 결정하는 키들. train.py의 재개 로직과 sample.py의 모델 재구성이
# 이 튜플 하나를 공유한다.
# WHY: 학습 하이퍼파라미터(lr, batch_size...)는 재개할 때 바꿔도 되지만, 형태를 바꾸면
# state_dict 로드가 그냥 깨진다. 두 부류를 튜플로 갈라 두면 실수할 여지가 없다.
MODEL_KEYS = (
"vocab_size",
"block_size",
"n_layer",
"n_head",
"n_kv_head",
"d_model",
"ffn_hidden",
)
class RMSNorm(nn.Module):
"""LayerNorm에서 평균 빼기를 없앤 정규화."""
def __init__(self, dim, eps=1e-5):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
in_dtype = x.dtype
# WHY[RMSNorm]: 벡터를 평균 대신 제곱평균제곱근(RMS)으로만 나눠 크기를 맞추는 정규화.
# 🧩 합창단 음량 맞추기. LayerNorm이 "평균 음정을 0으로 옮기고 음량도 맞춘다"면
# RMSNorm은 음정은 그대로 두고 음량만 맞춘다. 뺄셈 한 번을 통째로 아낀다.
# ⚠️ 제곱과 평균은 fp16에서 쉽게 오버플로/정밀도 손실을 낸다. autocast는 이 연산을
# fp16에 남겨두므로 반드시 float()로 승격해 계산하고 원래 dtype으로 되돌린다.
x = x.float()
x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
return x.to(in_dtype) * self.weight
def precompute_rope(head_dim, block_size, base=10000.0, device=None):
"""RoPE의 cos/sin 테이블을 미리 계산한다. 반환 shape은 각각 (block_size, head_dim//2)."""
# WHY[RoPE]: 위치 정보를 벡터에 더하는 대신, 벡터를 위치에 비례한 각도만큼 회전시키는 방식.
# 🧩 시곗바늘. 차원 쌍마다 회전 속도가 다른 바늘을 달아 두면, 두 토큰의 내적은
# 두 바늘의 각도 차 — 즉 상대 거리 — 에만 반응한다. 절대 위치가 사라진다.
# ⚠️ 학습된 위치 임베딩과 달리 테이블은 상수라 학습되지 않는다. 대신 block_size보다
# 긴 위치는 테이블에 아예 없으므로 생성 시 반드시 길이를 직접 막아야 한다.
inv_freq = 1.0 / (base ** (torch.arange(0, head_dim, 2, device=device).float() / head_dim))
t = torch.arange(block_size, device=device).float()
freqs = torch.outer(t, inv_freq)
return torch.cos(freqs), torch.sin(freqs)
def apply_rope(x, cos, sin, pos=0):
"""x: (B, n_head, T, head_dim). pos는 이 청크가 시퀀스에서 시작하는 절대 위치."""
seq_len = x.shape[-2]
# GOTCHA: pos 슬라이스가 KV캐시 디코드의 핵심이다. 캐시를 쓰면 매 스텝 T=1짜리 청크가
# 들어오는데, 그 토큰의 진짜 위치는 0이 아니라 이미 캐시에 쌓인 길이다.
c = cos[pos : pos + seq_len].to(x.dtype)
s = sin[pos : pos + seq_len].to(x.dtype)
x1, x2 = x[..., 0::2], x[..., 1::2]
out1 = x1 * c - x2 * s
out2 = x1 * s + x2 * c
return torch.stack((out1, out2), dim=-1).flatten(-2)
class CausalSelfAttention(nn.Module):
def __init__(self, cfg):
super().__init__()
d_model = cfg["d_model"]
# WHY[MHA]: 멀티헤드 어텐션. d_model 차원을 n_head개로 쪼개 각 헤드가 독립적으로
# "어느 토큰을 볼지"를 따로 결정하고, 결과를 다시 이어붙이는 구조.
# 🧩 같은 문장을 읽는 12명의 독자. 한 명은 주어-서술어 호응을, 한 명은 조사를,
# 한 명은 앞 문단의 화제를 좇는다. 12개의 시선을 합쳐 한 문장을 이해한다.
# ⚠️ 헤드를 늘려도 총 파라미터는 그대로다(차원을 쪼갤 뿐). 헤드가 많아질수록
# head_dim이 작아져 헤드 하나의 표현력은 오히려 떨어진다. 64가 관례적 하한.
self.n_head = cfg["n_head"]
self.n_kv_head = cfg["n_kv_head"]
assert d_model % self.n_head == 0, "d_model은 n_head로 나누어떨어져야 한다"
assert self.n_head % self.n_kv_head == 0, "n_head는 n_kv_head의 배수여야 한다"
self.head_dim = d_model // self.n_head
# WHY[GQA]: Grouped-Query Attention. Query 헤드는 다 두되 Key/Value 헤드만 줄여
# 여러 Q 헤드가 K/V 한 벌을 공유하게 만드는 절충안.
# 🧩 회의실 12명(Q)이 자료집 4권(K/V)을 3명씩 나눠 보는 것. 자료집 복사비(KV캐시
# 메모리)는 1/3로 줄지만 회의 참석자 수는 그대로다.
# ⚠️ n_kv_head=1이면 MQA, n_kv_head=n_head면 그냥 MHA다. 품질 손실은 거의 없지만
# "거의"이므로, 캐시 메모리가 병목이 아니라면 굳이 줄일 이유도 없다.
self.n_rep = self.n_head // self.n_kv_head
# WHY[Q/K/V]: Query는 "내가 찾는 것", Key는 "내가 가진 이름표", Value는 "내가 줄 내용".
# 🧩 도서관 검색. 검색어(Q)를 책등의 제목(K)들과 대조해 유사도를 매기고, 그 비율대로
# 책의 본문(V)을 섞어 가져온다. 정확히 한 권만 고르는 게 아니라 가중 평균이다.
# ⚠️ 세 벌 모두 같은 입력 x에서 나온다(그래서 "self"-attention). Q와 K의 head_dim이
# 같아야 내적이 되고, V의 head_dim은 사실 달라도 되지만 관례상 맞춘다.
self.q_proj = nn.Linear(d_model, self.n_head * self.head_dim, bias=False)
self.k_proj = nn.Linear(d_model, self.n_kv_head * self.head_dim, bias=False)
self.v_proj = nn.Linear(d_model, self.n_kv_head * self.head_dim, bias=False)
self.o_proj = nn.Linear(self.n_head * self.head_dim, d_model, bias=False)
def forward(self, x, cos, sin, cache=None, pos=0):
"""cache가 주어지면 (k_buf, v_buf)에 이번 청크의 k/v를 써 넣고 누적분 전체로 어텐션."""
B, T, _ = x.shape
q = self.q_proj(x).view(B, T, self.n_head, self.head_dim).transpose(1, 2)
k = self.k_proj(x).view(B, T, self.n_kv_head, self.head_dim).transpose(1, 2)
v = self.v_proj(x).view(B, T, self.n_kv_head, self.head_dim).transpose(1, 2)
# Q와 K에만 회전을 건다. V는 "내용"이라 위치를 섞을 이유가 없다.
q = apply_rope(q, cos, sin, pos)
k = apply_rope(k, cos, sin, pos)
if cache is not None:
k_buf, v_buf = cache
k_buf[:, :, pos : pos + T] = k
v_buf[:, :, pos : pos + T] = v
# GOTCHA: 캐시는 파라미터 dtype(보통 fp32)으로 할당돼 있는데 autocast 아래의 q는
# fp16이다. dtype이 어긋나면 SDPA가 그대로 터지므로 읽을 때 q에 맞춰 준다.
k = k_buf[:, :, : pos + T].to(q.dtype)
v = v_buf[:, :, : pos + T].to(q.dtype)
# GQA 확장: K/V 헤드를 n_rep번 복제해 Q 헤드 수에 맞춘다. repeat_interleave여야
# 그룹이 [0,0,0,1,1,1,...] 순으로 붙어 Q 헤드와 짝이 맞는다(repeat는 순서가 틀린다).
if self.n_rep > 1:
k = k.repeat_interleave(self.n_rep, dim=1)
v = v.repeat_interleave(self.n_rep, dim=1)
# WHY[인과마스크]: 각 위치가 자기 자신과 그 이전 토큰만 보게 어텐션 점수의 오른쪽
# 위 삼각형을 -inf로 막는 장치. 미래를 못 보게 하는 눈가리개.
# 🧩 시험지를 한 문항씩 가려 가며 푸는 것. 뒷장을 미리 봤다면 정답을 맞혀도
# 실력이 아니다. 마스크가 없으면 모델은 "다음 토큰"을 그냥 베껴 loss가 0으로 간다.
# ⚠️ SDPA의 is_causal=True는 마스크를 좌상단 기준으로 정렬한다. q_len=1, kv_len=100인
# 디코드에서 켜면 그 한 토큰이 position 0만 보게 되어 생성이 즉시 붕괴한다.
# GOTCHA: 그래서 정방행렬일 때만 켠다. 디코드(q_len<kv_len)는 이미 캐시에 과거만
# 들어 있으므로 마스크 자체가 불필요하다.
# GOTCHA: T4(SM75)는 FlashAttention-2를 지원하지 않는다. SDPA가 알아서
# mem-efficient 커널로 폴백하므로 코드는 그대로 두면 된다.
y = F.scaled_dot_product_attention(q, k, v, is_causal=(q.size(2) == k.size(2)))
y = y.transpose(1, 2).contiguous().view(B, T, -1)
return self.o_proj(y)
class SwiGLU(nn.Module):
def __init__(self, cfg):
super().__init__()
d_model, hidden = cfg["d_model"], cfg["ffn_hidden"]
self.w_gate = nn.Linear(d_model, hidden, bias=False)
self.w_up = nn.Linear(d_model, hidden, bias=False)
self.w_down = nn.Linear(hidden, d_model, bias=False)
def forward(self, x):
# WHY[SwiGLU]: 활성함수를 통과한 게이트(silu(w_gate·x))를 또 다른 선형 사영(w_up·x)에
# 원소별로 곱해, 어떤 채널을 얼마나 통과시킬지 입력 스스로 정하게 하는 FFN.
# 🧩 수도꼭지가 달린 파이프. w_up이 흘려보낼 물이라면 w_gate는 각 파이프의 밸브 개도다.
# ReLU가 "0 이하는 잠금"이라는 고정 규칙이라면 이쪽은 학습되는 밸브다.
# ⚠️ 행렬이 2개에서 3개로 늘어나므로 hidden을 4d 그대로 두면 파라미터가 1.5배가 된다.
# 그래서 (8/3)d로 줄인다 — 여기서는 768*8/3 = 2048.
# WHY: 3행렬 × 768 × 2048 = 4,718,592로, GPT-2식 2행렬 4d MLP(2 × 768 × 3072)와
# 파라미터 수가 정확히 같다. 공짜 성능이 아니라 등가 교환이라는 뜻이다.
return self.w_down(F.silu(self.w_gate(x)) * self.w_up(x))
class Block(nn.Module):
def __init__(self, cfg):
super().__init__()
self.attn_norm = RMSNorm(cfg["d_model"])
self.attn = CausalSelfAttention(cfg)
self.ffn_norm = RMSNorm(cfg["d_model"])
self.mlp = SwiGLU(cfg)
def forward(self, x, cos, sin, cache=None, pos=0):
# WHY[잔차연결]: 서브층의 출력을 입력에 덮어쓰지 않고 더하는 연결(x = x + f(x)).
# 🧩 원고에 빨간 펜으로 교정만 얹는 것. 원문을 지우고 새로 쓰는 게 아니라서
# 14층을 쌓아도 1층의 정보가 마지막까지 그대로 흘러간다.
# ⚠️ 역전파에서 기울기가 덧셈 경로를 타고 감쇠 없이 내려간다. 이게 없으면 14층은
# 학습이 안 된다. pre-norm(정규화를 f 안쪽에 두기)이어야 이 경로가 깨끗하다.
x = x + self.attn(self.attn_norm(x), cos, sin, cache, pos)
x = x + self.mlp(self.ffn_norm(x))
return x
class GPT(nn.Module):
def __init__(self, cfg):
super().__init__()
missing = [k for k in MODEL_KEYS if k not in cfg]
assert not missing, f"cfg에 모델 키가 없다: {missing}"
self.cfg = {k: cfg[k] for k in MODEL_KEYS}
d_model = cfg["d_model"]
self.block_size = cfg["block_size"]
self.n_layer = cfg["n_layer"]
self.n_kv_head = cfg["n_kv_head"]
self.head_dim = d_model // cfg["n_head"]
self.wte = nn.Embedding(cfg["vocab_size"], d_model)
self.blocks = nn.ModuleList([Block(cfg) for _ in range(self.n_layer)])
self.norm_f = RMSNorm(d_model)
self.lm_head = nn.Linear(d_model, cfg["vocab_size"], bias=False)
# WHY[tied embedding]: 입력 임베딩 행렬과 출력 lm_head 행렬을 같은 텐서로 묶는 것.
# 🧩 국어사전 한 권을 읽을 때도 쓰고 쓸 때도 쓰는 것. "단어 → 벡터"와
# "벡터 → 단어"는 결국 같은 사전을 양방향으로 보는 일이다.
# ⚠️ 여기서는 32768×768 = 25.2M 파라미터, 즉 전체의 22%를 통째로 아낀다. 다만 두
# 텐서가 물리적으로 같으므로 한쪽 기울기가 다른 쪽에 그대로 더해진다는 걸 잊지 말 것.
self.lm_head.weight = self.wte.weight
# 위치 임베딩은 없다 — 위치 정보는 전부 RoPE가 담당한다.
cos, sin = precompute_rope(self.head_dim, self.block_size)
# persistent=False: 상수 테이블이라 체크포인트에 넣을 이유가 없다.
self.register_buffer("rope_cos", cos, persistent=False)
self.register_buffer("rope_sin", sin, persistent=False)
self.apply(self._init_weights)
# 잔차 경로로 흘러들어가는 사영은 층 수만큼 누적되므로 표준편차를 1/sqrt(2L)로 줄인다.
# WHY: 층마다 두 번(attn, mlp) 더해지니 분산이 2L배로 커진다. GPT-2 논문의 처방.
residual_std = 0.02 / math.sqrt(2 * self.n_layer)
for name, param in self.named_parameters():
if name.endswith("o_proj.weight") or name.endswith("w_down.weight"):
nn.init.normal_(param, mean=0.0, std=residual_std)
@staticmethod
def _init_weights(module):
if isinstance(module, nn.Linear):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
elif isinstance(module, nn.Embedding):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
def num_params(self):
"""총 파라미터 수. 타잉된 lm_head는 wte와 같은 텐서라 자동으로 한 번만 센다."""
return sum(p.numel() for p in self.parameters())
def _backbone(self, idx, caches=None, pos=0):
x = self.wte(idx)
for i, block in enumerate(self.blocks):
cache = None if caches is None else caches[i]
x = block(x, self.rope_cos, self.rope_sin, cache, pos)
return self.norm_f(x)
def forward(self, idx, targets=None):
T = idx.size(1)
assert T <= self.block_size, f"시퀀스 길이 {T}가 block_size {self.block_size}를 넘는다"
x = self._backbone(idx)
logits = self.lm_head(x)
loss = None
if targets is not None:
loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.reshape(-1))
return logits, loss
# ---- KV캐시 추론 경로 -------------------------------------------------
def alloc_kv_cache(self, batch_size, device=None, dtype=None):
"""레이어마다 (k_buf, v_buf)를 block_size 길이로 미리 잡아 둔다."""
device = device or self.wte.weight.device
dtype = dtype or self.wte.weight.dtype
shape = (batch_size, self.n_kv_head, self.block_size, self.head_dim)
return [
(
torch.zeros(shape, device=device, dtype=dtype),
torch.zeros(shape, device=device, dtype=dtype),
)
for _ in range(self.n_layer)
]
def forward_cached(self, idx, caches, pos):
"""캐시를 갱신하며 청크 하나를 흘려보내고, 마지막 위치의 logits (B, vocab)만 돌려준다."""
x = self._backbone(idx, caches=caches, pos=pos)
return self.lm_head(x[:, -1])
@torch.no_grad()
def generate(self, idx, max_new_tokens, temperature=0.8, top_k=None):
# WHY[KV캐시]: 이미 계산한 토큰들의 Key/Value를 저장해 두고, 새 토큰 하나만 추가로
# 계산해 어텐션을 잇는 기법. 매 스텝 전체 시퀀스를 다시 돌리지 않게 해 준다.
# 🧩 회의록을 계속 이어 쓰는 것. 새 발언이 나올 때마다 회의 전체를 재연하는 대신
# 기록해 둔 회의록을 펼쳐 놓고 한 줄만 덧붙인다. O(T²)가 스텝당 O(T)가 된다.
# ⚠️ 속도를 메모리로 산다. 그리고 캐시를 쓰는 순간 "이 토큰의 위치"가 더 이상 0이
# 아니므로 RoPE 오프셋과 인과마스크를 손으로 맞춰 줘야 한다 — 여기가 늘 버그 나는 곳.
# WHY: 정직하게 말하면 이 규모에서 GQA의 캐시 절감은 실익이 없다. fp32·batch 1·T=1024
# 기준 GQA 29MB vs MHA 88MB(fp16이면 14.7MB / 44.0MB)로, 둘 다 그냥 작다.
# GQA를 넣은 이유는 메모리가 아니라 최신 아키텍처를 손으로 짜 보기 위해서다.
self.eval()
idx = idx[:, -self.block_size :]
caches = self.alloc_kv_cache(idx.size(0), device=idx.device)
pos = 0
for _ in range(max_new_tokens):
chunk = idx if pos == 0 else idx[:, -1:]
if pos + chunk.size(1) > self.block_size:
break # 컨텍스트 창을 다 썼다. RoPE 테이블 밖이므로 여기서 멈춘다.
logits = self.forward_cached(chunk, caches, pos)
pos += chunk.size(1)
idx_next = sample_from_logits(logits, temperature, top_k)
idx = torch.cat((idx, idx_next), dim=1)
return idx
def sample_from_logits(logits, temperature=0.8, top_k=None):
"""(B, vocab) logits에서 다음 토큰 (B, 1)을 뽑는다. temperature=0이면 greedy."""
if temperature == 0.0:
return torch.argmax(logits, dim=-1, keepdim=True)
logits = logits.float() / temperature
if top_k is not None:
k = min(top_k, logits.size(-1))
thresh = torch.topk(logits, k, dim=-1).values[:, -1:]
logits = logits.masked_fill(logits < thresh, float("-inf"))
probs = F.softmax(logits, dim=-1)
return torch.multinomial(probs, num_samples=1)