Maggio33 commited on
Commit
75bb59e
·
verified ·
1 Parent(s): ac68d22

Upload train_gpt_ref.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. train_gpt_ref.py +404 -0
train_gpt_ref.py ADDED
@@ -0,0 +1,404 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ # -*- coding: utf-8 -*-
3
+ """
4
+ Referencyjny ZWYKLY transformer (byte-level nanoGPT-style) ~25M — apples-to-apples vs BDH-25M.
5
+ Ta sama data (train.bin/val.bin uint8), ten sam scale (~25M), ten sam byte-level (vocab256).
6
+ Rozni sie TYLKO architektura (standard causal transformer vs BDH fast-weights) -> czysta referencja.
7
+
8
+ CLI mirror train_bdh.py. GPU ROCm/CUDA bf16, cosine+warmup+clip, ckpt/resume, logging.
9
+ Autor: Hart (N-02).
10
+
11
+ Smoke throughput (bez danych PII, syntetyczny bufor): --synthetic --steps 60
12
+ Realny: --data-dir . --run-id gpt25m_run1 --steps 30000
13
+ """
14
+ import argparse
15
+ import json
16
+ import math
17
+ import os
18
+ import time
19
+ import queue
20
+ import threading
21
+ from contextlib import nullcontext
22
+
23
+ import numpy as np
24
+ import torch
25
+ import torch.nn as nn
26
+ import torch.nn.functional as F
27
+
28
+
29
+ def get_args():
30
+ p = argparse.ArgumentParser()
31
+ p.add_argument("--data-dir", default=".")
32
+ p.add_argument("--out-dir", default=None)
33
+ p.add_argument("--run-id", default="gpt25m_run1")
34
+ p.add_argument("--steps", type=int, default=30000)
35
+ p.add_argument("--batch", type=int, default=32)
36
+ p.add_argument("--block", type=int, default=256)
37
+ p.add_argument("--n-layer", type=int, default=8)
38
+ p.add_argument("--n-embd", type=int, default=512)
39
+ p.add_argument("--n-head", type=int, default=8)
40
+ p.add_argument("--lr", type=float, default=6e-4)
41
+ p.add_argument("--min-lr", type=float, default=6e-5)
42
+ p.add_argument("--warmup", type=int, default=200)
43
+ p.add_argument("--wd", type=float, default=0.1)
44
+ p.add_argument("--grad-clip", type=float, default=1.0)
45
+ p.add_argument("--log-every", type=int, default=50)
46
+ p.add_argument("--eval-every", type=int, default=500)
47
+ p.add_argument("--eval-iters", type=int, default=50)
48
+ p.add_argument("--ckpt-every", type=int, default=1000)
49
+ p.add_argument("--seed", type=int, default=1337)
50
+ p.add_argument("--resume", action="store_true")
51
+ p.add_argument("--synthetic", action="store_true", help="smoke throughput na losowym uint8 (bez danych)")
52
+ p.add_argument("--vocab", type=int, default=256)
53
+ p.add_argument("--dtype", default="uint8", help="bin dtype: uint8 (byte) | uint16 (BPE)")
54
+ p.add_argument("--optimizer", choices=["adamw", "muon"], default="adamw",
55
+ help="adamw (default, backward-compat) | muon (Newton-Schulz ortho dla 2D-weights + AdamW dla reszty)")
56
+ p.add_argument("--muon-lr", type=float, default=0.02,
57
+ help="peak LR dla Muon (macierzowe params); AdamW-aux uzywa --lr. Muon skalowany ta sama cosine-schedule co AdamW przez lr_mult=muon_lr/lr")
58
+ return p.parse_args()
59
+
60
+
61
+ # ---- Muon (Keller Jordan) --------------------------------------------------
62
+ # Ref: https://github.com/KellerJordan/Muon (modded-nanogpt). Muon = momentum
63
+ # SGD, ale update ortogonalizowany przez ~5 krokow iteracji Newtona-Schulza
64
+ # (przyblizona ortogonalizacja macierzy gradientu). Stosowany TYLKO do
65
+ # macierzowych ukrytych wag (ndim>=2: qkv/proj/mlp). Embeddingi (tok/pos), head
66
+ # (tied), LayerNorm-gains i biasy ida do zwyklego AdamW.
67
+ def zeropower_via_newtonschulz5(G, steps=5):
68
+ """Ortogonalizacja macierzy G przez quintic Newton-Schulz (bf16). Zwraca
69
+ macierz ~ U V^T z SVD(G)=U S V^T. Wspolczynniki (a,b,c) z impl. Kellera."""
70
+ assert G.ndim == 2
71
+ a, b, c = (3.4445, -4.7750, 2.0315)
72
+ X = G.bfloat16()
73
+ transposed = G.size(0) > G.size(1)
74
+ if transposed:
75
+ X = X.T
76
+ X = X / (X.norm() + 1e-7)
77
+ for _ in range(steps):
78
+ A = X @ X.T
79
+ B = b * A + c * (A @ A)
80
+ X = a * X + B @ X
81
+ if transposed:
82
+ X = X.T
83
+ return X
84
+
85
+
86
+ class Muon(torch.optim.Optimizer):
87
+ """Momentum-SGD z ortogonalizowanym update. weight_decay domyslnie 0 (Muon-params
88
+ czysto; WD trzymamy na AdamW-aux). lr_mult pozwala petli lr-schedule skalowac
89
+ Muon proporcjonalnie do AdamW."""
90
+ def __init__(self, params, lr=0.02, lr_mult=1.0, momentum=0.95, nesterov=True,
91
+ ns_steps=5, weight_decay=0.0):
92
+ defaults = dict(lr=lr, lr_mult=lr_mult, momentum=momentum, nesterov=nesterov,
93
+ ns_steps=ns_steps, weight_decay=weight_decay)
94
+ super().__init__(params, defaults)
95
+
96
+ @torch.no_grad()
97
+ def step(self, closure=None):
98
+ loss = None
99
+ if closure is not None:
100
+ with torch.enable_grad():
101
+ loss = closure()
102
+ for group in self.param_groups:
103
+ lr = group["lr"]; momentum = group["momentum"]; wd = group["weight_decay"]
104
+ for p in group["params"]:
105
+ g = p.grad
106
+ if g is None:
107
+ continue
108
+ if g.ndim > 2:
109
+ g = g.reshape(g.size(0), -1)
110
+ state = self.state[p]
111
+ if "momentum_buffer" not in state:
112
+ state["momentum_buffer"] = torch.zeros_like(g)
113
+ buf = state["momentum_buffer"]
114
+ buf.mul_(momentum).add_(g)
115
+ g = g.add(buf, alpha=momentum) if group["nesterov"] else buf
116
+ u = zeropower_via_newtonschulz5(g, steps=group["ns_steps"])
117
+ if wd != 0:
118
+ p.mul_(1 - lr * wd)
119
+ # scale ~ sqrt(fan_out/fan_in): zrownuje RMS update niezaleznie od ksztaltu
120
+ scale = max(1.0, p.size(0) / p.size(1)) ** 0.5
121
+ p.add_(u.reshape(p.shape).to(p.dtype), alpha=-lr * scale)
122
+ return loss
123
+
124
+
125
+ class MuonWithAuxAdam:
126
+ """Kontener: Muon dla macierzowych ukrytych wag + AdamW dla reszty. Wystawia
127
+ param_groups/step/zero_grad/state_dict tak, by petla treningowa dzialala bez zmian."""
128
+ def __init__(self, muon, adamw):
129
+ self.muon = muon
130
+ self.adamw = adamw
131
+
132
+ @property
133
+ def param_groups(self):
134
+ return self.muon.param_groups + self.adamw.param_groups
135
+
136
+ @property
137
+ def state(self):
138
+ return {**self.muon.state, **self.adamw.state}
139
+
140
+ def step(self, closure=None):
141
+ self.muon.step()
142
+ self.adamw.step()
143
+
144
+ def zero_grad(self, set_to_none=True):
145
+ self.muon.zero_grad(set_to_none=set_to_none)
146
+ self.adamw.zero_grad(set_to_none=set_to_none)
147
+
148
+ def state_dict(self):
149
+ return {"muon": self.muon.state_dict(), "adamw": self.adamw.state_dict()}
150
+
151
+ def load_state_dict(self, sd):
152
+ self.muon.load_state_dict(sd["muon"])
153
+ self.adamw.load_state_dict(sd["adamw"])
154
+
155
+
156
+ def build_optimizer(a, model):
157
+ """--optimizer adamw -> DOKLADNIE poprzedni AdamW (backward-compat).
158
+ --optimizer muon -> Muon(2D-hidden) + AdamW(embeddingi/head/norm/bias)."""
159
+ if a.optimizer == "adamw":
160
+ return torch.optim.AdamW(model.parameters(), lr=a.lr, weight_decay=a.wd, betas=(0.9, 0.95))
161
+ muon_params, adamw_params, seen = [], [], set()
162
+ for name, p in model.named_parameters():
163
+ if not p.requires_grad or id(p) in seen:
164
+ continue
165
+ seen.add(id(p))
166
+ is_embed_or_head = name.startswith(("tok.", "pos.", "head."))
167
+ if p.ndim >= 2 and not is_embed_or_head:
168
+ muon_params.append(p)
169
+ else:
170
+ adamw_params.append(p)
171
+ lr_mult = a.muon_lr / a.lr if a.lr > 0 else 1.0
172
+ muon = Muon(muon_params, lr=a.muon_lr, lr_mult=lr_mult, weight_decay=0.0)
173
+ adamw = torch.optim.AdamW(adamw_params, lr=a.lr, weight_decay=a.wd, betas=(0.9, 0.95))
174
+ return MuonWithAuxAdam(muon, adamw)
175
+
176
+
177
+ class Block(nn.Module):
178
+ def __init__(self, d, nh, block):
179
+ super().__init__()
180
+ self.ln1 = nn.LayerNorm(d)
181
+ self.ln2 = nn.LayerNorm(d)
182
+ self.qkv = nn.Linear(d, 3 * d)
183
+ self.proj = nn.Linear(d, d)
184
+ self.mlp = nn.Sequential(nn.Linear(d, 4 * d), nn.GELU(), nn.Linear(4 * d, d))
185
+ self.nh = nh
186
+ self.d = d
187
+
188
+ def forward(self, x):
189
+ B, T, D = x.size()
190
+ h = self.ln1(x)
191
+ q, k, v = self.qkv(h).split(self.d, dim=2)
192
+ q = q.view(B, T, self.nh, D // self.nh).transpose(1, 2)
193
+ k = k.view(B, T, self.nh, D // self.nh).transpose(1, 2)
194
+ v = v.view(B, T, self.nh, D // self.nh).transpose(1, 2)
195
+ y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
196
+ y = y.transpose(1, 2).contiguous().view(B, T, D)
197
+ x = x + self.proj(y)
198
+ x = x + self.mlp(self.ln2(x))
199
+ return x
200
+
201
+
202
+ class GPT(nn.Module):
203
+ def __init__(self, vocab, n_layer, n_embd, n_head, block):
204
+ super().__init__()
205
+ self.tok = nn.Embedding(vocab, n_embd)
206
+ self.pos = nn.Embedding(block, n_embd)
207
+ self.blocks = nn.ModuleList([Block(n_embd, n_head, block) for _ in range(n_layer)])
208
+ self.lnf = nn.LayerNorm(n_embd)
209
+ self.head = nn.Linear(n_embd, vocab, bias=False)
210
+ self.head.weight = self.tok.weight # tie
211
+ self.block = block
212
+ self.apply(self._init)
213
+
214
+ def _init(self, m):
215
+ if isinstance(m, nn.Linear):
216
+ nn.init.normal_(m.weight, 0.0, 0.02)
217
+ if m.bias is not None:
218
+ nn.init.zeros_(m.bias)
219
+ elif isinstance(m, nn.Embedding):
220
+ nn.init.normal_(m.weight, 0.0, 0.02)
221
+
222
+ def forward(self, idx, targets=None):
223
+ B, T = idx.size()
224
+ pos = torch.arange(T, device=idx.device)
225
+ x = self.tok(idx) + self.pos(pos)[None]
226
+ for b in self.blocks:
227
+ x = b(x)
228
+ logits = self.head(self.lnf(x))
229
+ loss = None
230
+ if targets is not None:
231
+ loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1))
232
+ return logits, loss
233
+
234
+
235
+ def main():
236
+ a = get_args()
237
+ out_dir = a.out_dir or os.path.join(a.data_dir, "runs", a.run_id)
238
+ os.makedirs(out_dir, exist_ok=True)
239
+ log_path = os.path.join(out_dir, "train.log")
240
+ metrics_path = os.path.join(out_dir, "metrics.jsonl")
241
+ ckpt_path = os.path.join(out_dir, "ckpt.pt")
242
+
243
+ def log(msg):
244
+ line = f"[{time.strftime('%H:%M:%S')}] {msg}"
245
+ print(line, flush=True)
246
+ with open(log_path, "a", encoding="utf-8") as f:
247
+ f.write(line + "\n")
248
+
249
+ torch.manual_seed(a.seed)
250
+ torch.backends.cuda.matmul.allow_tf32 = True
251
+ torch.backends.cudnn.allow_tf32 = True
252
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
253
+ use_bf16 = device.type == "cuda" and torch.cuda.is_bf16_supported()
254
+ ptdtype = torch.bfloat16 if use_bf16 else torch.float32
255
+ ctx = torch.amp.autocast(device_type=device.type, dtype=ptdtype) if device.type == "cuda" else nullcontext()
256
+ log(f"device={device} bf16={use_bf16} dev={torch.cuda.get_device_name(0) if device.type=='cuda' else 'cpu'}")
257
+
258
+ if a.synthetic:
259
+ rng = np.random.default_rng(a.seed)
260
+ train_data = rng.integers(0, 256, size=8_000_000, dtype=np.uint8)
261
+ val_data = train_data[:200_000]
262
+ log("SYNTHETIC uint8 (smoke throughput, zero danych PII)")
263
+ else:
264
+ train_data = np.memmap(os.path.join(a.data_dir, "train.bin"), dtype=np.dtype(a.dtype), mode="r")
265
+ val_data = np.memmap(os.path.join(a.data_dir, "val.bin"), dtype=np.dtype(a.dtype), mode="r")
266
+ log(f"dane: train={len(train_data):,}B block={a.block} batch={a.batch} tok/step={a.block*a.batch:,}")
267
+
268
+ def _make_batch_cpu(split, generator=None):
269
+ """Wektoryzowane budowanie batcha: JEDEN numpy fancy-index zamiast
270
+ python-loop per-item. sliding_window_view daje strided-view (N-block, block+1)
271
+ BEZ kopiowania; windows[ix] materializuje tylko wybrane wiersze naraz.
272
+ Zwraca (x,y) long CPU (pinned jesli cuda). Rozklad batchy IDENTYCZNY jak
273
+ stary torch.stack-loop: x=data[i:i+block], y=data[i+1:i+1+block]."""
274
+ data = train_data if split == "train" else val_data
275
+ ix = torch.randint(len(data) - a.block - 1, (a.batch,), generator=generator)
276
+ # (N-block, block+1) view; jeden fancy-index kopiuje wybrane okna
277
+ windows = np.lib.stride_tricks.sliding_window_view(data, a.block + 1)
278
+ sel = windows[ix.numpy()] # (batch, block+1) materialized
279
+ x = torch.from_numpy(sel[:, :-1].astype(np.int64)) # astype -> contiguous copy
280
+ y = torch.from_numpy(sel[:, 1:].astype(np.int64))
281
+ if device.type == "cuda":
282
+ x = x.pin_memory(); y = y.pin_memory()
283
+ return x, y
284
+
285
+ def _to_device(x, y):
286
+ if device.type == "cuda":
287
+ return x.to(device, non_blocking=True), y.to(device, non_blocking=True)
288
+ return x.to(device), y.to(device)
289
+
290
+ def get_batch(split, generator=None):
291
+ return _to_device(*_make_batch_cpu(split, generator))
292
+
293
+ class Prefetcher:
294
+ """Async double-buffer: 1 background-thread buduje NASTEPNY batch na CPU
295
+ (pinned) podczas gdy GPU liczy biezacy. queue depth=2. Konsument robi
296
+ .next() -> H2D-copy (non_blocking) w watku glownym. Watek uzywa wlasnego
297
+ torch.Generator (seeded), wiec ciag train-batchy jest deterministyczny i
298
+ NIEZALEZNY od timingu watku oraz od RNG val-loopa (dystrybucja bez zmian)."""
299
+ def __init__(self, split, generator, depth=2):
300
+ self.split = split
301
+ self.gen = generator
302
+ self.q = queue.Queue(maxsize=depth)
303
+ self._stop = threading.Event()
304
+ self.t = threading.Thread(target=self._worker, daemon=True)
305
+ self.t.start()
306
+
307
+ def _worker(self):
308
+ while not self._stop.is_set():
309
+ try:
310
+ item = _make_batch_cpu(self.split, self.gen)
311
+ except Exception as e: # przekaz blad do konsumenta
312
+ self.q.put(e)
313
+ return
314
+ while not self._stop.is_set():
315
+ try:
316
+ self.q.put(item, timeout=0.5)
317
+ break
318
+ except queue.Full:
319
+ continue
320
+
321
+ def next(self):
322
+ item = self.q.get()
323
+ if isinstance(item, Exception):
324
+ raise item
325
+ return _to_device(*item)
326
+
327
+ def close(self):
328
+ self._stop.set()
329
+ # opróżnij kolejke zeby watek nie zawisl na put()
330
+ try:
331
+ self.q.get_nowait()
332
+ except queue.Empty:
333
+ pass
334
+
335
+ model = GPT(a.vocab, a.n_layer, a.n_embd, a.n_head, a.block).to(device)
336
+ nparam = sum(p.numel() for p in model.parameters())
337
+ log(f"model GPT-ref: {nparam/1e6:.1f}M param (L{a.n_layer} d{a.n_embd} h{a.n_head})")
338
+ opt = build_optimizer(a, model)
339
+ log(f"optimizer={a.optimizer}" + (f" muon_lr={a.muon_lr} (mult={a.muon_lr/a.lr:.1f}x)" if a.optimizer == "muon" else ""))
340
+
341
+ start_step = 0
342
+ if a.resume and os.path.exists(ckpt_path):
343
+ ck = torch.load(ckpt_path, map_location=device)
344
+ model.load_state_dict(ck["model"]); opt.load_state_dict(ck["opt"]); start_step = ck["step"]
345
+ log(f"RESUME @ {start_step}")
346
+
347
+ def lr_at(s):
348
+ if s < a.warmup:
349
+ return a.lr * (s + 1) / a.warmup
350
+ if s >= a.steps:
351
+ return a.min_lr
352
+ r = (s - a.warmup) / max(1, a.steps - a.warmup)
353
+ return a.min_lr + 0.5 * (a.lr - a.min_lr) * (1 + math.cos(math.pi * r))
354
+
355
+ @torch.no_grad()
356
+ def eval_val():
357
+ model.eval()
358
+ ls = []
359
+ for _ in range(a.eval_iters):
360
+ xb, yb = get_batch("val")
361
+ with ctx:
362
+ _, loss = model(xb, yb)
363
+ ls.append(loss.item())
364
+ model.train()
365
+ return sum(ls) / len(ls)
366
+
367
+ model.train()
368
+ log(f"START gpt-ref: steps={a.steps} (od {start_step}) lr={a.lr}->{a.min_lr}")
369
+ # dedykowany seeded generator dla train-prefetchera (determinizm niezalezny
370
+ # od RNG val-loopa i timingu watku; ta sama dystrybucja co global-RNG)
371
+ train_gen = torch.Generator()
372
+ train_gen.manual_seed(a.seed)
373
+ prefetcher = Prefetcher("train", train_gen)
374
+ t0 = time.time(); running = 0.0
375
+ for step in range(start_step, a.steps):
376
+ lr = lr_at(step)
377
+ for g in opt.param_groups:
378
+ g["lr"] = lr * g.get("lr_mult", 1.0)
379
+ xb, yb = prefetcher.next()
380
+ with ctx:
381
+ _, loss = model(xb, yb)
382
+ loss.backward()
383
+ gn = torch.nn.utils.clip_grad_norm_(model.parameters(), a.grad_clip) if a.grad_clip > 0 else 0.0
384
+ opt.step(); opt.zero_grad(set_to_none=True)
385
+ running += loss.item()
386
+ if (step + 1) % a.log_every == 0:
387
+ dt = time.time() - t0
388
+ tok_s = a.log_every * a.block * a.batch / dt
389
+ mem = torch.cuda.max_memory_allocated()/1e9 if device.type == "cuda" else 0.0
390
+ log(f"step {step+1}/{a.steps} loss {running/a.log_every:.4f} lr {lr:.2e} gnorm {float(gn):.2f} {tok_s:,.0f} tok/s peakVRAM {mem:.1f}GB")
391
+ with open(metrics_path, "a", encoding="utf-8") as f:
392
+ f.write(json.dumps({"step": step+1, "loss": running/a.log_every, "lr": lr, "tok_s": tok_s}) + "\n")
393
+ running = 0.0; t0 = time.time()
394
+ if (step + 1) % a.eval_every == 0:
395
+ log(f" >> VAL loss {eval_val():.4f} @ {step+1}")
396
+ if (step + 1) % a.ckpt_every == 0 and not a.synthetic:
397
+ torch.save({"model": model.state_dict(), "opt": opt.state_dict(), "step": step+1}, ckpt_path)
398
+ log(f"ckpt @ {step+1}")
399
+ prefetcher.close()
400
+ log("DONE")
401
+
402
+
403
+ if __name__ == "__main__":
404
+ main()