Download train_gpt_ref.py from SlayerLab/gollem-v5-ckpts: direct link, hf CLI and curl.
- Browser
- Download file 24 kB
-
https://huggingface.co/SlayerLab/gollem-v5-ckpts/resolve/437a1d431ea63562ee5c4bbf065b22b14616aaaf/train_gpt_ref.py
- Command line
-
hf download hf://SlayerLab/gollem-v5-ckpts@437a1d431ea63562ee5c4bbf065b22b14616aaaf/train_gpt_ref.py
-
curl -L -o train_gpt_ref.py https://huggingface.co/SlayerLab/gollem-v5-ckpts/resolve/437a1d431ea63562ee5c4bbf065b22b14616aaaf/train_gpt_ref.py
24 kB
| #!/usr/bin/env python | |
| # NOTE (GoLLeM-v5 leaderboard repo): this is a GENERAL nanoGPT-style causal transformer; | |
| # vocab/dtype are CLI-parameterized. The crown/Path-B leaderboard checkpoints were trained in | |
| # BPE-12k mode: --vocab 12288 --dtype uint16 (NOT the byte-level default below). The header | |
| # doc-comment reflects the file origin as a standard-GPT control vs experimental BDH; the | |
| # leaderboard models are the STANDARD transformer in BPE mode and do NOT use BDH. | |
| # -*- coding: utf-8 -*- | |
| """ | |
| Referencyjny ZWYKLY transformer (byte-level nanoGPT-style) ~25M — apples-to-apples vs BDH-25M. | |
| Ta sama data (train.bin/val.bin uint8), ten sam scale (~25M), ten sam byte-level (vocab256). | |
| Rozni sie TYLKO architektura (standard causal transformer vs BDH fast-weights) -> czysta referencja. | |
| CLI mirror train_bdh.py. GPU ROCm/CUDA bf16, cosine+warmup+clip, ckpt/resume, logging. | |
| Autor: Hart (N-02). | |
| Smoke throughput (bez danych PII, syntetyczny bufor): --synthetic --steps 60 | |
| Realny: --data-dir . --run-id gpt25m_run1 --steps 30000 | |
| """ | |
| import argparse | |
| import json | |
| import math | |
| import os | |
| import time | |
| import queue | |
| import threading | |
| from contextlib import nullcontext | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| def get_args(): | |
| p = argparse.ArgumentParser() | |
| p.add_argument("--data-dir", default=".") | |
| p.add_argument("--out-dir", default=None) | |
| p.add_argument("--run-id", default="gpt25m_run1") | |
| p.add_argument("--steps", type=int, default=30000) | |
| p.add_argument("--batch", type=int, default=32) | |
| p.add_argument("--block", type=int, default=256) | |
| p.add_argument("--n-layer", type=int, default=8) | |
| p.add_argument("--n-embd", type=int, default=512) | |
| p.add_argument("--n-head", type=int, default=8) | |
| p.add_argument("--lr", type=float, default=6e-4) | |
| p.add_argument("--min-lr", type=float, default=6e-5) | |
| p.add_argument("--warmup", type=int, default=200) | |
| p.add_argument("--wd", type=float, default=0.1) | |
| p.add_argument("--grad-clip", type=float, default=1.0) | |
| p.add_argument("--log-every", type=int, default=50) | |
| p.add_argument("--eval-every", type=int, default=500) | |
| p.add_argument("--eval-iters", type=int, default=50) | |
| p.add_argument("--ckpt-every", type=int, default=1000) | |
| p.add_argument("--seed", type=int, default=1337) | |
| p.add_argument("--resume", action="store_true") | |
| p.add_argument("--synthetic", action="store_true", help="smoke throughput na losowym uint8 (bez danych)") | |
| p.add_argument("--vocab", type=int, default=256) | |
| p.add_argument("--dtype", default="uint8", help="bin dtype: uint8 (byte) | uint16 (BPE)") | |
| p.add_argument("--optimizer", choices=["adamw", "muon"], default="adamw", | |
| help="adamw (default, backward-compat) | muon (Newton-Schulz ortho dla 2D-weights + AdamW dla reszty)") | |
| p.add_argument("--muon-lr", type=float, default=0.02, | |
| 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") | |
| p.add_argument("--norm", choices=["layernorm", "rmsnorm"], default="layernorm", | |
| help="layernorm (default, backward-compat) | rmsnorm (Qwen3-style, fp32-compute)") | |
| p.add_argument("--norm-eps", type=float, default=1e-6) | |
| p.add_argument("--pos", choices=["learned", "rope"], default="learned", | |
| help="learned (default) | rope (parameter-free RoPE na q,k; usuwa learned pos-embedding)") | |
| p.add_argument("--rope-theta", type=float, default=100000.0, | |
| help="RoPE theta; top-3-board (JugnuLM/GPT-X2) uzywaja 100000 (nie-default-10K)") | |
| p.add_argument("--ffn", choices=["gelu", "swiglu"], default="gelu", | |
| help="gelu (default 4x MLP) | swiglu (Qwen3 gated-MLP, hidden=--ffn-mult*d)") | |
| p.add_argument("--ffn-mult", type=float, default=2.667, | |
| help="mnoznik hidden dla swiglu (~param-parity z 4x-gelu przy 8/3)") | |
| p.add_argument("--value-residual", action="store_true", | |
| help="ResFormer value-residuals: v_l += lambda_l*v0 (lambda init 0), ARC-targeted") | |
| p.add_argument("--qk-norm", action="store_true", | |
| help="Qwen3 QK-Norm: RMSNorm per-head na Q,K przed-attention (stabilnosc z Muon/high-LR)") | |
| p.add_argument("--events-jsonl", default=None, | |
| help="jesli podane: emituj events.jsonl (fabryka-track sidecar-format: update/evaluation/checkpoint/end)") | |
| p.add_argument("--compile", action="store_true", | |
| help="torch.compile model (2-3x throughput; state_dict zapisywany bez _orig_mod prefix via raw_model)") | |
| return p.parse_args() | |
| # ---- Muon (Keller Jordan) -------------------------------------------------- | |
| # Ref: https://github.com/KellerJordan/Muon (modded-nanogpt). Muon = momentum | |
| # SGD, ale update ortogonalizowany przez ~5 krokow iteracji Newtona-Schulza | |
| # (przyblizona ortogonalizacja macierzy gradientu). Stosowany TYLKO do | |
| # macierzowych ukrytych wag (ndim>=2: qkv/proj/mlp). Embeddingi (tok/pos), head | |
| # (tied), LayerNorm-gains i biasy ida do zwyklego AdamW. | |
| def zeropower_via_newtonschulz5(G, steps=5): | |
| """Ortogonalizacja macierzy G przez quintic Newton-Schulz (bf16). Zwraca | |
| macierz ~ U V^T z SVD(G)=U S V^T. Wspolczynniki (a,b,c) z impl. Kellera.""" | |
| assert G.ndim == 2 | |
| a, b, c = (3.4445, -4.7750, 2.0315) | |
| X = G.bfloat16() | |
| transposed = G.size(0) > G.size(1) | |
| if transposed: | |
| X = X.T | |
| X = X / (X.norm() + 1e-7) | |
| for _ in range(steps): | |
| A = X @ X.T | |
| B = b * A + c * (A @ A) | |
| X = a * X + B @ X | |
| if transposed: | |
| X = X.T | |
| return X | |
| class Muon(torch.optim.Optimizer): | |
| """Momentum-SGD z ortogonalizowanym update. weight_decay domyslnie 0 (Muon-params | |
| czysto; WD trzymamy na AdamW-aux). lr_mult pozwala petli lr-schedule skalowac | |
| Muon proporcjonalnie do AdamW.""" | |
| def __init__(self, params, lr=0.02, lr_mult=1.0, momentum=0.95, nesterov=True, | |
| ns_steps=5, weight_decay=0.0): | |
| defaults = dict(lr=lr, lr_mult=lr_mult, momentum=momentum, nesterov=nesterov, | |
| ns_steps=ns_steps, weight_decay=weight_decay) | |
| super().__init__(params, defaults) | |
| def step(self, closure=None): | |
| loss = None | |
| if closure is not None: | |
| with torch.enable_grad(): | |
| loss = closure() | |
| for group in self.param_groups: | |
| lr = group["lr"]; momentum = group["momentum"]; wd = group["weight_decay"] | |
| for p in group["params"]: | |
| g = p.grad | |
| if g is None: | |
| continue | |
| if g.ndim > 2: | |
| g = g.reshape(g.size(0), -1) | |
| state = self.state[p] | |
| if "momentum_buffer" not in state: | |
| state["momentum_buffer"] = torch.zeros_like(g) | |
| buf = state["momentum_buffer"] | |
| buf.mul_(momentum).add_(g) | |
| g = g.add(buf, alpha=momentum) if group["nesterov"] else buf | |
| u = zeropower_via_newtonschulz5(g, steps=group["ns_steps"]) | |
| if wd != 0: | |
| p.mul_(1 - lr * wd) | |
| # scale ~ sqrt(fan_out/fan_in): zrownuje RMS update niezaleznie od ksztaltu | |
| scale = max(1.0, p.size(0) / p.size(1)) ** 0.5 | |
| p.add_(u.reshape(p.shape).to(p.dtype), alpha=-lr * scale) | |
| return loss | |
| class MuonWithAuxAdam: | |
| """Kontener: Muon dla macierzowych ukrytych wag + AdamW dla reszty. Wystawia | |
| param_groups/step/zero_grad/state_dict tak, by petla treningowa dzialala bez zmian.""" | |
| def __init__(self, muon, adamw): | |
| self.muon = muon | |
| self.adamw = adamw | |
| def param_groups(self): | |
| return self.muon.param_groups + self.adamw.param_groups | |
| def state(self): | |
| return {**self.muon.state, **self.adamw.state} | |
| def step(self, closure=None): | |
| self.muon.step() | |
| self.adamw.step() | |
| def zero_grad(self, set_to_none=True): | |
| self.muon.zero_grad(set_to_none=set_to_none) | |
| self.adamw.zero_grad(set_to_none=set_to_none) | |
| def state_dict(self): | |
| return {"muon": self.muon.state_dict(), "adamw": self.adamw.state_dict()} | |
| def load_state_dict(self, sd): | |
| self.muon.load_state_dict(sd["muon"]) | |
| self.adamw.load_state_dict(sd["adamw"]) | |
| def build_optimizer(a, model): | |
| """--optimizer adamw -> DOKLADNIE poprzedni AdamW (backward-compat). | |
| --optimizer muon -> Muon(2D-hidden) + AdamW(embeddingi/head/norm/bias).""" | |
| if a.optimizer == "adamw": | |
| return torch.optim.AdamW(model.parameters(), lr=a.lr, weight_decay=a.wd, betas=(0.9, 0.95)) | |
| muon_params, adamw_params, seen = [], [], set() | |
| for name, p in model.named_parameters(): | |
| if not p.requires_grad or id(p) in seen: | |
| continue | |
| seen.add(id(p)) | |
| is_embed_or_head = name.startswith(("tok.", "pos.", "head.")) | |
| if p.ndim >= 2 and not is_embed_or_head: | |
| muon_params.append(p) | |
| else: | |
| adamw_params.append(p) | |
| lr_mult = a.muon_lr / a.lr if a.lr > 0 else 1.0 | |
| muon = Muon(muon_params, lr=a.muon_lr, lr_mult=lr_mult, weight_decay=0.0) | |
| adamw = torch.optim.AdamW(adamw_params, lr=a.lr, weight_decay=a.wd, betas=(0.9, 0.95)) | |
| return MuonWithAuxAdam(muon, adamw) | |
| class RMSNorm(nn.Module): | |
| """Qwen3-style RMSNorm (fp32-compute dla stabilnosci). 1D weight -> AdamW w split-Muon.""" | |
| def __init__(self, d, eps=1e-6): | |
| super().__init__() | |
| self.weight = nn.Parameter(torch.ones(d)) | |
| self.eps = eps | |
| def forward(self, x): | |
| return x * torch.rsqrt(x.float().pow(2).mean(-1, keepdim=True) + self.eps).to(x.dtype) * self.weight | |
| def make_norm(d, cfg): | |
| return RMSNorm(d, cfg.norm_eps) if cfg.norm == "rmsnorm" else nn.LayerNorm(d) | |
| def apply_rope(x, base=100000.0): | |
| """Parameter-free RoPE na [B,H,T,D] (interleaved-conv, port z qwen_model.py). Train==eval | |
| MUSZA uzywac tej samej konwencji (self-contained eval -> spojne).""" | |
| _, _, T, dim = x.shape | |
| pos = torch.arange(T, device=x.device, dtype=torch.float32) | |
| freq = 1.0 / (base ** (torch.arange(0, dim, 2, device=x.device, dtype=torch.float32) / dim)) | |
| ang = torch.outer(pos, freq) | |
| cos, sin = ang.cos().to(x.dtype)[None, None], ang.sin().to(x.dtype)[None, None] | |
| even, odd = x[..., ::2], x[..., 1::2] | |
| return torch.stack((even * cos - odd * sin, even * sin + odd * cos), dim=-1).flatten(-2) | |
| class SwiGLU(nn.Module): | |
| """Qwen3 gated-MLP: down(silu(gate(x))*up(x)). 3x 2D bez-bias -> wszystkie do Muon.""" | |
| def __init__(self, d, hidden): | |
| super().__init__() | |
| self.gate = nn.Linear(d, hidden, bias=False) | |
| self.up = nn.Linear(d, hidden, bias=False) | |
| self.down = nn.Linear(hidden, d, bias=False) | |
| def forward(self, x): | |
| return self.down(F.silu(self.gate(x)) * self.up(x)) | |
| class Block(nn.Module): | |
| def __init__(self, d, nh, block, cfg, is_first=False): | |
| super().__init__() | |
| self.ln1 = make_norm(d, cfg) | |
| self.ln2 = make_norm(d, cfg) | |
| self.qkv = nn.Linear(d, 3 * d) | |
| self.proj = nn.Linear(d, d) | |
| if cfg.ffn == "swiglu": | |
| self.mlp = SwiGLU(d, int(round(cfg.ffn_mult * d))) | |
| else: | |
| self.mlp = nn.Sequential(nn.Linear(d, 4 * d), nn.GELU(), nn.Linear(4 * d, d)) | |
| self.nh = nh | |
| self.d = d | |
| self.cfg = cfg | |
| self.is_first = is_first | |
| if cfg.value_residual and not is_first: | |
| self.vr_lambda = nn.Parameter(torch.zeros(1)) | |
| if cfg.qk_norm: | |
| hd = d // nh | |
| self.q_norm = RMSNorm(hd, cfg.norm_eps) | |
| self.k_norm = RMSNorm(hd, cfg.norm_eps) | |
| def forward(self, x, v0=None): | |
| B, T, D = x.size() | |
| h = self.ln1(x) | |
| q, k, v = self.qkv(h).split(self.d, dim=2) | |
| hd = D // self.nh | |
| q = q.view(B, T, self.nh, hd).transpose(1, 2) | |
| k = k.view(B, T, self.nh, hd).transpose(1, 2) | |
| v = v.view(B, T, self.nh, hd).transpose(1, 2) | |
| if self.cfg.qk_norm: | |
| q = self.q_norm(q) | |
| k = self.k_norm(k) | |
| if self.cfg.pos == "rope": | |
| q = apply_rope(q, self.cfg.rope_theta) | |
| k = apply_rope(k, self.cfg.rope_theta) | |
| if self.cfg.value_residual: | |
| if self.is_first: | |
| v0 = v | |
| else: | |
| v = v + self.vr_lambda * v0 | |
| y = F.scaled_dot_product_attention(q, k, v, is_causal=True) | |
| y = y.transpose(1, 2).contiguous().view(B, T, D) | |
| x = x + self.proj(y) | |
| x = x + self.mlp(self.ln2(x)) | |
| return x, v0 | |
| class GPT(nn.Module): | |
| def __init__(self, vocab, n_layer, n_embd, n_head, block, cfg): | |
| super().__init__() | |
| self.cfg = cfg | |
| self.tok = nn.Embedding(vocab, n_embd) | |
| self.use_rope = cfg.pos == "rope" | |
| if not self.use_rope: | |
| self.pos = nn.Embedding(block, n_embd) | |
| self.blocks = nn.ModuleList([Block(n_embd, n_head, block, cfg, is_first=(i == 0)) for i in range(n_layer)]) | |
| self.lnf = make_norm(n_embd, cfg) | |
| self.head = nn.Linear(n_embd, vocab, bias=False) | |
| self.head.weight = self.tok.weight # tie | |
| self.block = block | |
| self.apply(self._init) | |
| def _init(self, m): | |
| if isinstance(m, nn.Linear): | |
| nn.init.normal_(m.weight, 0.0, 0.02) | |
| if m.bias is not None: | |
| nn.init.zeros_(m.bias) | |
| elif isinstance(m, nn.Embedding): | |
| nn.init.normal_(m.weight, 0.0, 0.02) | |
| def forward(self, idx, targets=None): | |
| B, T = idx.size() | |
| x = self.tok(idx) | |
| if not self.use_rope: | |
| pos = torch.arange(T, device=idx.device) | |
| x = x + self.pos(pos)[None] | |
| v0 = None | |
| for b in self.blocks: | |
| x, v0 = b(x, v0) | |
| logits = self.head(self.lnf(x)) | |
| loss = None | |
| if targets is not None: | |
| loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1)) | |
| return logits, loss | |
| def main(): | |
| a = get_args() | |
| out_dir = a.out_dir or os.path.join(a.data_dir, "runs", a.run_id) | |
| os.makedirs(out_dir, exist_ok=True) | |
| log_path = os.path.join(out_dir, "train.log") | |
| metrics_path = os.path.join(out_dir, "metrics.jsonl") | |
| ckpt_path = os.path.join(out_dir, "ckpt.pt") | |
| def log(msg): | |
| line = f"[{time.strftime('%H:%M:%S')}] {msg}" | |
| print(line, flush=True) | |
| with open(log_path, "a", encoding="utf-8") as f: | |
| f.write(line + "\n") | |
| events_path = a.events_jsonl | |
| def emit(kind, step, metrics=None, **extra): | |
| if not events_path: | |
| return | |
| rec = {"kind": kind, "updates": int(step), "tokens": int(step) * a.block * a.batch} | |
| if metrics: | |
| rec["metrics"] = metrics | |
| rec.update(extra) | |
| with open(events_path, "a", encoding="utf-8") as f: | |
| f.write(json.dumps(rec) + "\n") | |
| torch.manual_seed(a.seed) | |
| torch.backends.cuda.matmul.allow_tf32 = True | |
| torch.backends.cudnn.allow_tf32 = True | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| use_bf16 = device.type == "cuda" and torch.cuda.is_bf16_supported() | |
| ptdtype = torch.bfloat16 if use_bf16 else torch.float32 | |
| ctx = torch.amp.autocast(device_type=device.type, dtype=ptdtype) if device.type == "cuda" else nullcontext() | |
| log(f"device={device} bf16={use_bf16} dev={torch.cuda.get_device_name(0) if device.type=='cuda' else 'cpu'}") | |
| if a.synthetic: | |
| rng = np.random.default_rng(a.seed) | |
| train_data = rng.integers(0, 256, size=8_000_000, dtype=np.uint8) | |
| val_data = train_data[:200_000] | |
| log("SYNTHETIC uint8 (smoke throughput, zero danych PII)") | |
| else: | |
| train_data = np.memmap(os.path.join(a.data_dir, "train.bin"), dtype=np.dtype(a.dtype), mode="r") | |
| val_data = np.memmap(os.path.join(a.data_dir, "val.bin"), dtype=np.dtype(a.dtype), mode="r") | |
| log(f"dane: train={len(train_data):,}B block={a.block} batch={a.batch} tok/step={a.block*a.batch:,}") | |
| def _make_batch_cpu(split, generator=None): | |
| """Wektoryzowane budowanie batcha: JEDEN numpy fancy-index zamiast | |
| python-loop per-item. sliding_window_view daje strided-view (N-block, block+1) | |
| BEZ kopiowania; windows[ix] materializuje tylko wybrane wiersze naraz. | |
| Zwraca (x,y) long CPU (pinned jesli cuda). Rozklad batchy IDENTYCZNY jak | |
| stary torch.stack-loop: x=data[i:i+block], y=data[i+1:i+1+block].""" | |
| data = train_data if split == "train" else val_data | |
| ix = torch.randint(len(data) - a.block - 1, (a.batch,), generator=generator) | |
| # (N-block, block+1) view; jeden fancy-index kopiuje wybrane okna | |
| windows = np.lib.stride_tricks.sliding_window_view(data, a.block + 1) | |
| sel = windows[ix.numpy()] # (batch, block+1) materialized | |
| x = torch.from_numpy(sel[:, :-1].astype(np.int64)) # astype -> contiguous copy | |
| y = torch.from_numpy(sel[:, 1:].astype(np.int64)) | |
| if device.type == "cuda": | |
| x = x.pin_memory(); y = y.pin_memory() | |
| return x, y | |
| def _to_device(x, y): | |
| if device.type == "cuda": | |
| return x.to(device, non_blocking=True), y.to(device, non_blocking=True) | |
| return x.to(device), y.to(device) | |
| def get_batch(split, generator=None): | |
| return _to_device(*_make_batch_cpu(split, generator)) | |
| class Prefetcher: | |
| """Async double-buffer: 1 background-thread buduje NASTEPNY batch na CPU | |
| (pinned) podczas gdy GPU liczy biezacy. queue depth=2. Konsument robi | |
| .next() -> H2D-copy (non_blocking) w watku glownym. Watek uzywa wlasnego | |
| torch.Generator (seeded), wiec ciag train-batchy jest deterministyczny i | |
| NIEZALEZNY od timingu watku oraz od RNG val-loopa (dystrybucja bez zmian).""" | |
| def __init__(self, split, generator, depth=2): | |
| self.split = split | |
| self.gen = generator | |
| self.q = queue.Queue(maxsize=depth) | |
| self._stop = threading.Event() | |
| self.t = threading.Thread(target=self._worker, daemon=True) | |
| self.t.start() | |
| def _worker(self): | |
| while not self._stop.is_set(): | |
| try: | |
| item = _make_batch_cpu(self.split, self.gen) | |
| except Exception as e: # przekaz blad do konsumenta | |
| self.q.put(e) | |
| return | |
| while not self._stop.is_set(): | |
| try: | |
| self.q.put(item, timeout=0.5) | |
| break | |
| except queue.Full: | |
| continue | |
| def next(self): | |
| item = self.q.get() | |
| if isinstance(item, Exception): | |
| raise item | |
| return _to_device(*item) | |
| def close(self): | |
| self._stop.set() | |
| # opróżnij kolejke zeby watek nie zawisl na put() | |
| try: | |
| self.q.get_nowait() | |
| except queue.Empty: | |
| pass | |
| raw_model = GPT(a.vocab, a.n_layer, a.n_embd, a.n_head, a.block, a).to(device) | |
| nparam = sum(p.numel() for p in raw_model.parameters()) | |
| log(f"model GPT-ref: {nparam/1e6:.1f}M param (L{a.n_layer} d{a.n_embd} h{a.n_head})") | |
| opt = build_optimizer(a, raw_model) | |
| log(f"optimizer={a.optimizer}" + (f" muon_lr={a.muon_lr} (mult={a.muon_lr/a.lr:.1f}x)" if a.optimizer == "muon" else "")) | |
| model = torch.compile(raw_model) if a.compile else raw_model | |
| if a.compile: | |
| log("torch.compile enabled") | |
| start_step = 0 | |
| if a.resume and os.path.exists(ckpt_path): | |
| ck = torch.load(ckpt_path, map_location=device) | |
| raw_model.load_state_dict(ck["model"]); opt.load_state_dict(ck["opt"]); start_step = ck["step"] | |
| log(f"RESUME @ {start_step}") | |
| def lr_at(s): | |
| if s < a.warmup: | |
| return a.lr * (s + 1) / a.warmup | |
| if s >= a.steps: | |
| return a.min_lr | |
| r = (s - a.warmup) / max(1, a.steps - a.warmup) | |
| return a.min_lr + 0.5 * (a.lr - a.min_lr) * (1 + math.cos(math.pi * r)) | |
| def eval_val(): | |
| model.eval() | |
| ls = [] | |
| for _ in range(a.eval_iters): | |
| xb, yb = get_batch("val") | |
| with ctx: | |
| _, loss = model(xb, yb) | |
| ls.append(loss.item()) | |
| model.train() | |
| return sum(ls) / len(ls) | |
| model.train() | |
| log(f"START gpt-ref: steps={a.steps} (od {start_step}) lr={a.lr}->{a.min_lr}") | |
| # dedykowany seeded generator dla train-prefetchera (determinizm niezalezny | |
| # od RNG val-loopa i timingu watku; ta sama dystrybucja co global-RNG) | |
| train_gen = torch.Generator() | |
| train_gen.manual_seed(a.seed) | |
| prefetcher = Prefetcher("train", train_gen) | |
| t0 = time.time(); running = 0.0 | |
| for step in range(start_step, a.steps): | |
| lr = lr_at(step) | |
| for g in opt.param_groups: | |
| g["lr"] = lr * g.get("lr_mult", 1.0) | |
| xb, yb = prefetcher.next() | |
| with ctx: | |
| _, loss = model(xb, yb) | |
| loss.backward() | |
| gn = torch.nn.utils.clip_grad_norm_(model.parameters(), a.grad_clip) if a.grad_clip > 0 else 0.0 | |
| opt.step(); opt.zero_grad(set_to_none=True) | |
| running += loss.item() | |
| if (step + 1) % a.log_every == 0: | |
| dt = time.time() - t0 | |
| tok_s = a.log_every * a.block * a.batch / dt | |
| mem = torch.cuda.max_memory_allocated()/1e9 if device.type == "cuda" else 0.0 | |
| 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") | |
| with open(metrics_path, "a", encoding="utf-8") as f: | |
| f.write(json.dumps({"step": step+1, "loss": running/a.log_every, "lr": lr, "tok_s": tok_s}) + "\n") | |
| emit("update", step + 1, metrics={"loss": running / a.log_every, "tokens_per_second": tok_s, | |
| "gradient_norm": float(gn), "learning_rate": lr}) | |
| running = 0.0; t0 = time.time() | |
| if (step + 1) % a.eval_every == 0: | |
| vloss = eval_val() | |
| log(f" >> VAL loss {vloss:.4f} @ {step+1}") | |
| emit("evaluation", step + 1, metrics={"loss": vloss}) | |
| if (step + 1) % a.ckpt_every == 0 and not a.synthetic: | |
| torch.save({"model": raw_model.state_dict(), "opt": opt.state_dict(), "step": step + 1, | |
| "config": {"vocab": a.vocab, "n_layer": a.n_layer, "n_embd": a.n_embd, | |
| "n_head": a.n_head, "block": a.block, "norm": a.norm, | |
| "norm_eps": a.norm_eps, "pos": a.pos, "rope_theta": a.rope_theta, | |
| "ffn": a.ffn, "ffn_mult": a.ffn_mult, "value_residual": a.value_residual, | |
| "qk_norm": a.qk_norm}}, ckpt_path) | |
| log(f"ckpt @ {step+1}") | |
| if a.events_jsonl: | |
| import hashlib as _hl | |
| _h = _hl.sha256() | |
| with open(ckpt_path, "rb") as _cf: | |
| for _chunk in iter(lambda: _cf.read(1 << 20), b""): | |
| _h.update(_chunk) | |
| emit("checkpoint", step + 1, sha256=_h.hexdigest()) | |
| prefetcher.close() | |
| log("DONE") | |
| emit("end", a.steps, status="completed") | |
| if __name__ == "__main__": | |
| main() | |