"""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 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) 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): """Muon(2D block matrices) + AdamW(everything else), with linear-warmup scheds for both.""" 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) oa = torch.optim.AdamW(rest, lr=lr_adamw, weight_decay=1e-4) 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)