diff options
Diffstat (limited to 'ep_run/muon.py')
| -rw-r--r-- | ep_run/muon.py | 14 |
1 files changed, 11 insertions, 3 deletions
diff --git a/ep_run/muon.py b/ep_run/muon.py index 5ac6fc6..3116030 100644 --- a/ep_run/muon.py +++ b/ep_run/muon.py @@ -52,13 +52,21 @@ class MultiSched: for s in self.scheds: s.step() -def build_hybrid(blocks, other_params, lr_adamw, lr_muon, warmup): - """Muon(2D block matrices) + AdamW(everything else), with linear-warmup scheds for both.""" +def build_hybrid(blocks, other_params, lr_adamw, lr_muon, warmup, total_steps=0, lr_min_ratio=0.1): + """Muon(2D block matrices) + AdamW(everything else). Scheds: linear warmup, then cosine decay to + lr_min_ratio*peak if total_steps>0 (long runs), else constant after warmup (legacy).""" + import math as _m mats = [p for p in blocks.parameters() if p.ndim == 2] mat_ids = {id(p) for p in mats} rest = [p for p in other_params if id(p) not in mat_ids] om = Muon(mats, lr=lr_muon) oa = torch.optim.AdamW(rest, lr=lr_adamw, weight_decay=1e-4) - fn = lambda s: min(1.0, (s + 1) / max(warmup, 1)) + if total_steps > 0: + def fn(s): + if s < warmup: return (s + 1) / max(warmup, 1) + p = min(1.0, (s - warmup) / max(1, total_steps - warmup)) + return lr_min_ratio + 0.5 * (1 - lr_min_ratio) * (1 + _m.cos(_m.pi * p)) + else: + fn = lambda s: min(1.0, (s + 1) / max(warmup, 1)) scheds = [torch.optim.lr_scheduler.LambdaLR(om, fn), torch.optim.lr_scheduler.LambdaLR(oa, fn)] return MultiOpt([om, oa]), MultiSched(scheds) |
