summaryrefslogtreecommitdiff
path: root/ep_run/muon.py
diff options
context:
space:
mode:
authorYuren Hao <yurenh2@illinois.edu>2026-07-13 12:01:46 -0500
committerYuren Hao <yurenh2@illinois.edu>2026-07-13 12:01:46 -0500
commit907588538ead4d005b6cf4c0fa26ac2450fa5d7e (patch)
tree6b50012cb5e090b4f4106f1152ec7ee6c576362d /ep_run/muon.py
parentf71f449e89b505286861021bde1af8faa8e41660 (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.py6
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: