Maggio33 commited on
Commit
e88580a
·
verified ·
1 Parent(s): 0364e96

model-def (GPT-ref Qwen3-arch) for same-harness loader

Browse files
Files changed (1) hide show
  1. 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 = 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
@@ -226,10 +310,13 @@ class GPT(nn.Module):
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:
@@ -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
- 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):
@@ -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
- 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__":
 
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__":