refactor: extract composite optimizer helpers and unify naming

- add astrai/optim/composite.py with shared step/zero_grad/state_dict/param_groups helpers and OptimizerFactory
- rename MuonMix to MuonAdamW (matches registered name muon_adamw) and file to muon_adamw.py
- use @OptimizerFactory.register decorator in each optimizer module instead of post-import registration in __init__
- fix closure being invoked once per sub-optimizer in MuonAdamW.step (now exactly once via composite_step)
- NoraNAdamW.step now forwards closure correctly
This commit is contained in:
2026-08-01 08:07:45 +08:00
parent 25c9e81b2b
commit 11073bd1d2
5 changed files with 119 additions and 52 deletions
+13 -14
View File
@@ -1,15 +1,13 @@
"""Optimizer implementations and factory registration."""
from torch.optim import Optimizer
from astrai.factory import BaseFactory
class OptimizerFactory(BaseFactory[Optimizer]):
"""Factory for built-in training optimizers."""
from astrai.optim.muon_mix import MuonMix
from astrai.optim.composite import (
OptimizerFactory,
composite_state_dict,
composite_step,
composite_zero_grad,
refresh_param_groups,
)
from astrai.optim.muon_adamw import MuonAdamW
from astrai.optim.nora_nadamw import (
NAdamW,
Nora,
@@ -20,17 +18,18 @@ from astrai.optim.nora_nadamw import (
partition_optimizer_parameters,
)
OptimizerFactory.register("nora_nadamw")(NoraNAdamW)
OptimizerFactory.register("muon_adamw")(MuonMix)
__all__ = [
"MuonMix",
"MuonAdamW",
"NAdamW",
"Nora",
"NoraNAdamW",
"OptimizerFactory",
"OptimizerParameterGroups",
"composite_state_dict",
"composite_step",
"composite_zero_grad",
"nora_direction",
"nora_lr_scale",
"partition_optimizer_parameters",
"refresh_param_groups",
]
+71
View File
@@ -0,0 +1,71 @@
"""Shared infrastructure for the optim package.
This module hosts two things:
* ``OptimizerFactory`` — the registry for built-in optimizers. Defining it
here (rather than in ``__init__.py``) lets each optimizer module import it
and register itself with a decorator, avoiding circular imports.
* Composite-optimizer helpers — ``step``/``zero_grad``/``state_dict``/
``param_groups`` delegation shared by every optimizer that routes different
parameter groups through distinct sub-optimizers.
"""
from typing import Any
import torch
from torch.optim import Optimizer
from astrai.factory import BaseFactory
class OptimizerFactory(BaseFactory[Optimizer]):
"""Factory for built-in training optimizers."""
def composite_step(
sub_optimizers: list[Optimizer],
closure=None,
) -> torch.Tensor | None:
"""Run ``step`` on every sub-optimizer, invoking the closure once.
The closure (if given) is executed inside ``torch.enable_grad`` exactly
once before any sub-optimizer steps, matching the contract of a single
``Optimizer.step``. Sub-optimizers receive ``None`` so they do not
re-execute it.
"""
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for sub in sub_optimizers:
sub.step()
return loss
def composite_zero_grad(
sub_optimizers: list[Optimizer],
set_to_none: bool = True,
) -> None:
for sub in sub_optimizers:
sub.zero_grad(set_to_none=set_to_none)
def composite_state_dict(
named_sub_optimizers: dict[str, Optimizer | None],
) -> dict[str, Any]:
"""Serialize sub-optimizers, preserving ``None`` slots."""
return {
name: sub.state_dict() if sub is not None else None
for name, sub in named_sub_optimizers.items()
}
def refresh_param_groups(
sub_optimizers: list[Optimizer],
) -> list[dict]:
"""Concatenate param_groups from every non-None sub-optimizer."""
groups: list[dict] = []
for sub in sub_optimizers:
if sub is not None:
groups.extend(sub.param_groups)
return groups
@@ -5,8 +5,17 @@ from typing import Any
import torch
from torch import Tensor, nn, optim
from astrai.optim.composite import (
OptimizerFactory,
composite_state_dict,
composite_step,
composite_zero_grad,
refresh_param_groups,
)
class MuonMix(optim.Optimizer):
@OptimizerFactory.register("muon_adamw")
class MuonAdamW(optim.Optimizer):
"""Combined Muon (matrix) + AdamW (non-matrix) optimizer."""
optimizer_name = "muon_adamw"
@@ -64,22 +73,17 @@ class MuonMix(optim.Optimizer):
fused=True,
)
self.param_groups = [*self.muon.param_groups, *self.adamw.param_groups]
self.param_groups = refresh_param_groups([self.muon, self.adamw])
@torch.no_grad()
def step(self, closure=None):
self.muon.step(closure)
self.adamw.step(closure)
return composite_step([self.muon, self.adamw], closure)
def zero_grad(self, set_to_none: bool = True):
self.muon.zero_grad(set_to_none=set_to_none)
self.adamw.zero_grad(set_to_none=set_to_none)
composite_zero_grad([self.muon, self.adamw], set_to_none)
def state_dict(self) -> dict[str, Any]:
return {
"muon": self.muon.state_dict(),
"adamw": self.adamw.state_dict(),
}
return composite_state_dict({"muon": self.muon, "adamw": self.adamw})
def load_state_dict(self, state_dict: dict[str, Any]):
if "muon" not in state_dict or "adamw" not in state_dict:
@@ -88,4 +92,4 @@ class MuonMix(optim.Optimizer):
)
self.muon.load_state_dict(state_dict["muon"])
self.adamw.load_state_dict(state_dict["adamw"])
self.param_groups = [*self.muon.param_groups, *self.adamw.param_groups]
self.param_groups = refresh_param_groups([self.muon, self.adamw])
+19 -26
View File
@@ -13,6 +13,13 @@ from astrai.model.components.embedding import Embedding
from astrai.model.components.linear import Linear
from astrai.model.components.lora import LoRALinear
from astrai.model.components.norm import RMSNorm
from astrai.optim.composite import (
OptimizerFactory,
composite_state_dict,
composite_step,
composite_zero_grad,
refresh_param_groups,
)
NORA_EPS = 1e-10
@@ -271,6 +278,7 @@ def partition_optimizer_parameters(model: nn.Module) -> OptimizerParameterGroups
return OptimizerParameterGroups(nora, nadamw_decay, nadamw_no_decay)
@OptimizerFactory.register("nora_nadamw")
class NoraNAdamW(Optimizer):
"""Nora for internal linear weights and NAdamW for remaining parameters."""
@@ -320,38 +328,23 @@ class NoraNAdamW(Optimizer):
{"params": groups.nadamw_no_decay, "weight_decay": 0.0}
)
self.nadamw = NAdamW(nadamw_groups, lr=lr) if nadamw_groups else None
self._refresh_param_groups()
def _refresh_param_groups(self) -> None:
self.param_groups = []
if self.nora is not None:
self.param_groups.extend(self.nora.param_groups)
if self.nadamw is not None:
self.param_groups.extend(self.nadamw.param_groups)
self.param_groups = refresh_param_groups([self.nora, self.nadamw])
@torch.no_grad()
def step(self, closure=None):
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
if self.nora is not None:
self.nora.step()
if self.nadamw is not None:
self.nadamw.step()
return loss
return composite_step(
[opt for opt in (self.nora, self.nadamw) if opt is not None],
closure,
)
def zero_grad(self, set_to_none: bool = True):
if self.nora is not None:
self.nora.zero_grad(set_to_none=set_to_none)
if self.nadamw is not None:
self.nadamw.zero_grad(set_to_none=set_to_none)
composite_zero_grad(
[opt for opt in (self.nora, self.nadamw) if opt is not None],
set_to_none,
)
def state_dict(self) -> dict[str, Any]:
return {
"nora": self.nora.state_dict() if self.nora is not None else None,
"nadamw": self.nadamw.state_dict() if self.nadamw is not None else None,
}
return composite_state_dict({"nora": self.nora, "nadamw": self.nadamw})
def load_state_dict(self, state_dict: dict[str, Any]):
if "muon" in state_dict or "adamw" in state_dict:
@@ -376,4 +369,4 @@ class NoraNAdamW(Optimizer):
self.nora.load_state_dict(saved_nora)
if self.nadamw is not None:
self.nadamw.load_state_dict(saved_nadamw)
self._refresh_param_groups()
self.param_groups = refresh_param_groups([self.nora, self.nadamw])
+1 -1
View File
@@ -56,7 +56,7 @@ under DTensor sharding and rejects layouts sharded along the last dimension.
| `--nora_weight_decay` | Nora matrix weight decay | 0.0 |
Optimizer identity and hyperparameters are saved in checkpoint metadata. Optimizer
states are intentionally not interchangeable: resume older MuonMix checkpoints
states are intentionally not interchangeable: resume older MuonAdamW checkpoints
with `--optimizer=muon_adamw`.
### Data Loading