summaryrefslogtreecommitdiff
path: root/ep_run/muon.py
diff options
context:
space:
mode:
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)
+
+