Maggio33 commited on
Commit
b3d0010
·
verified ·
1 Parent(s): 02a1315

Upload train_gpt_ref.py with huggingface_hub

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