model-def (GPT-ref Qwen3-arch) for same-harness loader
Browse files- train_gpt_ref.py +142 -23
train_gpt_ref.py
CHANGED
|
@@ -60,6 +60,25 @@ def get_args():
|
|
| 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 |
|
|
@@ -179,38 +198,103 @@ def build_optimizer(a, model):
|
|
| 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 =
|
| 186 |
-
self.ln2 =
|
| 187 |
self.qkv = nn.Linear(d, 3 * d)
|
| 188 |
self.proj = nn.Linear(d, d)
|
| 189 |
-
|
|
|
|
|
|
|
|
|
|
| 190 |
self.nh = nh
|
| 191 |
self.d = d
|
| 192 |
-
|
| 193 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 194 |
B, T, D = x.size()
|
| 195 |
h = self.ln1(x)
|
| 196 |
q, k, v = self.qkv(h).split(self.d, dim=2)
|
| 197 |
-
|
| 198 |
-
|
| 199 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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.
|
| 212 |
-
|
| 213 |
-
|
|
|
|
|
|
|
| 214 |
self.head = nn.Linear(n_embd, vocab, bias=False)
|
| 215 |
self.head.weight = self.tok.weight # tie
|
| 216 |
self.block = block
|
|
@@ -226,10 +310,13 @@ class GPT(nn.Module):
|
|
| 226 |
|
| 227 |
def forward(self, idx, targets=None):
|
| 228 |
B, T = idx.size()
|
| 229 |
-
|
| 230 |
-
|
|
|
|
|
|
|
|
|
|
| 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:
|
|
@@ -251,6 +338,18 @@ def main():
|
|
| 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
|
|
@@ -337,16 +436,19 @@ def main():
|
|
| 337 |
except queue.Empty:
|
| 338 |
pass
|
| 339 |
|
| 340 |
-
|
| 341 |
-
nparam = sum(p.numel() for p in
|
| 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,
|
| 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 |
-
|
| 350 |
log(f"RESUME @ {start_step}")
|
| 351 |
|
| 352 |
def lr_at(s):
|
|
@@ -395,14 +497,31 @@ def main():
|
|
| 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 |
-
|
|
|
|
|
|
|
| 401 |
if (step + 1) % a.ckpt_every == 0 and not a.synthetic:
|
| 402 |
-
torch.save({"model":
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 403 |
log(f"ckpt @ {step+1}")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 404 |
prefetcher.close()
|
| 405 |
log("DONE")
|
|
|
|
| 406 |
|
| 407 |
|
| 408 |
if __name__ == "__main__":
|
|
|
|
| 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 |
+
p.add_argument("--norm", choices=["layernorm", "rmsnorm"], default="layernorm",
|
| 64 |
+
help="layernorm (default, backward-compat) | rmsnorm (Qwen3-style, fp32-compute)")
|
| 65 |
+
p.add_argument("--norm-eps", type=float, default=1e-6)
|
| 66 |
+
p.add_argument("--pos", choices=["learned", "rope"], default="learned",
|
| 67 |
+
help="learned (default) | rope (parameter-free RoPE na q,k; usuwa learned pos-embedding)")
|
| 68 |
+
p.add_argument("--rope-theta", type=float, default=100000.0,
|
| 69 |
+
help="RoPE theta; top-3-board (JugnuLM/GPT-X2) uzywaja 100000 (nie-default-10K)")
|
| 70 |
+
p.add_argument("--ffn", choices=["gelu", "swiglu"], default="gelu",
|
| 71 |
+
help="gelu (default 4x MLP) | swiglu (Qwen3 gated-MLP, hidden=--ffn-mult*d)")
|
| 72 |
+
p.add_argument("--ffn-mult", type=float, default=2.667,
|
| 73 |
+
help="mnoznik hidden dla swiglu (~param-parity z 4x-gelu przy 8/3)")
|
| 74 |
+
p.add_argument("--value-residual", action="store_true",
|
| 75 |
+
help="ResFormer value-residuals: v_l += lambda_l*v0 (lambda init 0), ARC-targeted")
|
| 76 |
+
p.add_argument("--qk-norm", action="store_true",
|
| 77 |
+
help="Qwen3 QK-Norm: RMSNorm per-head na Q,K przed-attention (stabilnosc z Muon/high-LR)")
|
| 78 |
+
p.add_argument("--events-jsonl", default=None,
|
| 79 |
+
help="jesli podane: emituj events.jsonl (fabryka-track sidecar-format: update/evaluation/checkpoint/end)")
|
| 80 |
+
p.add_argument("--compile", action="store_true",
|
| 81 |
+
help="torch.compile model (2-3x throughput; state_dict zapisywany bez _orig_mod prefix via raw_model)")
|
| 82 |
return p.parse_args()
|
| 83 |
|
| 84 |
|
|
|
|
| 198 |
return MuonWithAuxAdam(muon, adamw)
|
| 199 |
|
| 200 |
|
| 201 |
+
class RMSNorm(nn.Module):
|
| 202 |
+
"""Qwen3-style RMSNorm (fp32-compute dla stabilnosci). 1D weight -> AdamW w split-Muon."""
|
| 203 |
+
def __init__(self, d, eps=1e-6):
|
| 204 |
+
super().__init__()
|
| 205 |
+
self.weight = nn.Parameter(torch.ones(d))
|
| 206 |
+
self.eps = eps
|
| 207 |
+
|
| 208 |
+
def forward(self, x):
|
| 209 |
+
return x * torch.rsqrt(x.float().pow(2).mean(-1, keepdim=True) + self.eps).to(x.dtype) * self.weight
|
| 210 |
+
|
| 211 |
+
|
| 212 |
+
def make_norm(d, cfg):
|
| 213 |
+
return RMSNorm(d, cfg.norm_eps) if cfg.norm == "rmsnorm" else nn.LayerNorm(d)
|
| 214 |
+
|
| 215 |
+
|
| 216 |
+
def apply_rope(x, base=100000.0):
|
| 217 |
+
"""Parameter-free RoPE na [B,H,T,D] (interleaved-conv, port z qwen_model.py). Train==eval
|
| 218 |
+
MUSZA uzywac tej samej konwencji (self-contained eval -> spojne)."""
|
| 219 |
+
_, _, T, dim = x.shape
|
| 220 |
+
pos = torch.arange(T, device=x.device, dtype=torch.float32)
|
| 221 |
+
freq = 1.0 / (base ** (torch.arange(0, dim, 2, device=x.device, dtype=torch.float32) / dim))
|
| 222 |
+
ang = torch.outer(pos, freq)
|
| 223 |
+
cos, sin = ang.cos().to(x.dtype)[None, None], ang.sin().to(x.dtype)[None, None]
|
| 224 |
+
even, odd = x[..., ::2], x[..., 1::2]
|
| 225 |
+
return torch.stack((even * cos - odd * sin, even * sin + odd * cos), dim=-1).flatten(-2)
|
| 226 |
+
|
| 227 |
+
|
| 228 |
+
class SwiGLU(nn.Module):
|
| 229 |
+
"""Qwen3 gated-MLP: down(silu(gate(x))*up(x)). 3x 2D bez-bias -> wszystkie do Muon."""
|
| 230 |
+
def __init__(self, d, hidden):
|
| 231 |
+
super().__init__()
|
| 232 |
+
self.gate = nn.Linear(d, hidden, bias=False)
|
| 233 |
+
self.up = nn.Linear(d, hidden, bias=False)
|
| 234 |
+
self.down = nn.Linear(hidden, d, bias=False)
|
| 235 |
+
|
| 236 |
+
def forward(self, x):
|
| 237 |
+
return self.down(F.silu(self.gate(x)) * self.up(x))
|
| 238 |
+
|
| 239 |
+
|
| 240 |
class Block(nn.Module):
|
| 241 |
+
def __init__(self, d, nh, block, cfg, is_first=False):
|
| 242 |
super().__init__()
|
| 243 |
+
self.ln1 = make_norm(d, cfg)
|
| 244 |
+
self.ln2 = make_norm(d, cfg)
|
| 245 |
self.qkv = nn.Linear(d, 3 * d)
|
| 246 |
self.proj = nn.Linear(d, d)
|
| 247 |
+
if cfg.ffn == "swiglu":
|
| 248 |
+
self.mlp = SwiGLU(d, int(round(cfg.ffn_mult * d)))
|
| 249 |
+
else:
|
| 250 |
+
self.mlp = nn.Sequential(nn.Linear(d, 4 * d), nn.GELU(), nn.Linear(4 * d, d))
|
| 251 |
self.nh = nh
|
| 252 |
self.d = d
|
| 253 |
+
self.cfg = cfg
|
| 254 |
+
self.is_first = is_first
|
| 255 |
+
if cfg.value_residual and not is_first:
|
| 256 |
+
self.vr_lambda = nn.Parameter(torch.zeros(1))
|
| 257 |
+
if cfg.qk_norm:
|
| 258 |
+
hd = d // nh
|
| 259 |
+
self.q_norm = RMSNorm(hd, cfg.norm_eps)
|
| 260 |
+
self.k_norm = RMSNorm(hd, cfg.norm_eps)
|
| 261 |
+
|
| 262 |
+
def forward(self, x, v0=None):
|
| 263 |
B, T, D = x.size()
|
| 264 |
h = self.ln1(x)
|
| 265 |
q, k, v = self.qkv(h).split(self.d, dim=2)
|
| 266 |
+
hd = D // self.nh
|
| 267 |
+
q = q.view(B, T, self.nh, hd).transpose(1, 2)
|
| 268 |
+
k = k.view(B, T, self.nh, hd).transpose(1, 2)
|
| 269 |
+
v = v.view(B, T, self.nh, hd).transpose(1, 2)
|
| 270 |
+
if self.cfg.qk_norm:
|
| 271 |
+
q = self.q_norm(q)
|
| 272 |
+
k = self.k_norm(k)
|
| 273 |
+
if self.cfg.pos == "rope":
|
| 274 |
+
q = apply_rope(q, self.cfg.rope_theta)
|
| 275 |
+
k = apply_rope(k, self.cfg.rope_theta)
|
| 276 |
+
if self.cfg.value_residual:
|
| 277 |
+
if self.is_first:
|
| 278 |
+
v0 = v
|
| 279 |
+
else:
|
| 280 |
+
v = v + self.vr_lambda * v0
|
| 281 |
y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
|
| 282 |
y = y.transpose(1, 2).contiguous().view(B, T, D)
|
| 283 |
x = x + self.proj(y)
|
| 284 |
x = x + self.mlp(self.ln2(x))
|
| 285 |
+
return x, v0
|
| 286 |
|
| 287 |
|
| 288 |
class GPT(nn.Module):
|
| 289 |
+
def __init__(self, vocab, n_layer, n_embd, n_head, block, cfg):
|
| 290 |
super().__init__()
|
| 291 |
+
self.cfg = cfg
|
| 292 |
self.tok = nn.Embedding(vocab, n_embd)
|
| 293 |
+
self.use_rope = cfg.pos == "rope"
|
| 294 |
+
if not self.use_rope:
|
| 295 |
+
self.pos = nn.Embedding(block, n_embd)
|
| 296 |
+
self.blocks = nn.ModuleList([Block(n_embd, n_head, block, cfg, is_first=(i == 0)) for i in range(n_layer)])
|
| 297 |
+
self.lnf = make_norm(n_embd, cfg)
|
| 298 |
self.head = nn.Linear(n_embd, vocab, bias=False)
|
| 299 |
self.head.weight = self.tok.weight # tie
|
| 300 |
self.block = block
|
|
|
|
| 310 |
|
| 311 |
def forward(self, idx, targets=None):
|
| 312 |
B, T = idx.size()
|
| 313 |
+
x = self.tok(idx)
|
| 314 |
+
if not self.use_rope:
|
| 315 |
+
pos = torch.arange(T, device=idx.device)
|
| 316 |
+
x = x + self.pos(pos)[None]
|
| 317 |
+
v0 = None
|
| 318 |
for b in self.blocks:
|
| 319 |
+
x, v0 = b(x, v0)
|
| 320 |
logits = self.head(self.lnf(x))
|
| 321 |
loss = None
|
| 322 |
if targets is not None:
|
|
|
|
| 338 |
with open(log_path, "a", encoding="utf-8") as f:
|
| 339 |
f.write(line + "\n")
|
| 340 |
|
| 341 |
+
events_path = a.events_jsonl
|
| 342 |
+
|
| 343 |
+
def emit(kind, step, metrics=None, **extra):
|
| 344 |
+
if not events_path:
|
| 345 |
+
return
|
| 346 |
+
rec = {"kind": kind, "updates": int(step), "tokens": int(step) * a.block * a.batch}
|
| 347 |
+
if metrics:
|
| 348 |
+
rec["metrics"] = metrics
|
| 349 |
+
rec.update(extra)
|
| 350 |
+
with open(events_path, "a", encoding="utf-8") as f:
|
| 351 |
+
f.write(json.dumps(rec) + "\n")
|
| 352 |
+
|
| 353 |
torch.manual_seed(a.seed)
|
| 354 |
torch.backends.cuda.matmul.allow_tf32 = True
|
| 355 |
torch.backends.cudnn.allow_tf32 = True
|
|
|
|
| 436 |
except queue.Empty:
|
| 437 |
pass
|
| 438 |
|
| 439 |
+
raw_model = GPT(a.vocab, a.n_layer, a.n_embd, a.n_head, a.block, a).to(device)
|
| 440 |
+
nparam = sum(p.numel() for p in raw_model.parameters())
|
| 441 |
log(f"model GPT-ref: {nparam/1e6:.1f}M param (L{a.n_layer} d{a.n_embd} h{a.n_head})")
|
| 442 |
+
opt = build_optimizer(a, raw_model)
|
| 443 |
log(f"optimizer={a.optimizer}" + (f" muon_lr={a.muon_lr} (mult={a.muon_lr/a.lr:.1f}x)" if a.optimizer == "muon" else ""))
|
| 444 |
+
model = torch.compile(raw_model) if a.compile else raw_model
|
| 445 |
+
if a.compile:
|
| 446 |
+
log("torch.compile enabled")
|
| 447 |
|
| 448 |
start_step = 0
|
| 449 |
if a.resume and os.path.exists(ckpt_path):
|
| 450 |
ck = torch.load(ckpt_path, map_location=device)
|
| 451 |
+
raw_model.load_state_dict(ck["model"]); opt.load_state_dict(ck["opt"]); start_step = ck["step"]
|
| 452 |
log(f"RESUME @ {start_step}")
|
| 453 |
|
| 454 |
def lr_at(s):
|
|
|
|
| 497 |
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")
|
| 498 |
with open(metrics_path, "a", encoding="utf-8") as f:
|
| 499 |
f.write(json.dumps({"step": step+1, "loss": running/a.log_every, "lr": lr, "tok_s": tok_s}) + "\n")
|
| 500 |
+
emit("update", step + 1, metrics={"loss": running / a.log_every, "tokens_per_second": tok_s,
|
| 501 |
+
"gradient_norm": float(gn), "learning_rate": lr})
|
| 502 |
running = 0.0; t0 = time.time()
|
| 503 |
if (step + 1) % a.eval_every == 0:
|
| 504 |
+
vloss = eval_val()
|
| 505 |
+
log(f" >> VAL loss {vloss:.4f} @ {step+1}")
|
| 506 |
+
emit("evaluation", step + 1, metrics={"loss": vloss})
|
| 507 |
if (step + 1) % a.ckpt_every == 0 and not a.synthetic:
|
| 508 |
+
torch.save({"model": raw_model.state_dict(), "opt": opt.state_dict(), "step": step + 1,
|
| 509 |
+
"config": {"vocab": a.vocab, "n_layer": a.n_layer, "n_embd": a.n_embd,
|
| 510 |
+
"n_head": a.n_head, "block": a.block, "norm": a.norm,
|
| 511 |
+
"norm_eps": a.norm_eps, "pos": a.pos, "rope_theta": a.rope_theta,
|
| 512 |
+
"ffn": a.ffn, "ffn_mult": a.ffn_mult, "value_residual": a.value_residual,
|
| 513 |
+
"qk_norm": a.qk_norm}}, ckpt_path)
|
| 514 |
log(f"ckpt @ {step+1}")
|
| 515 |
+
if a.events_jsonl:
|
| 516 |
+
import hashlib as _hl
|
| 517 |
+
_h = _hl.sha256()
|
| 518 |
+
with open(ckpt_path, "rb") as _cf:
|
| 519 |
+
for _chunk in iter(lambda: _cf.read(1 << 20), b""):
|
| 520 |
+
_h.update(_chunk)
|
| 521 |
+
emit("checkpoint", step + 1, sha256=_h.hexdigest())
|
| 522 |
prefetcher.close()
|
| 523 |
log("DONE")
|
| 524 |
+
emit("end", a.steps, status="completed")
|
| 525 |
|
| 526 |
|
| 527 |
if __name__ == "__main__":
|