diff options
| author | Yuren Hao <yurenh2@illinois.edu> | 2026-08-05 19:14:33 -0500 |
|---|---|---|
| committer | Yuren Hao <yurenh2@illinois.edu> | 2026-08-05 19:14:33 -0500 |
| commit | 2de7789ecf05eba74cc52dbd5f4d8298e4084e50 (patch) | |
| tree | 5546f171017966ea7294c4c09a67e4a09778a891 /ep_run/muon.py | |
| parent | 0df61c1d36faf3e985394c943853b5bab5adaa8c (diff) | |
PSGD-QUAD入列: vendor psgd.py + Optimizer子类适配器(合成损失closure桥接EP预填梯度), 13路全量回归冒烟PASS
修复过程记录: 首版适配器非Optimizer子类被LambdaLR拒; 重写时splice错序造成muon.py重复段
(旧类后定义胜出), 按结构图手术去重并断言全符号在位。三轴定位: GPU配方候选(追踪慢变统计=
R4幸存哲学+噪声鲁棒卖点), 硬件不合适(d²记忆), 偏置公理最差档(逐方向白化放大低方差偏置)。
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014FAPDWQ49M5Ye3NpTndTpn
Diffstat (limited to 'ep_run/muon.py')
| -rw-r--r-- | ep_run/muon.py | 28 |
1 files changed, 27 insertions, 1 deletions
diff --git a/ep_run/muon.py b/ep_run/muon.py index 7bad679..86c81d3 100644 --- a/ep_run/muon.py +++ b/ep_run/muon.py @@ -157,6 +157,29 @@ class Adafactor2D(torch.optim.Optimizer): p.add_(u, alpha=-g_['lr']) +class PSGDQuadWrap(torch.optim.Optimizer): + """Xi-Lin Li's KronWhiten (dQ='QUAD') as a proper torch Optimizer so LambdaLR accepts it. + Our EP flow pre-populates p.grad; the closure hands PSGD a synthetic scalar whose autograd + gradient equals the stored p.grad (loss = sum <p, p.grad_detached>). Screening-tier: the + inner preconditioner state is NOT checkpointed (base state_dict covers param_groups only).""" + def __init__(self, params, lr=1e-3, momentum=0.95): + params = list(params) + super().__init__(params, dict(lr=lr)) + from psgd_vendor import KronWhiten + self._flat = [p for g_ in self.param_groups for p in g_['params']] + self.inner = KronWhiten(self._flat, preconditioner_init_scale=1.0, + lr_params=lr, lr_preconditioner=0.1, momentum=momentum, + whiten_grad=True, dQ="QUAD") + + @torch.no_grad() + def step(self, closure=None): + self.inner.lr_params = self.param_groups[0]['lr'] + flat = self._flat + def _closure(): + return sum((p * p.grad.detach()).sum() for p in flat if p.grad is not None) + self.inner.step(_closure) + + def build_alt(opt_name, blocks, other_params, lr, warmup, total_steps=0, lr_min_ratio=0.1, lr_matrix=None, wd=0.0): """Screening-tier builder for the optimizer price list: OPT on block matrices + AdamW on the @@ -181,7 +204,8 @@ def build_alt(opt_name, blocks, other_params, lr, warmup, total_steps=0, lr_min_ 'olionns': lambda: OLion(mats, lr=lm, wd=wd, ns_steps=0), 'olionk1': lambda: OLion(mats, lr=lm, wd=wd, ns_steps=1), 'olionk2': lambda: OLion(mats, lr=lm, wd=wd, ns_steps=2), - 'olionk3': lambda: OLion(mats, lr=lm, wd=wd, ns_steps=3)}[opt_name]() + 'olionk3': lambda: OLion(mats, lr=lm, wd=wd, ns_steps=3), + 'psgdquad': lambda: PSGDQuadWrap(mats, lr=lm)}[opt_name]() oa = torch.optim.AdamW(rest, lr=lr, weight_decay=1e-4) opts = [om, oa] if total_steps > 0: @@ -296,3 +320,5 @@ class DitherLion(torch.optim.Optimizer): if g_['wd'] > 0: p.mul_(1 - g_['lr'] * g_['wd']) p.add_((u + d).sign_(), alpha=-g_['lr']) m.mul_(b2).add_(p.grad, alpha=1 - b2) + + |
