From 2de7789ecf05eba74cc52dbd5f4d8298e4084e50 Mon Sep 17 00:00:00 2001 From: Yuren Hao Date: Wed, 5 Aug 2026 19:14:33 -0500 Subject: =?UTF-8?q?PSGD-QUAD=E5=85=A5=E5=88=97:=20vendor=20psgd.py=20+=20O?= =?UTF-8?q?ptimizer=E5=AD=90=E7=B1=BB=E9=80=82=E9=85=8D=E5=99=A8(=E5=90=88?= =?UTF-8?q?=E6=88=90=E6=8D=9F=E5=A4=B1closure=E6=A1=A5=E6=8E=A5EP=E9=A2=84?= =?UTF-8?q?=E5=A1=AB=E6=A2=AF=E5=BA=A6),=2013=E8=B7=AF=E5=85=A8=E9=87=8F?= =?UTF-8?q?=E5=9B=9E=E5=BD=92=E5=86=92=E7=83=9FPASS?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 修复过程记录: 首版适配器非Optimizer子类被LambdaLR拒; 重写时splice错序造成muon.py重复段 (旧类后定义胜出), 按结构图手术去重并断言全符号在位。三轴定位: GPU配方候选(追踪慢变统计= R4幸存哲学+噪声鲁棒卖点), 硬件不合适(d²记忆), 偏置公理最差档(逐方向白化放大低方差偏置)。 Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_014FAPDWQ49M5Ye3NpTndTpn --- ep_run/muon.py | 28 +++++++++++++++++++++++++++- 1 file changed, 27 insertions(+), 1 deletion(-) (limited to 'ep_run/muon.py') 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 ). 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) + + -- cgit v1.2.3