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
|
"""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)
|