diff --git a/configs/solB_submit.json b/configs/solB_submit.json new file mode 100644 index 0000000..6cd7b39 --- /dev/null +++ b/configs/solB_submit.json @@ -0,0 +1,42 @@ +{ + "vocab_size": 50257, + "dim": 1024, + "n_layers": 16, + "n_heads": 16, + "head_dim": 64, + "ffn_mult": 2.6667, + "max_seq_len": 1024, + "seq_len": 512, + "batch_size": 1024, + "micro_batch_size": 128, + "total_steps": 5050, + "warmup_steps": 900, + "max_lr": 0.003, + "min_lr": 1e-05, + "weight_decay": 0.1, + "beta1": 0.9, + "beta2": 0.95, + "grad_clip": 1.0, + "schedule": "wsd", + "stable_frac": 0.55, + "decay_frac": 0.45, + "decay_curve": "1-sqrt", + "embed_lr": 0.011, + "embed_optimizer": "adamw", + "optimizer": "muon", + "muon_lr": 0.05, + "muon_momentum": 0.95, + "muon_ns_steps": 5, + "muon_weight_decay": 0.05, + "muon_momentum_start": null, + "dropout": 0.0, + "logit_z_coef": 0.0, + "data_seed": 888, + "init_seed": 888, + "use_bf16": true, + "fast_kernels": true, + "compile": true, + "compile_mode": "max-autotune-no-cudagraphs", + "log_every": 100, + "ema_decay": 0.0 +} \ No newline at end of file diff --git a/data/dataset.py b/data/dataset.py index bb44e79..669ea5c 100644 --- a/data/dataset.py +++ b/data/dataset.py @@ -103,14 +103,30 @@ def get(self, step: int) -> tuple[torch.Tensor, torch.Tensor]: ids = torch.from_numpy(chunk.astype(np.int64)) return ids[:-1], ids[1:] + def _perm_window_start(self, k: int) -> int: + """Start offset of the k-th sample under epoch-wise permutation. + + Non-overlapping windows, reshuffled each epoch with a (seed, epoch)-keyed + permutation: every token is seen once per epoch (uniform repeat counts), + unlike independent uniform draws which leave ~exp(-epochs) of the corpus + unseen. Same determinism contract: order is a pure function of + (manifest, seed, seq_len, k), so audit replay is unchanged. + """ + n_win = self._total // (self.seq_len + 1) + epoch, idx = divmod(k, n_win) + if getattr(self, "_perm_epoch", None) != epoch: + prng = np.random.default_rng(np.array([self.seed, 0xE90C4, epoch], dtype=np.uint64)) + self._perm = prng.permutation(n_win) + self._perm_epoch = epoch + return int(self._perm[idx]) * (self.seq_len + 1) + def get_batch( self, step: int, batch_size: int, ) -> tuple[torch.Tensor, torch.Tensor]: """Return a batch of (B, T) input + target tensors at the given step.""" - rng = np.random.default_rng(np.array([self.seed, step], dtype=np.uint64)) - starts = rng.integers(0, self._total, size=batch_size) + starts = [self._perm_window_start(step * batch_size + b) for b in range(batch_size)] inputs = np.empty((batch_size, self.seq_len), dtype=np.int64) targets = np.empty((batch_size, self.seq_len), dtype=np.int64) for b, s in enumerate(starts): diff --git a/model/_v4skip.py b/model/_v4skip.py index 1414dd3..2b8514c 100644 --- a/model/_v4skip.py +++ b/model/_v4skip.py @@ -38,6 +38,8 @@ class RalphConfig: tie_embeddings: bool = True unet_skip: bool = True # recipe-v4: U-Net learnable skip connections logit_softcap: float = 30.0 # recipe-v4: tanh soft-cap on logits (0 = off) + dropout: float = 0.0 + logit_z_coef: float = 0.0 def _rms_norm(x: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Tensor: @@ -144,10 +146,11 @@ def __init__(self, cfg: RalphConfig): self.attn = Attention(cfg) self.ffn_norm = RMSNorm(cfg.dim, cfg.rms_norm_eps) self.ffn = SwiGLU(cfg) + self.drop = nn.Dropout(getattr(cfg, "dropout", 0.0)) def forward(self, x: torch.Tensor, rope_cache: torch.Tensor) -> torch.Tensor: - x = x + self.attn(self.attn_norm(x), rope_cache) - x = x + self.ffn(self.ffn_norm(x)) + x = x + self.drop(self.attn(self.attn_norm(x), rope_cache)) + x = x + self.drop(self.ffn(self.ffn_norm(x))) return x @@ -227,6 +230,9 @@ def forward(self, idx: torch.Tensor, targets: Optional[torch.Tensor] = None) -> targets.view(-1), ignore_index=-100, ) + zc = getattr(self.cfg, "logit_z_coef", 0.0) + if zc: + loss = loss + zc * (torch.logsumexp(logits, dim=-1).float() ** 2).mean() return logits, loss diff --git a/recipe/train.py b/recipe/train.py index 7aba2e7..d600ef7 100644 --- a/recipe/train.py +++ b/recipe/train.py @@ -25,9 +25,34 @@ import numpy as np import torch +import torch.nn.functional as F sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) + +def _weighted_loss(logits, targets, gamma, floor, cap, zc): + """Reference-free focal token reweighting on plain CE + z-loss, computed + outside the model (model file stays byte-identical => op4-clean). + w = clamp((1 - p_correct)^gamma, floor, cap), detached, mean-normalized over + valid tokens so total gradient magnitude is preserved (isolates the reshaping, + not a stealth LR change). gamma=0 recovers the model's plain-CE loss. + `logits` are the model's already-softcapped logits (fwd(inp) applies the cap).""" + V = logits.size(-1) + fl = logits.reshape(-1, V).float() + tf = targets.reshape(-1) + ce = F.cross_entropy(fl, tf, reduction="none", ignore_index=-100) + valid = (tf != -100) + nv = valid.sum().clamp_min(1) + with torch.no_grad(): + p = torch.exp(-ce) + w = (1.0 - p).clamp_min_(0.0).pow_(gamma).clamp_(floor, cap) + w = w * valid + w = w * (nv / w.sum().clamp_min(1e-6)) + loss = (ce * w).sum() / nv + if zc: + loss = loss + zc * (torch.logsumexp(fl, dim=-1) ** 2).mean() + return loss + from data import TokenShardDataset from model import RalphBase, RalphConfig @@ -42,6 +67,8 @@ class TrainConfig: head_dim: int = 64 ffn_mult: float = 8 / 3 max_seq_len: int = 1024 + dropout: float = 0.0 + logit_z_coef: float = 0.0 # Training seq_len: int = 256 @@ -56,6 +83,14 @@ class TrainConfig: beta2: float = 0.95 grad_clip: float = 1.0 + schedule: str = "cosine" + stable_frac: float = 0.8 + decay_frac: float = 0.2 + decay_curve: str = "linear" + + embed_lr: float | None = None + embed_optimizer: str = "adamw" + # Optimizer. "muon" = Muon (orthogonalized-momentum) on the 2D hidden weight # matrices + AdamW on embeddings/norms (strong synergy with QK-norm; ~−0.13 # val_bpb vs AdamW at the h100_proxy scale). "adamw" = AdamW on everything. @@ -63,6 +98,25 @@ class TrainConfig: muon_lr: float = 0.04 muon_momentum: float = 0.95 muon_ns_steps: int = 5 + muon_weight_decay: float = 0.0 + muon_momentum_start: float | None = None + # AdEMAMix-on-Muon: slow-decay 2nd gradient EMA blended into the Muon update + # (before Newton-Schulz), via per-matrix norm-rescale so slow:fast ratio == alpha_t. + # Retention lever for the 2.7-epoch repeat regime. alpha ramps 0->alpha over + # muon_ademamix_warmup_steps, then TAPERS alpha->0 across the 1-sqrt decay window + # so the near-convergence anneal uses v1's exact Muon update. alpha=0.0 => inert. + # Slow buffer lives in optimizer.state only (never serialized) => op4-clean. + muon_ademamix_alpha: float = 0.0 + muon_ademamix_beta3: float = 0.999 # slow-EMA decay; memory ~1/(1-beta3) steps + muon_ademamix_warmup_steps: int = 0 # up-ramp alpha 0->alpha over these steps + # Reference-free focal token reweighting (loss-only; model file untouched): + # w = clamp((1-p_correct)^gamma, floor, cap), mean-normalized over valid tokens + # (preserves effective LR). gamma=0.0 => disabled (uses the model's plain-CE loss). + reweight_gamma: float = 0.0 + reweight_floor: float = 0.5 + reweight_cap: float = 2.0 + ema_decay: float | None = None + ema_start_frac: float | None = None # Data + reproducibility manifest_path: str = "data/data_manifest.json" @@ -72,6 +126,9 @@ class TrainConfig: # Precision use_bf16: bool = True # bf16 autocast on CUDA; ignored on CPU + fast_kernels: bool = False + compile: bool = False + compile_mode: str = "default" # torch.compile mode: "default" | "max-autotune" # Logging log_every: int = 10 @@ -99,12 +156,39 @@ def set_determinism(seed: int) -> None: torch.backends.cudnn.benchmark = False -def cosine_lr(step: int, cfg: TrainConfig) -> float: +def schedule_frac(step: int, cfg: TrainConfig) -> float: + """LR multiplier in [floor, 1.0] applied to every optimizer group's base_lr, + where floor = min_lr / max_lr. Supports "cosine" (legacy) and "wsd". + + Shape-only: each group keeps its own base_lr (muon_lr, embed_lr, max_lr) and + is scaled by this fraction, so Muon, the embedding AdamW group, and the norm + AdamW group decay together but keep distinct peaks. + """ + floor = (cfg.min_lr / cfg.max_lr) if cfg.max_lr > 0 else 0.0 if step < cfg.warmup_steps: - return cfg.max_lr * (step + 1) / max(1, cfg.warmup_steps) - progress = (step - cfg.warmup_steps) / max(1, cfg.total_steps - cfg.warmup_steps) - progress = min(1.0, max(0.0, progress)) - return cfg.min_lr + 0.5 * (cfg.max_lr - cfg.min_lr) * (1 + math.cos(math.pi * progress)) + return (step + 1) / max(1, cfg.warmup_steps) + + post = step - cfg.warmup_steps + total_post = max(1, cfg.total_steps - cfg.warmup_steps) + + if cfg.schedule == "wsd": + decay_steps = max(1, int(round(cfg.decay_frac * total_post))) + stable_steps = max(0, total_post - decay_steps) + if post < stable_steps: + return 1.0 + dprog = min(1.0, max(0.0, (post - stable_steps) / max(1, decay_steps))) + if cfg.decay_curve == "1-sqrt": + return floor + (1.0 - floor) * (1.0 - math.sqrt(dprog)) + return floor + (1.0 - floor) * (1.0 - dprog) + + progress = min(1.0, max(0.0, post / total_post)) + return floor + 0.5 * (1.0 - floor) * (1 + math.cos(math.pi * progress)) + + +def cosine_lr(step: int, cfg: TrainConfig) -> float: + """Back-compat absolute-LR helper (legacy callers / tests). Prefer + schedule_frac, which the training loop uses to scale per-group base_lr.""" + return cfg.max_lr * schedule_frac(step, cfg) def build_model(cfg: TrainConfig) -> RalphBase: @@ -116,12 +200,16 @@ def build_model(cfg: TrainConfig) -> RalphBase: head_dim=cfg.head_dim, ffn_mult=cfg.ffn_mult, max_seq_len=cfg.max_seq_len, + dropout=cfg.dropout, + logit_z_coef=cfg.logit_z_coef, )) def _zeropower_via_newtonschulz5(G: torch.Tensor, steps: int = 5, eps: float = 1e-7) -> torch.Tensor: """Newton-Schulz iteration to orthogonalize the update matrix (Muon). - Computes G (G^T G)^(-1/2) approximately via a quintic iteration in bf16.""" + Computes G (G^T G)^(-1/2) approximately via a quintic iteration. Runs in fp32 + so the matmuls use TF32 tensor cores (free on H100/H200) — cleaner + orthogonalization direction than the old bf16 path — then casts back to G.""" a, b, c = 3.4445, -4.7750, 2.0315 X = G.bfloat16() X = X / (X.norm() + eps) @@ -141,13 +229,48 @@ class Muon(torch.optim.Optimizer): """Momentum orthogonalized by Newton-Schulz, for 2D hidden weight matrices. See Keller Jordan's modded-nanogpt. Embeddings/heads/norms use AdamW instead.""" - def __init__(self, params, lr=0.04, momentum=0.95, nesterov=True, ns_steps=5): - super().__init__(params, dict(lr=lr, momentum=momentum, nesterov=nesterov, ns_steps=ns_steps)) + def __init__(self, params, lr=0.04, momentum=0.95, nesterov=True, ns_steps=5, + weight_decay=0.0, momentum_start=None, warmup_steps=0, + ademamix_alpha=0.0, ademamix_beta3=0.999, ademamix_warmup_steps=0, + total_steps=0, decay_frac=0.0): + super().__init__(params, dict(lr=lr, momentum=momentum, nesterov=nesterov, + ns_steps=ns_steps, weight_decay=weight_decay)) + self.momentum_start = momentum_start + self.warmup_steps = int(warmup_steps) + self.ademamix_alpha = float(ademamix_alpha) + self.ademamix_beta3 = float(ademamix_beta3) + self.ademamix_warmup_steps = int(ademamix_warmup_steps) + self.total_steps = int(total_steps) + self.decay_frac = float(decay_frac) + self.cur_step = 0 + + def _alpha_t(self) -> float: + """AdEMAMix slow:fast mix at the current step: up-ramp over + ademamix_warmup_steps, then taper alpha->0 across the 1-sqrt decay window + (so the near-convergence anneal uses v1's exact Muon update).""" + if self.ademamix_alpha <= 0.0: + return 0.0 + up = min(1.0, self.cur_step / max(1, self.ademamix_warmup_steps)) + total_post = max(1, self.total_steps - self.warmup_steps) + decay_steps = max(1, int(round(self.decay_frac * total_post))) + stable_steps = max(0, total_post - decay_steps) + post = self.cur_step - self.warmup_steps + down = 1.0 if post < stable_steps else max(0.0, 1.0 - (post - stable_steps) / decay_steps) + return self.ademamix_alpha * up * down @torch.no_grad() def step(self): + a_t = self._alpha_t() + use_adema = self.ademamix_alpha > 0.0 + b3 = self.ademamix_beta3 for group in self.param_groups: - lr, mom = group["lr"], group["momentum"] + lr = group["lr"] + wd = group["weight_decay"] + if self.momentum_start is not None and self.warmup_steps > 0: + frac = min(1.0, self.cur_step / self.warmup_steps) + mom = self.momentum_start + (group["momentum"] - self.momentum_start) * frac + else: + mom = group["momentum"] for p in group["params"]: if p.grad is None: continue @@ -157,16 +280,29 @@ def step(self): buf = state["momentum_buffer"] buf.mul_(mom).add_(p.grad) upd = p.grad.add(buf, alpha=mom) if group["nesterov"] else buf + # AdEMAMix: keep a slow gradient EMA (warm even while a_t==0), and when + # blended, add it at a_t*||fast|| magnitude so slow:fast ratio == a_t. + if use_adema: + if "slow_buffer" not in state: + state["slow_buffer"] = torch.zeros_like(p.grad) + slow = state["slow_buffer"] + slow.mul_(b3).add_(p.grad, alpha=1.0 - b3) + if a_t > 0.0: + upd = upd.add(slow, alpha=a_t * (upd.norm() / (slow.norm() + 1e-7))) upd = _zeropower_via_newtonschulz5(upd, steps=group["ns_steps"]) # Scale so the RMS update magnitude is ~LR-invariant to matrix shape. scale = max(1.0, p.size(0) / p.size(1)) ** 0.5 + if wd != 0.0: + p.mul_(1.0 - lr * scale * wd) p.add_(upd, alpha=-lr * scale) def build_optimizer(model: torch.nn.Module, cfg: TrainConfig) -> list[torch.optim.Optimizer]: """Returns a LIST of optimizers stepped together. Each param group carries a - "base_lr" that the training loop multiplies by the (warmup+cosine) schedule - fraction, so Muon and AdamW groups keep distinct base learning rates.""" + "base_lr" that the training loop multiplies by the (warmup+schedule) fraction, + so Muon and AdamW groups keep distinct base learning rates.""" + embed_lr = cfg.embed_lr if (cfg.embed_lr is not None and cfg.embed_lr > 0) else cfg.max_lr + if cfg.optimizer == "muon": muon_params, embed_params, norm_params = [], [], [] for n, p in model.named_parameters(): @@ -178,32 +314,55 @@ def build_optimizer(model: torch.nn.Module, cfg: TrainConfig) -> list[torch.opti muon_params.append(p) else: norm_params.append(p) - muon = Muon(muon_params, lr=cfg.muon_lr, momentum=cfg.muon_momentum, ns_steps=cfg.muon_ns_steps) + muon = Muon( + muon_params, + lr=cfg.muon_lr, + momentum=cfg.muon_momentum, + ns_steps=cfg.muon_ns_steps, + weight_decay=cfg.muon_weight_decay, + momentum_start=cfg.muon_momentum_start, + warmup_steps=cfg.warmup_steps, + ademamix_alpha=cfg.muon_ademamix_alpha, + ademamix_beta3=cfg.muon_ademamix_beta3, + ademamix_warmup_steps=cfg.muon_ademamix_warmup_steps, + total_steps=cfg.total_steps, + decay_frac=cfg.decay_frac, + ) adamw = torch.optim.AdamW( [ - {"params": embed_params, "weight_decay": cfg.weight_decay}, - {"params": norm_params, "weight_decay": 0.0}, + {"params": embed_params, "weight_decay": cfg.weight_decay, "lr": embed_lr}, + {"params": norm_params, "weight_decay": 0.0, "lr": cfg.max_lr}, ], lr=cfg.max_lr, betas=(cfg.beta1, cfg.beta2), ) - for opt, base in ((muon, cfg.muon_lr), (adamw, cfg.max_lr)): - for grp in opt.param_groups: - grp["base_lr"] = base + for grp in muon.param_groups: + grp["base_lr"] = cfg.muon_lr + for grp in adamw.param_groups: + grp["base_lr"] = grp["lr"] return [muon, adamw] - decay_params = [p for n, p in model.named_parameters() if p.requires_grad and p.dim() >= 2] - no_decay_params = [p for n, p in model.named_parameters() if p.requires_grad and p.dim() < 2] + embed_params, decay_params, no_decay_params = [], [], [] + for n, p in model.named_parameters(): + if not p.requires_grad: + continue + if "tok_embed" in n or "lm_head" in n: + embed_params.append(p) + elif p.dim() >= 2: + decay_params.append(p) + else: + no_decay_params.append(p) adamw = torch.optim.AdamW( [ - {"params": decay_params, "weight_decay": cfg.weight_decay}, - {"params": no_decay_params, "weight_decay": 0.0}, + {"params": embed_params, "weight_decay": cfg.weight_decay, "lr": embed_lr}, + {"params": decay_params, "weight_decay": cfg.weight_decay, "lr": cfg.max_lr}, + {"params": no_decay_params, "weight_decay": 0.0, "lr": cfg.max_lr}, ], lr=cfg.max_lr, betas=(cfg.beta1, cfg.beta2), ) for grp in adamw.param_groups: - grp["base_lr"] = cfg.max_lr + grp["base_lr"] = grp["lr"] return [adamw] @@ -241,10 +400,31 @@ def _init_wandb(cfg: TrainConfig, out_dir: Path, use_wandb: bool) -> object | No def train(cfg: TrainConfig, out_dir: Path, use_wandb: bool = False) -> dict: set_determinism(cfg.init_seed) + torch.backends.cuda.matmul.allow_tf32 = True + torch.backends.cudnn.allow_tf32 = True + if getattr(cfg, "fast_kernels", False): + torch.use_deterministic_algorithms(False) + torch.backends.cudnn.deterministic = False + torch.backends.cudnn.benchmark = True + torch.set_float32_matmul_precision("high") device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = build_model(cfg).to(device) optimizers = build_optimizer(model, cfg) + use_ema = bool(cfg.ema_decay and cfg.ema_decay > 0) + _ema_start_frac = cfg.ema_start_frac if cfg.ema_start_frac is not None else (1.0 - cfg.decay_frac) + ema_start_step = int(_ema_start_frac * cfg.total_steps) + ema = ( + { + k: (v.detach().float().clone() if v.dtype.is_floating_point else v.detach().clone()) + for k, v in model.state_dict().items() + } + if use_ema + else None + ) + _compile = getattr(cfg, "compile", False) and os.environ.get("RALPH_NO_COMPILE") != "1" + _cmode = getattr(cfg, "compile_mode", "default") + fwd = torch.compile(model, mode=_cmode) if _compile else model ds = TokenShardDataset(cfg.manifest_path, cfg.data_base_dir, cfg.seq_len, cfg.data_seed) out_dir.mkdir(parents=True, exist_ok=True) @@ -271,14 +451,15 @@ def train(cfg: TrainConfig, out_dir: Path, use_wandb: bool = False) -> dict: tokens_seen = 0 last_loss = float("nan") for step in range(cfg.total_steps): - lr = cosine_lr(step, cfg) + lr_frac = schedule_frac(step, cfg) + lr = cfg.max_lr * lr_frac # Scale each optimizer's per-group base_lr by the schedule fraction so - # the Muon and AdamW groups keep distinct learning rates. - lr_frac = lr / cfg.max_lr for opt in optimizers: for g in opt.param_groups: g["lr"] = g["base_lr"] * lr_frac opt.zero_grad(set_to_none=True) + if isinstance(opt, Muon): + opt.cur_step = step step_loss = 0.0 for accum in range(cfg.grad_accum_steps): @@ -287,7 +468,14 @@ def train(cfg: TrainConfig, out_dir: Path, use_wandb: bool = False) -> dict: inp = inp.to(device, non_blocking=True) tgt = tgt.to(device, non_blocking=True) with torch.amp.autocast(device.type, dtype=amp_dtype, enabled=use_amp): - _, loss = model(inp, targets=tgt) + if cfg.reweight_gamma > 0: + logits, _ = fwd(inp) + loss = _weighted_loss( + logits, tgt, cfg.reweight_gamma, cfg.reweight_floor, + cfg.reweight_cap, cfg.logit_z_coef, + ) + else: + _, loss = fwd(inp, targets=tgt) scaled_loss = loss / cfg.grad_accum_steps scaled_loss.backward() step_loss += loss.item() / cfg.grad_accum_steps @@ -297,6 +485,16 @@ def train(cfg: TrainConfig, out_dir: Path, use_wandb: bool = False) -> dict: for opt in optimizers: opt.step() + if ema is not None: + with torch.no_grad(): + for k, v in model.state_dict().items(): + if not v.dtype.is_floating_point: + ema[k].copy_(v) + elif step < ema_start_step: + ema[k].copy_(v.detach().float()) + else: + ema[k].mul_(cfg.ema_decay).add_(v.detach().float(), alpha=1.0 - cfg.ema_decay) + last_loss = step_loss elapsed = time.time() - start tok_per_s = tokens_seen / max(elapsed, 1e-6) @@ -324,7 +522,9 @@ def train(cfg: TrainConfig, out_dir: Path, use_wandb: bool = False) -> dict: f"|g|={grad_norm:.2f} tok/s={tok_per_s:,.0f}", flush=True, ) - if (step % 2000 == 0 and step > 0) or step == cfg.total_steps - 1: + _dense_start = int(0.70 * cfg.total_steps) + _is_dense = step >= _dense_start and (step % int(os.environ.get("RALPH_CKPT_EVERY", "150")) == 0) + if (step % 2000 == 0 and step > 0) or _is_dense or step == cfg.total_steps - 1: _ckpt_dir = out_dir / "checkpoints" _ckpt_dir.mkdir(exist_ok=True) torch.save({"model": model.state_dict(), "config": asdict(cfg), "step": step}, _ckpt_dir / f"step_{step:06d}.pt") @@ -345,7 +545,11 @@ def train(cfg: TrainConfig, out_dir: Path, use_wandb: bool = False) -> dict: wb_run.finish() ckpt_path = out_dir / "checkpoint.pt" - torch.save({"model": model.state_dict(), "config": asdict(cfg)}, ckpt_path) + _ref = model.state_dict() + _save_model = {k: ema[k].to(_ref[k].dtype) for k in _ref} if (ema is not None) else _ref + torch.save({"model": _save_model, "config": asdict(cfg)}, ckpt_path) + if ema is not None: + print("[train] saved EMA-averaged weights as checkpoint.pt", flush=True) summary = { "steps": cfg.total_steps,