summaryrefslogtreecommitdiff
path: root/ep_run/muon.py
diff options
context:
space:
mode:
authorYuren Hao <yurenh2@illinois.edu>2026-08-05 19:14:33 -0500
committerYuren Hao <yurenh2@illinois.edu>2026-08-05 19:14:33 -0500
commit2de7789ecf05eba74cc52dbd5f4d8298e4084e50 (patch)
tree5546f171017966ea7294c4c09a67e4a09778a891 /ep_run/muon.py
parent0df61c1d36faf3e985394c943853b5bab5adaa8c (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.py28
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)
+
+