"""Muon optimizer (Newton-Schulz orthogonalized momentum) + hybrid helpers. Convention: Muon on 2D hidden matrices, AdamW on everything else (emb/pos/LN/bias).""" import torch def newton_schulz(G, steps=5, eps=1e-7): """approximate polar factor of G via the quintic NS iteration (Keller Jordan coefficients).""" a, b, c = 3.4445, -4.7750, 2.0315 X = G / (G.norm() + eps) transposed = X.size(0) > X.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 class Muon(torch.optim.Optimizer): def __init__(self, params, lr=0.02, momentum=0.95, ns_steps=5, nesterov=True): super().__init__(params, dict(lr=lr, momentum=momentum, ns_steps=ns_steps, nesterov=nesterov)) @torch.no_grad() def step(self, closure=None): for group in self.param_groups: for p in group['params']: if p.grad is None: continue g = p.grad st = self.state[p] if 'mom' not in st: st['mom'] = torch.zeros_like(g) buf = st['mom'] buf.mul_(group['momentum']).add_(g) u = g.add(buf, alpha=group['momentum']) if group['nesterov'] else buf if u.ndim == 2: u = newton_schulz(u, group['ns_steps']) u = u * max(1.0, u.size(0) / u.size(1)) ** 0.5 # rms-matched scaling p.add_(u, alpha=-group['lr']) class MultiOpt: """duck-typed bundle of optimizers (step/zero_grad/state_dict API-compatible).""" def __init__(self, opts): self.optimizers = opts def step(self): for o in self.optimizers: o.step() def zero_grad(self, set_to_none=True): for o in self.optimizers: o.zero_grad(set_to_none=set_to_none) def state_dict(self): return [o.state_dict() for o in self.optimizers] def load_state_dict(self, sds): for o, sd in zip(self.optimizers, sds): o.load_state_dict(sd) class MultiSched: def __init__(self, scheds): self.scheds = scheds def step(self): for s in self.scheds: s.step() def build_hybrid(blocks, other_params, lr_adamw, lr_muon, warmup, total_steps=0, lr_min_ratio=0.1, muon_mom=0.95, adam_b1=0.9, head_param=None, head_lr_mult=1.0): """Muon(2D block matrices) + AdamW(everything else). Scheds: linear warmup, then cosine decay to lr_min_ratio*peak if total_steps>0 (long runs), else constant after warmup (legacy). muon_mom/adam_b1: momentum knobs (late-SNR noise-averaging arms, 2026-07-13).""" import math as _m mats = [p for p in blocks.parameters() if p.ndim == 2] mat_ids = {id(p) for p in mats} rest = [p for p in other_params if id(p) not in mat_ids] om = Muon(mats, lr=lr_muon, momentum=muon_mom) if head_param is not None and head_lr_mult != 1.0: hid = id(head_param) groups = [{'params': [p for p in rest if id(p) != hid], 'lr': lr_adamw}, {'params': [p for p in rest if id(p) == hid], 'lr': lr_adamw * head_lr_mult}] oa = torch.optim.AdamW(groups, lr=lr_adamw, weight_decay=1e-4, betas=(adam_b1, 0.999)) else: oa = torch.optim.AdamW(rest, lr=lr_adamw, weight_decay=1e-4, betas=(adam_b1, 0.999)) if total_steps > 0: def fn(s): if s < warmup: return (s + 1) / max(warmup, 1) p = min(1.0, (s - warmup) / max(1, total_steps - warmup)) return lr_min_ratio + 0.5 * (1 - lr_min_ratio) * (1 + _m.cos(_m.pi * p)) else: fn = lambda s: min(1.0, (s + 1) / max(warmup, 1)) scheds = [torch.optim.lr_scheduler.LambdaLR(om, fn), torch.optim.lr_scheduler.LambdaLR(oa, fn)] return MultiOpt([om, oa]), MultiSched(scheds)