diff options
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) + + |
