File size: 35,782 Bytes
d8c733f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 | """Guido-small — pretrain ~200M (24L×768), math-first chattabile.
Fork di `train_distill_start.py` SENZA KD (il teacher non aiuta a questa scala —
vedi SESSION_HANDOFF.md §A). Differenze chiave:
- shape 200M (24L × 768 × 12h × 64hd × 3072ff)
- MultiShardMixture reader: legge corpus_v2 (7 shard per-dataset, NON shufflati a
write-time) con chunk-shuffle a read-time (NO rimpiazzo, copertura 100%, ordine
casuale). Pesi mixture 70% math / 22% fineweb / 8% cosmopedia in HP.mixture.
- loss-trace per-step → CSV (loss, gnorm, lr, math_frac) con UN sync ogni
`train_log_every` step (buffer su GPU). math_frac correla spike↔batch math-heavy.
- looping opzionale (env LOOP_STYLE / LOOP_ACTIVATION_FRAC), default OFF per la baseline.
Run:
torchrun --standalone --nproc_per_node=4 train_guido_small.py
Pre-req: corpus_v2 tokenizzato (Mathstral SPM 32k, uint32) in $FAST/corpus_v2/<name>/.
"""
from __future__ import annotations
import csv
import os
import re
import sys
import time
from pathlib import Path
import numpy as np
import torch
import torch.distributed as dist
import torch.nn as nn
import torch.nn.functional as F
from torch import Tensor
from torch.nn.parallel import DistributedDataParallel as DDP
# ---- Vathos backbone ----
APLOS_PATH = os.environ.get("APLOS_PATH", "/leonardo_work/IscrC_YENDRI/paerle/PiCO/aplos")
if APLOS_PATH not in sys.path:
sys.path.insert(0, APLOS_PATH)
from Vathos._basics import (Builder, RMSNorm as VRMSNorm, ReLU2, LeakyReLU2,
VariableUDLP, VariableGatedUDLP, set_vathos_mode)
from Vathos._spatials import MultiheadAttentionMixer, MultiheadGatedAttentionMixer, RoPE
from Vathos.blocks import PiCOFormer as VathosPiCOFormer, SmearGate
set_vathos_mode("production")
# FAST RMSNorm fused
def _fast_rmsnorm_forward(self, x):
return F.rms_norm(x, (x.size(-1),), self.weight, self.eps)
VRMSNorm.forward = _fast_rmsnorm_forward
# ---- Shard reader ----
sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "PiCO2"))
from pico2.shards import ShardReader
from cut_cross_entropy import linear_cross_entropy
# =============================================================================
# HYPERPARAMETERS
# =============================================================================
class HP:
# data — corpus_v2 multi-shard mixture (Mathstral SPM 32k)
corpus_root = os.environ.get(
"CORPUS_ROOT",
"/leonardo_scratch/fast/IscrC_YENDRI/mprignan/corpus_v2",
)
# Mixture preset selezionabile via MIXTURE_PRESET env. Default = 70/22/8 (run baseline 200M).
# "v3_noloops" = fineweb 22%→15% (≈2/3), delta freed → math (77/15/8 total).
_MIXTURE_PRESETS = {
"default": {
"openmath_full": 0.34, "tinygsm": 0.14, "openmathreasoning": 0.10,
"numina_15": 0.07, "numina_cot": 0.05, # = 0.70 math
"fineweb_edu": 0.22, "cosmopedia": 0.08, # = 0.30 NL
},
"v3_noloops": {
"openmath_full": 0.37, "tinygsm": 0.15, "openmathreasoning": 0.11,
"numina_15": 0.08, "numina_cot": 0.06, # = 0.77 math
"fineweb_edu": 0.15, "cosmopedia": 0.08, # = 0.23 NL (fineweb ridotto a 2/3)
},
}
mixture = _MIXTURE_PRESETS[os.environ.get("MIXTURE_PRESET", "default")]
nl_datasets = ("fineweb_edu", "cosmopedia") # per math_frac logging
seq_len = 2048
bs_per_dev = int(os.environ.get("BS_PER_DEV", 32))
train_batch_tokens = 4 * bs_per_dev * 2048 # 4×32×2048 = 262k tok/step (4 GPU)
# model — shape 200M (scale-up dal 92M: d_model 512→768, d_ff 2048→3072, heads 8→12)
n_layers = int(os.environ.get("N_LAYERS", 24))
d_model = int(os.environ.get("D_MODEL", 768))
n_heads = int(os.environ.get("N_HEADS", 12))
head_dim = 64
d_ff = int(os.environ.get("D_FF", 3072))
rope_base = 1_000_000.0
logit_softcap = 30.0
tied_embeddings = True
qk_norm = True
# v3 features (locked)
use_gated_attn = True
# Channel mixer (FFN) variant. Knob aperto dall'audit A/B 2026-05-28 (vedi analysis/guido_ab/REPORT).
# "udlp" : VariableUDLP, contract(LeakyReLU²(expand(x))) (current default)
# "gated_udlp" : VariableGatedUDLP, UDLP·sigmoid(gate_proj(x[..., :K])) (sparse output gate)
# "leaky_reglu2" : LeakyReGLU² GLU, contract(LeakyReLU²(expand(x)) · up(x)) (act·value, +1 proj)
mlp_kind = os.environ.get("MLP_KIND", "udlp")
use_smear_gate = True
# gate_input_dim: 12 era stato scelto stile parameter-golf/sparse. Sul d=768 è ~1.5% di x:
# per il prossimo run testare 64 (8%) o 128 (17%) — vedi audit gate-activity (40% dei gate<0.1).
gate_input_dim = int(os.environ.get("GATE_INPUT_DIM", 12))
# optimizer (Muon LOCKED)
optimizer_kind = "muon"
# Orthogonalizer per Muon. Knob aperto dall'audit A/B 2026-05-28.
# "ns5" : Newton-Schulz quintic, coeff (3.4445,-4.7750,2.0315), Keller Jordan
# "polar_express": Polar Express (arxiv 2505.16932) — schedule di 8 coeff/step,
# convergenza più rapida. RACCOMANDATO MUON_BACKEND_STEPS=8 (default 5
# di NS5 è sotto-iterato per PE → convergenza parziale).
orthogonalizer = os.environ.get("ORTHOGONALIZER", "ns5")
muon_backend_steps = int(os.environ.get("MUON_BACKEND_STEPS", 5))
matrix_lr = 0.01
tied_embed_lr = 0.005
scalar_lr = 0.04
muon_momentum = 0.95
muon_momentum_warmup_start = 0.85
muon_momentum_warmup_steps = 200
beta1 = 0.9
beta2 = 0.95
adam_eps = 1e-8
grad_clip = 1.0
# schedule (WSD: warmup → plateau → linear warmdown a lr_min_scale)
iterations = int(os.environ.get("ITERATIONS", 15000)) # ~4B tok; CALIBRA dopo lo smoke
warmup_steps = int(os.environ.get("WARMUP", 300))
warmdown_iters = int(os.environ.get("WARMDOWN", int(0.60 * iterations))) # 60%
lr_min_scale = float(os.environ.get("LR_MIN_SCALE", 0.001))
# === Layer looping (opzionale, default OFF per baseline) ===
# loop_v2 sul 92M: +4.5pp GSM8K @ activation_frac 0.30 (durante plateau LR).
loop_style = os.environ.get("LOOP_STYLE", "")
loop_activation_frac = float(os.environ.get("LOOP_ACTIVATION_FRAC", 0.30))
loop_pre_warm = os.environ.get("LOOP_PRE_WARM", "1") == "1"
# logging / checkpoint
train_log_every = int(os.environ.get("LOG_EVERY", 50)) # dump CSV ogni N (loss PER-STEP nel buffer)
loss_trace_csv = os.environ.get(
"LOSS_TRACE_CSV",
"/leonardo_work/IscrC_YENDRI/paerle/PiCO/Guido-1/PiCO2_test/logs/guido_small_loss_trace.csv",
)
ckpt_dir = os.environ.get(
"CKPT_DIR",
"/leonardo_work/IscrC_YENDRI/paerle/PiCO/ckpts/pico_guido_small",
)
save_final = True
save_interval = int(os.environ.get("SAVE_INTERVAL", 2000))
keep_n_recent = 5
seed = int(os.environ.get("SEED", 1337))
# =============================================================================
# MUON OPTIMIZER
# =============================================================================
def zeropower_via_newtonschulz5(G: Tensor, steps: int = 5, eps: float = 1e-7) -> Tensor:
a, b, c = (3.4445, -4.7750, 2.0315)
X = G.bfloat16()
X /= X.norm() + eps
transposed = G.size(0) > G.size(1)
if transposed:
X = X.T
for _ in range(steps):
A = X @ X.T
B = b * A + c * A @ A
X = a * X + B @ X
return X.T if transposed else X
# Polar Express coefficients (arxiv 2505.16932). Stessa forma di NS5 ma SCHEDULE di 8 set:
# iter 1 aggressivo per matrici mal condizionate, iter 7-8 = steady-state Halley (3/8, -10/8, 15/8).
# Source: NVIDIA NeMo Emerging-Optimizers (authoritative, production-tested).
# https://docs.nvidia.com/nemo/emerging-optimizers/0.1.0/_modules/emerging_optimizers/orthogonalized_optimizers/muon_utils.html
_POLAR_EXPRESS_COEFFS = (
(8.2051, -22.9019, 16.4607),
(4.0664, -2.8612, 0.5184),
(3.9096, -2.8234, 0.5250),
(3.2856, -2.4153, 0.4853),
(2.2779, -1.6198, 0.3985),
(1.8726, -1.2307, 0.3585),
(1.8564, -1.2132, 0.3568),
(1.8750, -1.2500, 0.3750),
)
def zeropower_via_polar_express(G: Tensor, steps: int = 8, eps: float = 1e-7) -> Tensor:
"""Polar Express orthogonalizer — convergenza più rapida di NS5 quintic.
Iteration: X = a*X + (b*A + c*A@A)@X con A=X@X^T, a/b/c SCHEDULATI per step.
Normalizzazione Frobenius (identica a NS5). Per convergenza piena servono ≥7 step
(raccomandato MUON_BACKEND_STEPS=8). Oltre 8 step usa lo steady-state (Halley)."""
X = G.bfloat16()
X /= X.norm() + eps
transposed = G.size(0) > G.size(1)
if transposed:
X = X.T
n_coeffs = len(_POLAR_EXPRESS_COEFFS)
for k in range(steps):
a, b, c = _POLAR_EXPRESS_COEFFS[k if k < n_coeffs else n_coeffs - 1]
A = X @ X.T
B = b * A + c * A @ A
X = a * X + B @ X
return X.T if transposed else X
class Muon(torch.optim.Optimizer):
def __init__(self, params, lr, momentum, backend_steps, nesterov=True):
super().__init__(params, dict(lr=lr, momentum=momentum,
backend_steps=backend_steps, nesterov=nesterov))
@torch.no_grad()
def step(self, closure=None):
distributed = dist.is_available() and dist.is_initialized()
world_size = dist.get_world_size() if distributed else 1
rank = dist.get_rank() if distributed else 0
for group in self.param_groups:
params = group["params"]
if not params:
continue
lr, mom, ns_steps, nesterov = group["lr"], group["momentum"], group["backend_steps"], group["nesterov"]
total = sum(int(p.numel()) for p in params)
updates_flat = torch.zeros(total, device=params[0].device, dtype=torch.bfloat16)
curr = 0
for i, p in enumerate(params):
if i % world_size == rank and p.grad is not None:
g = p.grad
state = self.state[p]
if "momentum_buffer" not in state:
state["momentum_buffer"] = torch.zeros_like(g)
buf = state["momentum_buffer"]
buf.mul_(mom).add_(g)
if nesterov:
g = g.add(buf, alpha=mom)
g = zeropower_via_newtonschulz5(g, steps=ns_steps)
g *= max(1, g.size(0) / g.size(1)) ** 0.5
updates_flat[curr:curr + p.numel()] = g.reshape(-1)
curr += p.numel()
if distributed:
dist.all_reduce(updates_flat, op=dist.ReduceOp.SUM)
curr = 0
for p in params:
u = updates_flat[curr:curr + p.numel()].view_as(p).to(p.dtype)
p.add_(u, alpha=-lr)
curr += p.numel()
# =============================================================================
# CHANNEL MIXER VARIANT — LeakyReGLU² GLU-style FFN
# =============================================================================
class LeakyReGLU2_FFN(nn.Module):
"""LeakyReLU² GLU: contract(LeakyReLU²(expand(x)) · up(x)).
GLU-style channel mixer (à la SwiGLU/ReGLU) con activation LeakyReLU². Aggiunge UN extra
proiezione `up` (d_model→M) rispetto a VariableUDLP → +50% param nell'FFN. Identity-at-init:
contract=0 ⇒ branch nullo all'init (gradiente fluisce attraverso `expand` e `up`).
Compatibile con Vathos `Builder`: signature (d_model, d_output, M, activation, dropout).
Param naming: `expand` (act path, ndim==2 → Muon), `up` (value path), `contract` (output).
"""
def __init__(self, d_model, d_output, M, activation=None, dropout=0.0):
super().__init__()
self.expand = nn.Linear(d_model, M, bias=False) # act path (the W in act(xW))
self.up = nn.Linear(d_model, M, bias=False) # value path (the V in ·xV)
self.contract = nn.Linear(M, d_output, bias=False)
self.activation = activation() if activation is not None else LeakyReLU2()
self.dropout = nn.Dropout(dropout)
self._init_weights()
def _init_weights(self):
torch.nn.init.orthogonal_(self.expand.weight)
torch.nn.init.orthogonal_(self.up.weight)
torch.nn.init.zeros_(self.contract.weight) # identity-at-init: branch nullo
def forward(self, x):
return self.dropout(self.contract(self.activation(self.expand(x)) * self.up(x)))
# =============================================================================
# DATA — MultiShardMixture: chunk-shuffle a read-time (NO rimpiazzo)
# =============================================================================
class ShuffledStream:
"""Un dataset → finestre (seq_len+1) in ordine PERMUTATO (shuffle senza rimpiazzo).
Chunk disgiunti per rank; re-permuta a fine epoca (sub-epoca = nessun wrap col budget attuale)."""
def __init__(self, shard_dir, rank, world_size, seq_len, seed):
self.reader = ShardReader(shard_dir)
self.win = seq_len + 1
self.n_chunks = self.reader.total_tokens // self.win
self.rank, self.ws, self.seed, self._epoch = rank, world_size, seed, 0
self._reperm()
def _reperm(self):
rng = np.random.default_rng(self.seed + 104729 * self._epoch)
self.my = rng.permutation(self.n_chunks)[self.rank::self.ws]
self.ptr = 0
self._epoch += 1
def next_window(self):
if self.ptr >= len(self.my):
self._reperm()
ci = int(self.my[self.ptr])
self.ptr += 1
return self.reader.read(ci * self.win, self.win) # np.uint32 [win]
class MultiShardMixture:
"""Weighted-pick del dataset per-sample + finestra chunk-shuffled. Ogni doc ≤1 volta/epoca."""
def __init__(self, mixture, corpus_root, rank, world_size, device, seq_len, seed):
self.names = list(mixture.keys())
self.w = np.array([mixture[n] for n in self.names], float)
self.w /= self.w.sum()
self.streams = {
n: ShuffledStream(str(Path(corpus_root) / n), rank, world_size, seq_len, seed + i * 7919)
for i, n in enumerate(self.names)
}
self.rng = np.random.default_rng(seed * 31 + rank)
self.device = device
def next_batch(self, bs):
picks = self.rng.choice(len(self.names), size=bs, p=self.w)
xs, ys, ids = [], [], []
for pi in picks:
w = self.streams[self.names[pi]].next_window()
t = torch.from_numpy(w.astype(np.int64))
xs.append(t[:-1])
ys.append(t[1:])
ids.append(self.names[pi])
return (torch.stack(xs).to(self.device, non_blocking=True),
torch.stack(ys).to(self.device, non_blocking=True), ids)
# =============================================================================
# LAYER LOOPING — parser + scalar gate init=0 (Universal Transformer style)
# =============================================================================
_LOOP_SEG_RE = re.compile(r"\(\s*([\d\s,]+?)\s*\)\s*x\s*(\d+)")
def parse_loop_style(s: str, n_layers: int):
"""Parse loop style stringa → lista di (indices, n_total_executions).
Esempio: "(3,4,5)x2, (8,9)x3" → [([3,4,5], 2), ([8,9], 3)]. Semantica xN: N esecuzioni
totali (1 originale + N-1 loop extra con gate). Validation: in-range, contigui, disgiunti, N>=2."""
s = s.strip()
if not s:
return []
stripped = re.sub(r"\s+", "", s)
pattern_strict = re.compile(r"\(([\d,]+)\)x(\d+)")
consumed = "".join(f"({a})x{b}" for a, b in pattern_strict.findall(stripped))
consumed_with_commas = ",".join(f"({a})x{b}" for a, b in pattern_strict.findall(stripped))
if stripped != consumed and stripped != consumed_with_commas:
raise ValueError(f"loop_style mal formato: {s!r}. Atteso: '(i,j,k)xN, (l,m)xM, ...'")
segments = []
used_indices = set()
for indices_str, n_str in _LOOP_SEG_RE.findall(s):
indices = sorted({int(x.strip()) for x in indices_str.split(",") if x.strip()})
n = int(n_str)
if n < 2:
raise ValueError(f"loop count xN deve avere N>=2 (segmento ({indices}) ha x{n}). "
f"Usa loop_style='' per nessun looping.")
if not indices:
raise ValueError(f"loop segment vuoto in {s!r}")
if indices != list(range(min(indices), max(indices) + 1)):
raise ValueError(f"segmento {tuple(indices)} NON contiguo: i layer in un loop "
f"devono essere contigui (e.g. (3,4,5), non (3,5)).")
for i in indices:
if not 0 <= i < n_layers:
raise ValueError(f"layer {i} fuori range [0, {n_layers}) nel segmento {tuple(indices)}")
if i in used_indices:
raise ValueError(f"layer {i} appare in più segmenti — i segmenti devono essere disgiunti")
used_indices.add(i)
segments.append((indices, n))
segments.sort(key=lambda seg: seg[0][0])
for s1, s2 in zip(segments, segments[1:]):
if s1[0][-1] >= s2[0][0]:
raise ValueError(f"segmenti sovrapposti dopo ordinamento: {s1[0]} e {s2[0]}")
return segments
class LoopGate(nn.Module):
"""Scalar gate init=0 (additive): h_main + α·(h_post − h_main). α=0 → identity."""
def __init__(self):
super().__init__()
self.alpha = nn.Parameter(torch.zeros(()))
def forward(self, h_main: Tensor, h_post: Tensor) -> Tensor:
return h_main + self.alpha * (h_post - h_main)
# =============================================================================
# MODEL — wrapper Vathos + layer looping (CE only, no KD)
# =============================================================================
class PiCOFormerLM(nn.Module):
def __init__(self, hp: HP, vocab_size: int, loop_style: str = ""):
super().__init__()
self.vocab_size = vocab_size
self.logit_softcap = hp.logit_softcap
self.tied_embeddings = hp.tied_embeddings
self.loop_style = loop_style
self.loop_segments = parse_loop_style(loop_style, hp.n_layers) if loop_style else []
self.loop_gates = nn.ModuleList([
nn.ModuleList([LoopGate() for _ in range(n - 1)])
for indices, n in self.loop_segments
])
self._seg_starts = {indices[0]: seg_id for seg_id, (indices, _) in enumerate(self.loop_segments)}
self.looped_mode = False
rope = RoPE(dim=hp.head_dim, max_len=8192, base=hp.rope_base)
if hp.use_gated_attn:
attn_builder = Builder(MultiheadGatedAttentionMixer,
n_heads=hp.n_heads, causal=True, dropout=0.0,
qk_norm=hp.qk_norm, pos_emb=rope, gate_input_dim=hp.gate_input_dim)
else:
attn_builder = Builder(MultiheadAttentionMixer,
n_heads=hp.n_heads, causal=True, dropout=0.0,
qk_norm=hp.qk_norm, pos_emb=rope)
if hp.mlp_kind == "udlp":
chan_builder = Builder(VariableUDLP,
d_output=hp.d_model, M=hp.d_ff, activation=LeakyReLU2)
elif hp.mlp_kind == "gated_udlp":
chan_builder = Builder(VariableGatedUDLP,
d_output=hp.d_model, M=hp.d_ff, activation=LeakyReLU2,
gate_input_dim=hp.gate_input_dim)
elif hp.mlp_kind == "leaky_reglu2":
chan_builder = Builder(LeakyReGLU2_FFN,
d_output=hp.d_model, M=hp.d_ff, activation=LeakyReLU2)
else:
raise ValueError(f"unknown mlp_kind={hp.mlp_kind!r}; usa udlp|gated_udlp|leaky_reglu2")
smear = SmearGate(hp.d_model, gate_input_dim=hp.gate_input_dim) if hp.use_smear_gate else None
smear_lookback = 1 if hp.use_smear_gate else 0
self.backbone = VathosPiCOFormer(
vocab_size=vocab_size, d_model=hp.d_model, n_layers=hp.n_layers,
spatials=attn_builder, channel=chan_builder, norm=VRMSNorm,
ve_groups=None, smear_gate=smear, smear_gate_lookback=smear_lookback,
logit_softcap=hp.logit_softcap, tied_embeddings=hp.tied_embeddings,
)
@property
def classifier_weight(self):
if self.tied_embeddings:
return self.backbone.embedder.embedding.weight
return self.backbone.unembedder.linear.weight
def _hidden(self, input_ids: Tensor) -> Tensor:
if self.looped_mode and self.loop_segments:
return self._hidden_looped(input_ids)
return self._hidden_normal(input_ids)
def _hidden_normal(self, input_ids: Tensor) -> Tensor:
bb = self.backbone
x0 = bb.embedder(input_ids)
if getattr(bb, "smear_gate", None) is not None:
x0 = bb.smear_gate(x0)
ves = bb._compute_ves(input_ids) if hasattr(bb, "_compute_ves") else [None] * len(bb.blocks)
h = x0
for i, block in enumerate(bb.blocks):
h = block(h, x0, ve=ves[i] if ves is not None else None)
h = bb.final_norm(h)
return h
def _hidden_looped(self, input_ids: Tensor) -> Tensor:
bb = self.backbone
x0 = bb.embedder(input_ids)
if getattr(bb, "smear_gate", None) is not None:
x0 = bb.smear_gate(x0)
ves = bb._compute_ves(input_ids) if hasattr(bb, "_compute_ves") else [None] * len(bb.blocks)
h = x0
n_layers = len(bb.blocks)
i = 0
while i < n_layers:
seg_id = self._seg_starts.get(i)
if seg_id is None:
h = bb.blocks[i](h, x0, ve=ves[i] if ves is not None else None)
i += 1
continue
indices, n_total = self.loop_segments[seg_id]
for li in indices:
h = bb.blocks[li](h, x0, ve=ves[li] if ves is not None else None)
h_main = h
for k in range(n_total - 1):
h_post = h_main
for li in indices:
h_post = bb.blocks[li](h_post, x0, ve=ves[li] if ves is not None else None)
h_main = self.loop_gates[seg_id][k](h_main, h_post)
h = h_main
i = indices[-1] + 1
h = bb.final_norm(h)
return h
def forward(self, input_ids: Tensor, targets: Tensor) -> Tensor:
"""CE only (cce fused, no logits materialized)."""
h = self._hidden(input_ids)
W = self.classifier_weight
# cce vuole input bf16/fp16; con master fp32 + autocast h/W possono essere fp32 → cast esplicito
# (input cce identici a prima; il grad risale al master fp32 attraverso il cast).
return linear_cross_entropy(
h.reshape(-1, h.size(-1)).bfloat16(), W.bfloat16(), targets.reshape(-1),
reduction="mean",
softcap=self.logit_softcap if self.logit_softcap > 0 else None,
)
# =============================================================================
# MAIN
# =============================================================================
def main():
hp = HP()
# Dispatch orthogonalizer (Muon.step chiama il nome globale `zeropower_via_newtonschulz5`,
# quindi se l'utente sceglie polar_express ri-bindo quel nome alla nuova fn).
global zeropower_via_newtonschulz5
if hp.orthogonalizer == "ns5":
zeropower_via_newtonschulz5 = torch.compile(zeropower_via_newtonschulz5)
elif hp.orthogonalizer == "polar_express":
# NotImplementedError lo lancia al primo call dentro Muon.step.
zeropower_via_newtonschulz5 = torch.compile(zeropower_via_polar_express)
else:
raise ValueError(f"unknown orthogonalizer={hp.orthogonalizer!r}; usa ns5|polar_express")
distributed = "RANK" in os.environ and "WORLD_SIZE" in os.environ
rank = int(os.environ.get("RANK", "0"))
world_size = int(os.environ.get("WORLD_SIZE", "1"))
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
device = torch.device("cuda", local_rank)
torch.cuda.set_device(device)
if distributed:
dist.init_process_group(backend="nccl", device_id=device)
dist.barrier()
is_main = (rank == 0)
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
torch.backends.cuda.enable_flash_sdp(True)
torch.backends.cuda.enable_mem_efficient_sdp(False)
torch.backends.cuda.enable_math_sdp(False)
torch.backends.cuda.enable_cudnn_sdp(False)
import torch._inductor.config as _ind
_ind.coordinate_descent_tuning = True
_ind.fx_graph_cache = True
_ind.triton.cudagraphs = False
torch.manual_seed(hp.seed)
torch.cuda.manual_seed_all(hp.seed)
np.random.seed(hp.seed)
def log(msg):
if is_main:
print(msg, flush=True)
log(f"world_size={world_size} rank={rank} device={device}")
log(f"corpus_root={hp.corpus_root}")
log(f"mixture={hp.mixture}")
# ---- data: verifica vocab consistency su tutti gli shard ----
vocab_size = None
eos_id = None
for name in hp.mixture:
idx = ShardReader(str(Path(hp.corpus_root) / name)).index
if vocab_size is None:
vocab_size, eos_id = idx.vocab_size, idx.eos_id
assert idx.vocab_size == vocab_size, (
f"vocab mismatch: {name} ha {idx.vocab_size}, atteso {vocab_size}")
assert idx.eos_id == eos_id, f"eos mismatch su {name}"
log(f"vocab={vocab_size} eos={eos_id} (consistente su {len(hp.mixture)} shard)")
assert vocab_size > eos_id
train_loader = MultiShardMixture(hp.mixture, hp.corpus_root, rank, world_size,
device, hp.seq_len, hp.seed)
if is_main:
for n, s in train_loader.streams.items():
log(f" [stream] {n:<20} chunks={s.n_chunks:,} (per-rank {len(s.my):,}) w={hp.mixture[n]:.2f}")
# ---- student model ----
# FIX (audit bf16): pesi in fp32 = MASTER weights. Il forward gira comunque in bf16 via autocast.
# Senza master fp32, al floor del warmdown gli update (~lr·1/√fan) cadono sotto la ULP bf16 e
# p.add_ diventa un no-op → il warmdown non abbassa più la loss. fp32 master lo risolve.
model = PiCOFormerLM(hp, vocab_size=vocab_size, loop_style=hp.loop_style).to(device)
n_params = sum(p.numel() for p in model.parameters())
log(f"model_params={n_params/1e6:.2f}M shape={hp.n_layers}L×{hp.d_model}×{hp.n_heads}h×{hp.d_ff}ff dtype=fp32-master/bf16-autocast")
log(f"knobs: mlp_kind={hp.mlp_kind} gate_input_dim={hp.gate_input_dim} orthogonalizer={hp.orthogonalizer} muon_steps={hp.muon_backend_steps} mixture_preset={os.environ.get('MIXTURE_PRESET','default')}")
if hp.loop_style:
loop_summary = ", ".join(f"({','.join(map(str, idx))})x{n}" for idx, n in model.loop_segments)
n_extra_passes = sum(len(idx) * (n - 1) for idx, n in model.loop_segments)
log(f"[loop] style={hp.loop_style!r} segments={loop_summary}")
log(f"[loop] gates={sum(len(g) for g in model.loop_gates)} (init α=0) "
f"extra_passes_per_fwd={n_extra_passes} activation_frac={hp.loop_activation_frac}")
compiled = torch.compile(model, dynamic=False, fullgraph=False, mode="max-autotune-no-cudagraphs")
ddp_kwargs = dict(device_ids=[local_rank], broadcast_buffers=False, gradient_as_bucket_view=True)
if hp.loop_style:
ddp_kwargs["find_unused_parameters"] = True # loop_gates senza grad in NORMAL fwd
model_for_train = DDP(compiled, **ddp_kwargs) if distributed else compiled
model.train()
# ---- optimizer ----
block_params = list(model.backbone.blocks.named_parameters())
matrix_params = [p for n, p in block_params if p.ndim == 2]
scalar_params = [p for n, p in block_params if p.ndim < 2]
for n, p in model.backbone.final_norm.named_parameters():
scalar_params.append(p)
for n, p in model.loop_gates.named_parameters():
scalar_params.append(p)
# FIX (audit): smear_gate è attributo top-level del backbone (NON dentro blocks) → senza questo
# i suoi param cadono fuori da ogni optimizer e restano congelati a init (smear_lambda=0 → gate
# no-op per tutto il run). Routing per ndim, coerente col resto (gate.weight 2D→Muon, lambda→scalar).
if getattr(model.backbone, "smear_gate", None) is not None:
for n, p in model.backbone.smear_gate.named_parameters():
(matrix_params if p.ndim == 2 else scalar_params).append(p)
embed_param = model.backbone.embedder.embedding.weight
opt_muon = Muon(matrix_params, lr=hp.matrix_lr, momentum=hp.muon_momentum,
backend_steps=hp.muon_backend_steps)
opt_embed = torch.optim.Adam(
[{"params": [embed_param], "lr": hp.tied_embed_lr, "base_lr": hp.tied_embed_lr}],
betas=(hp.beta1, hp.beta2), eps=hp.adam_eps, fused=True,
)
opt_scalar = torch.optim.Adam(
[{"params": scalar_params, "lr": hp.scalar_lr, "base_lr": hp.scalar_lr}],
betas=(hp.beta1, hp.beta2), eps=hp.adam_eps, fused=True,
)
optimizers = [opt_muon, opt_embed, opt_scalar]
for g in opt_muon.param_groups:
g["base_lr"] = hp.matrix_lr
log(f"opt: Muon={sum(p.numel() for p in matrix_params)/1e6:.2f}M "
f"Adam embed={embed_param.numel()/1e6:.2f}M scalars={sum(p.numel() for p in scalar_params)/1e6:.2f}M")
local_bs = hp.bs_per_dev
log(f"bs/rank={local_bs} global_bs={local_bs*world_size} tok/step={local_bs*world_size*hp.seq_len:,}")
log(f"iterations={hp.iterations} warmup={hp.warmup_steps} warmdown={hp.warmdown_iters} "
f"→ target_tok={hp.iterations*local_bs*world_size*hp.seq_len/1e9:.2f}B")
def lr_scale(step):
if step < hp.warmup_steps:
return step / max(hp.warmup_steps, 1)
decay_start = hp.iterations - hp.warmdown_iters
if step < decay_start:
return 1.0
progress = (step - decay_start) / max(hp.warmdown_iters, 1)
return max(1.0 - progress * (1.0 - hp.lr_min_scale), hp.lr_min_scale)
Path(hp.ckpt_dir).mkdir(parents=True, exist_ok=True)
# ---- loss-trace CSV (per-step; UN sync ogni train_log_every) ----
csv_f = csv_w = None
if is_main:
Path(hp.loss_trace_csv).parent.mkdir(parents=True, exist_ok=True)
csv_f = open(hp.loss_trace_csv, "w", newline="")
csv_w = csv.writer(csv_f)
csv_w.writerow(["step", "loss", "gnorm", "lr", "math_frac"])
log(f"loss_trace_csv={hp.loss_trace_csv}")
loss_buf, gn_buf, lr_buf, step_buf, mf_buf = [], [], [], [], []
# ---- Layer looping: pre-compile both variants ----
loop_activation_step = int(hp.iterations * hp.loop_activation_frac) if hp.loop_style else None
loop_active = False
if hp.loop_style and hp.loop_pre_warm:
log(f"[loop] pre-compiling NORMAL + LOOPED graphs (one-shot, può richiedere 5-15 min)")
x_warm = torch.zeros((local_bs, hp.seq_len), dtype=torch.long, device=device)
y_warm = torch.zeros_like(x_warm)
with torch.no_grad():
model.looped_mode = False
with torch.autocast(device_type="cuda", dtype=torch.bfloat16, enabled=True):
_ = compiled(x_warm, y_warm)
torch.cuda.synchronize()
log(f"[loop] NORMAL fwd graph compiled")
model.looped_mode = True
with torch.autocast(device_type="cuda", dtype=torch.bfloat16, enabled=True):
_ = compiled(x_warm, y_warm)
torch.cuda.synchronize()
log(f"[loop] LOOPED fwd graph compiled")
model.looped_mode = False
del x_warm, y_warm
torch.cuda.synchronize()
t0 = time.perf_counter()
total_tok_seen = 0
nan_abort = False
for step in range(1, hp.iterations + 1):
if loop_activation_step is not None and not loop_active and step >= loop_activation_step:
log(f"[loop] ACTIVATING looping at step {step} (gates α=0, identity passthrough)")
model.looped_mode = True
loop_active = True
for opt in optimizers:
opt.zero_grad(set_to_none=True)
x, y, ds_ids = train_loader.next_batch(local_bs)
total_tok_seen += local_bs * world_size * hp.seq_len
with torch.autocast(device_type="cuda", dtype=torch.bfloat16, enabled=True):
loss = model_for_train(x, y)
loss.backward()
scale = lr_scale(step)
for opt in optimizers:
for g in opt.param_groups:
g["lr"] = g["base_lr"] * scale
frac = min(step / hp.muon_momentum_warmup_steps, 1.0) if hp.muon_momentum_warmup_steps > 0 else 1.0
mom = (1 - frac) * hp.muon_momentum_warmup_start + frac * hp.muon_momentum
for g in opt_muon.param_groups:
g["momentum"] = mom
gn = torch.nn.utils.clip_grad_norm_(model.parameters(), hp.grad_clip)
for opt in optimizers:
opt.step()
# per-step buffers (NO sync qui — un solo sync ogni train_log_every al flush)
loss_buf.append(loss.detach())
gn_buf.append(gn.detach())
lr_buf.append(opt_muon.param_groups[0]["lr"])
step_buf.append(step)
mf_buf.append(sum(d not in hp.nl_datasets for d in ds_ids) / len(ds_ids))
if step % hp.train_log_every == 0 or step == 1 or step == hp.iterations:
torch.cuda.synchronize()
ls = torch.stack(loss_buf).float().cpu().tolist()
gs = torch.stack(gn_buf).float().cpu().tolist()
if not all(np.isfinite(ls)):
log(f"!! NaN/Inf loss intorno a step {step}, abort")
nan_abort = True
if is_main:
for s, l, g, lr_, mf in zip(step_buf, ls, gs, lr_buf, mf_buf):
csv_w.writerow([s, f"{l:.5f}", f"{g:.3f}", f"{lr_:.3e}", f"{mf:.2f}"])
csv_f.flush()
elapsed = time.perf_counter() - t0
tps = total_tok_seen / max(elapsed, 1e-9)
loop_tag = ""
if loop_active:
alphas = [g.alpha.detach().abs().item()
for seg_gates in model.loop_gates for g in seg_gates]
loop_tag = f" [LOOP |α|avg={sum(alphas)/max(len(alphas),1):.3f}]"
log(f"step {step:>5d}/{hp.iterations}{loop_tag} loss={ls[-1]:.4f} "
f"gnorm={gs[-1]:.3f} lr={lr_buf[-1]:.2e} mathfrac={mf_buf[-1]:.2f} "
f"tok={total_tok_seen/1e9:.2f}B tok/s={tps:>10,.0f} elapsed={elapsed:.1f}s")
loss_buf.clear(); gn_buf.clear(); lr_buf.clear(); step_buf.clear(); mf_buf.clear()
if nan_abort:
break
if (hp.save_interval > 0 and is_main and step % hp.save_interval == 0
and step != hp.iterations):
ckpt_path = Path(hp.ckpt_dir) / f"step_{step:06d}.pt"
torch.save({"step": step, "model": model.state_dict(), "vocab_size": vocab_size}, ckpt_path)
log(f" [ckpt] saved {ckpt_path.name}")
intermediates = [p for p in sorted(Path(hp.ckpt_dir).glob("step_*.pt")) if "_final" not in p.name]
for old in intermediates[:-hp.keep_n_recent]:
old.unlink()
if hp.save_final and is_main and not nan_abort:
ckpt_path = Path(hp.ckpt_dir) / f"step_{hp.iterations:06d}_final.pt"
torch.save({"step": hp.iterations, "model": model.state_dict(), "vocab_size": vocab_size}, ckpt_path)
log(f"saved {ckpt_path}")
if is_main and csv_f is not None:
csv_f.close()
if distributed:
dist.destroy_process_group()
if __name__ == "__main__":
main()
|