diff options
| author | Yuren Hao <yurenh2@illinois.edu> | 2026-07-13 12:01:46 -0500 |
|---|---|---|
| committer | Yuren Hao <yurenh2@illinois.edu> | 2026-07-13 12:01:46 -0500 |
| commit | 907588538ead4d005b6cf4c0fa26ac2450fa5d7e (patch) | |
| tree | 6b50012cb5e090b4f4106f1152ec7ee6c576362d /ep_run/muon.py | |
| parent | f71f449e89b505286861021bde1af8faa8e41660 (diff) | |
fw72m crown run launched (NCCL DDP prod), exact resume (opt state), RESULT 17, shuffle+prep scripts
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.py | 6 |
1 files changed, 5 insertions, 1 deletions
diff --git a/ep_run/muon.py b/ep_run/muon.py index 3116030..7281811 100644 --- a/ep_run/muon.py +++ b/ep_run/muon.py @@ -38,12 +38,16 @@ class Muon(torch.optim.Optimizer): class MultiOpt: - """duck-typed bundle of optimizers (step/zero_grad API-compatible).""" + """duck-typed bundle of optimizers (step/zero_grad/state_dict API-compatible).""" def __init__(self, opts): self.optimizers = opts def step(self): for o in self.optimizers: o.step() def zero_grad(self, set_to_none=True): for o in self.optimizers: o.zero_grad(set_to_none=set_to_none) + def state_dict(self): + return [o.state_dict() for o in self.optimizers] + def load_state_dict(self, sds): + for o, sd in zip(self.optimizers, sds): o.load_state_dict(sd) class MultiSched: |
