diff --git a/astrai/optim/__init__.py b/astrai/optim/__init__.py index 7fcc7aa..a5ccb93 100644 --- a/astrai/optim/__init__.py +++ b/astrai/optim/__init__.py @@ -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", ] diff --git a/astrai/optim/composite.py b/astrai/optim/composite.py new file mode 100644 index 0000000..1d40ccf --- /dev/null +++ b/astrai/optim/composite.py @@ -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 diff --git a/astrai/optim/muon_mix.py b/astrai/optim/muon_adamw.py similarity index 79% rename from astrai/optim/muon_mix.py rename to astrai/optim/muon_adamw.py index c48b287..f27f5f3 100644 --- a/astrai/optim/muon_mix.py +++ b/astrai/optim/muon_adamw.py @@ -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]) diff --git a/astrai/optim/nora_nadamw.py b/astrai/optim/nora_nadamw.py index 2a4512f..1400429 100644 --- a/astrai/optim/nora_nadamw.py +++ b/astrai/optim/nora_nadamw.py @@ -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]) diff --git a/docs/guides/params.md b/docs/guides/params.md index c8e103c..70b016a 100644 --- a/docs/guides/params.md +++ b/docs/guides/params.md @@ -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