GUIDO_test_200M / modeling_guido.py
Paerle's picture
Initial upload: GUIDO-small 200M bugfix ckpt + Vathos + modeling + README
d8c733f verified
Raw
History Blame Contribute Delete
35.8 kB
"""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}{hp.d_model}×{hp.n_heads}{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()