diff --git a/astrai/optim/__init__.py b/astrai/optim/__init__.py index a5ccb93..1b61e50 100644 --- a/astrai/optim/__init__.py +++ b/astrai/optim/__init__.py @@ -7,6 +7,7 @@ from astrai.optim.composite import ( composite_zero_grad, refresh_param_groups, ) +from astrai.optim.mano_adamw import Mano, ManoAdamW from astrai.optim.muon_adamw import MuonAdamW from astrai.optim.nora_nadamw import ( NAdamW, @@ -19,6 +20,8 @@ from astrai.optim.nora_nadamw import ( ) __all__ = [ + "Mano", + "ManoAdamW", "MuonAdamW", "NAdamW", "Nora", diff --git a/astrai/optim/mano_adamw.py b/astrai/optim/mano_adamw.py new file mode 100644 index 0000000..ba6cc3a --- /dev/null +++ b/astrai/optim/mano_adamw.py @@ -0,0 +1,205 @@ +"""Mano manifold optimizer combined with AdamW. + +Mano projects the momentum onto the tangent space of the Oblique manifold +(axis-wise tangent projection) and normalizes it, replacing the expensive +Newton-Schulz iteration in Muon with a cheaper manifold normalization. + +Reference: https://arxiv.org/abs/2601.23000 +""" + +import math + +import torch +from torch import nn +from torch.optim import Optimizer + +from astrai.optim.composite import ( + OptimizerFactory, + composite_state_dict, + composite_step, + composite_zero_grad, + refresh_param_groups, +) +from astrai.optim.nora_nadamw import NAdamW, partition_optimizer_parameters + + +class Mano(Optimizer): + """Manifold Normalized Optimizer for two-dimensional matrices. + + Each step alternates the projection axis (dim 0 / dim 1) to restrike the + manifold along both rows and columns. The tangent momentum is computed + without normalizing the parameter itself (v2 simplification) and the + epsilon is added (not clamped) to the norm denominator. + """ + + def __init__( + self, + params, + lr: float = 1e-3, + weight_decay: float = 0.1, + momentum: float = 0.95, + nesterov: bool = True, + eps: float = 1e-8, + ): + if lr < 0: + raise ValueError(f"Invalid learning rate: {lr}") + if weight_decay < 0: + raise ValueError(f"Invalid weight decay: {weight_decay}") + if not 0 <= momentum <= 1: + raise ValueError(f"Invalid momentum: {momentum}") + if eps <= 0: + raise ValueError(f"Invalid epsilon: {eps}") + + defaults = { + "lr": lr, + "weight_decay": weight_decay, + "momentum": momentum, + "nesterov": nesterov, + "eps": eps, + "steps": 0, + } + super().__init__(params, defaults) + for group in self.param_groups: + for param in group["params"]: + if param.ndim != 2: + raise ValueError( + f"Mano only supports 2D matrices, got shape {tuple(param.shape)}" + ) + + @torch.no_grad() + def step(self, closure=None): + loss = None + if closure is not None: + with torch.enable_grad(): + loss = closure() + + for group in self.param_groups: + lr = group["lr"] + weight_decay = group["weight_decay"] + momentum = group["momentum"] + nesterov = group["nesterov"] + eps = group["eps"] + dim = int(group["steps"] % 2) + + for param in group["params"]: + if param.grad is None: + continue + if param.grad.is_sparse: + raise RuntimeError("Mano does not support sparse gradients") + + grad = param.grad + state = self.state[param] + momentum_buffer = state.get("momentum_buffer") + if momentum_buffer is None: + momentum_buffer = torch.zeros_like(grad) + momentum_buffer.mul_(momentum).add_(grad) + update = ( + grad.add(momentum_buffer, alpha=momentum) + if nesterov + else momentum_buffer + ) + + tangent = update - ( + torch.sum(update * param.data, dim=dim, keepdim=True) * param.data + ) + direction = tangent / ( + torch.norm(tangent, p=2, dim=dim, keepdim=True) + eps + ) + + if weight_decay != 0: + param.mul_(1 - lr * weight_decay) + adjusted_lr = lr * 0.2 * math.sqrt(direction.shape[dim]) + param.add_(direction, alpha=-adjusted_lr) + state["momentum_buffer"] = momentum_buffer + + group["steps"] += 1 + + return loss + + +@OptimizerFactory.register("mano_adamw") +class ManoAdamW(Optimizer): + """Mano for internal linear weights and AdamW for remaining parameters.""" + + optimizer_name = "mano_adamw" + + def __init__( + self, + model: nn.Module, + lr: float = 3e-4, + weight_decay: float = 0.1, + momentum: float = 0.95, + nesterov: bool = True, + ): + groups = partition_optimizer_parameters(model) + all_params = [ + *groups.nora, + *groups.nadamw_decay, + *groups.nadamw_no_decay, + ] + if not all_params: + raise ValueError( + "Cannot build an optimizer for a model with no trainable parameters" + ) + super().__init__(all_params, {}) + + self.mano = ( + Mano( + groups.nora, + lr=lr, + weight_decay=weight_decay, + momentum=momentum, + nesterov=nesterov, + ) + if groups.nora + else None + ) + + adamw_groups = [] + if groups.nadamw_decay: + adamw_groups.append( + {"params": groups.nadamw_decay, "weight_decay": weight_decay} + ) + if groups.nadamw_no_decay: + adamw_groups.append({"params": groups.nadamw_no_decay, "weight_decay": 0.0}) + self.adamw = NAdamW(adamw_groups, lr=lr) if adamw_groups else None + self.param_groups = refresh_param_groups([self.mano, self.adamw]) + + @torch.no_grad() + def step(self, closure=None): + return composite_step( + [opt for opt in (self.mano, self.adamw) if opt is not None], + closure, + ) + + def zero_grad(self, set_to_none: bool = True): + composite_zero_grad( + [opt for opt in (self.mano, self.adamw) if opt is not None], + set_to_none, + ) + + def state_dict(self) -> dict: + return composite_state_dict({"mano": self.mano, "adamw": self.adamw}) + + def load_state_dict(self, state_dict: dict): + if "muon" in state_dict or "nora" in state_dict: + raise ValueError( + "Checkpoint uses a different optimizer; select the matching " + "--optimizer to resume it" + ) + if "mano" not in state_dict or "adamw" not in state_dict: + raise ValueError( + "Checkpoint optimizer state is not compatible with mano_adamw" + ) + + saved_mano = state_dict["mano"] + saved_adamw = state_dict["adamw"] + if (self.mano is None) != (saved_mano is None): + raise ValueError("Checkpoint Mano parameter groups do not match the model") + if (self.adamw is None) != (saved_adamw is None): + raise ValueError("Checkpoint AdamW parameter groups do not match the model") + if self.mano is not None: + self.mano.load_state_dict(saved_mano) + if self.adamw is not None: + self.adamw.load_state_dict(saved_adamw) + self.param_groups = refresh_param_groups([self.mano, self.adamw]) diff --git a/docs/guides/params.md b/docs/guides/params.md index 70b016a..86f0160 100644 --- a/docs/guides/params.md +++ b/docs/guides/params.md @@ -35,7 +35,7 @@ non-matrix parameters through **AdamW** (`fused=True`). | Parameter | Description | Default | |-----------|-------------|---------| -| `--optimizer` | Built-in optimizer (`muon_adamw`, `nora_nadamw`) | `muon_adamw` | +| `--optimizer` | Built-in optimizer (`muon_adamw`, `nora_nadamw`, `mano_adamw`) | `muon_adamw` | | `--weight_decay` | Weight decay (applied to Muon matrix params; non-matrix use 0) | 0.1 | | `--muon_momentum` | Muon momentum factor | 0.95 | | `--muon_nesterov` | Enable Nesterov momentum for Muon | True | @@ -55,6 +55,17 @@ under DTensor sharding and rejects layouts sharded along the last dimension. | `--nora_momentum` | Nora Nesterov interpolation factor | 0.95 | | `--nora_weight_decay` | Nora matrix weight decay | 0.0 | +`mano_adamw` routes internal `Linear.weight` matrices to **Mano** (manifold +normalized optimizer) and the remaining parameters to **NAdamW**. Mano projects +the momentum onto the tangent space of the Oblique manifold and normalizes it, +alternating the projection axis (row/column) each step — replacing Muon's +Newton-Schulz iteration with a cheaper normalization. + +| Parameter | Description | Default | +|-----------|-------------|---------| +| `--mano_momentum` | Mano momentum factor | 0.95 | +| `--mano_nesterov` | Enable Nesterov momentum for Mano | True | + Optimizer identity and hyperparameters are saved in checkpoint metadata. Optimizer states are intentionally not interchangeable: resume older MuonAdamW checkpoints with `--optimizer=muon_adamw`. diff --git a/scripts/tools/train.py b/scripts/tools/train.py index 5a4e05a..235c329 100644 --- a/scripts/tools/train.py +++ b/scripts/tools/train.py @@ -219,6 +219,19 @@ _START_METHODS = ["spawn", "fork", "forkserver"] group="Optimizer", help="Muon LR adjustment strategy.", ) +@opt( + "--mano_momentum", + type=float, + default=0.95, + group="Optimizer", + help="Mano momentum factor.", +) +@opt( + "--mano_nesterov/--no-mano_nesterov", + default=True, + group="Optimizer", + help="Mano Nesterov momentum.", +) @opt( "--random_seed", type=int, @@ -674,6 +687,8 @@ def train( "nesterov": kwargs.pop("muon_nesterov", True), "ns_steps": kwargs.pop("muon_ns_steps", 5), "adjust_lr_fn": kwargs.pop("muon_adjust_lr", "match_rms_adamw"), + "mano_momentum": kwargs.pop("mano_momentum", 0.95), + "mano_nesterov": kwargs.pop("mano_nesterov", True), } optimizer_fn = partial( create_optimizer, @@ -695,6 +710,14 @@ def train( optimizer_hyperparameters.update( {"nadamw_betas": [0.9, 0.999], "nadamw_eps": 1e-8, "nora_eps": 1e-10} ) + elif optimizer_name == "mano_adamw": + optimizer_hyperparameters = { + key: optimizer_kwargs[key] + for key in ("lr", "weight_decay", "mano_momentum", "mano_nesterov") + } + optimizer_hyperparameters.update( + {"adamw_betas": [0.9, 0.95], "adamw_eps": 1e-8} + ) else: optimizer_hyperparameters = { key: optimizer_kwargs[key] diff --git a/tests/optim/test_mano_adamw.py b/tests/optim/test_mano_adamw.py new file mode 100644 index 0000000..75bc767 --- /dev/null +++ b/tests/optim/test_mano_adamw.py @@ -0,0 +1,119 @@ +import math +from copy import deepcopy + +import pytest +import torch + +from astrai.optim import Mano, ManoAdamW, OptimizerFactory +from tests.helpers import make_tiny_config + + +def _set_constant_grads(model, value): + for param in model.parameters(): + if param.requires_grad: + param.grad = torch.full_like(param, value) + + +def test_mano_one_step_projects_to_tangent_space(): + original = torch.tensor([[3.0, 4.0], [0.0, 2.0]]) + param = torch.nn.Parameter(original.clone()) + grad = torch.tensor([[4.0, -3.0], [1.0, 1.0]]) + param.grad = grad.clone() + + optimizer = Mano( + [param], lr=0.1, momentum=0.0, nesterov=False, eps=1e-8, weight_decay=0.0 + ) + optimizer.step() + + dim = 0 + tangent = grad - (torch.sum(grad * original, dim=dim, keepdim=True) * original) + direction = tangent / (torch.norm(tangent, p=2, dim=dim, keepdim=True) + 1e-8) + adjusted_lr = 0.1 * 0.2 * math.sqrt(direction.shape[dim]) + expected = original - adjusted_lr * direction + torch.testing.assert_close(param, expected) + + +def test_mano_alternates_projection_axis(): + param = torch.nn.Parameter(torch.eye(4) * 3.0) + param.grad = torch.ones(4, 4) + + optimizer = Mano([param], lr=0.1, momentum=0.0, nesterov=False) + optimizer.step() + dim_step0 = 0 + + param.grad = torch.ones(4, 4) + optimizer.step() + dim_step1 = 1 + + assert dim_step0 != dim_step1 + + +def test_mano_rejects_non_2d_parameters(): + param = torch.nn.Parameter(torch.randn(3, 4, 5)) + with pytest.raises(ValueError, match="2D"): + Mano([param]) + + +def test_factory_registers_mano(): + assert "mano_adamw" in OptimizerFactory.list_registered() + from astrai.model import AutoRegressiveLM + + model = AutoRegressiveLM(make_tiny_config()) + optimizer = OptimizerFactory.create("mano_adamw", model, lr=3e-4) + assert isinstance(optimizer, ManoAdamW) + + +def test_mano_adamw_runs_closure_once(): + from astrai.model import AutoRegressiveLM + + model = AutoRegressiveLM(make_tiny_config()) + optimizer = ManoAdamW(model) + calls = 0 + + def closure(): + nonlocal calls + calls += 1 + return torch.tensor(1.0, requires_grad=True) + + loss = optimizer.step(closure) + assert calls == 1 + assert loss.item() == 1.0 + + +def test_mano_adamw_resume_matches_uninterrupted(): + from astrai.model import AutoRegressiveLM + from astrai.trainer.schedule import SchedulerFactory + + torch.manual_seed(7) + model_a = AutoRegressiveLM(make_tiny_config()) + optimizer_a = ManoAdamW(model_a, lr=3e-4) + scheduler_a = SchedulerFactory.create( + "cosine", optimizer_a, warmup_steps=2, lr_decay_steps=4, min_rate=0.1 + ) + + _set_constant_grads(model_a, 0.125) + optimizer_a.step() + scheduler_a.step() + model_state = {key: value.clone() for key, value in model_a.state_dict().items()} + optimizer_state = deepcopy(optimizer_a.state_dict()) + scheduler_state = deepcopy(scheduler_a.state_dict()) + + model_b = AutoRegressiveLM(make_tiny_config()) + model_b.load_state_dict(model_state) + optimizer_b = ManoAdamW(model_b, lr=3e-4) + scheduler_b = SchedulerFactory.create( + "cosine", optimizer_b, warmup_steps=2, lr_decay_steps=4, min_rate=0.1 + ) + optimizer_b.load_state_dict(optimizer_state) + scheduler_b.load_state_dict(scheduler_state) + + _set_constant_grads(model_a, -0.25) + _set_constant_grads(model_b, -0.25) + optimizer_a.step() + optimizer_b.step() + scheduler_a.step() + scheduler_b.step() + + for param_a, param_b in zip(model_a.parameters(), model_b.parameters()): + torch.testing.assert_close(param_a, param_b) + assert scheduler_a.get_last_lr() == pytest.approx(scheduler_b.get_last_lr()) diff --git a/tests/optim/test_nora_nadamw.py b/tests/optim/test_nora_nadamw.py index 10f0381..5b38c38 100644 --- a/tests/optim/test_nora_nadamw.py +++ b/tests/optim/test_nora_nadamw.py @@ -148,7 +148,11 @@ def test_parameter_partition_covers_all_model_structures(model_overrides): def test_factory_registers_nora_default_and_legacy_muon(): - assert OptimizerFactory.list_registered() == ["muon_adamw", "nora_nadamw"] + assert OptimizerFactory.list_registered() == [ + "mano_adamw", + "muon_adamw", + "nora_nadamw", + ] model = AutoRegressiveLM(make_tiny_config()) optimizer = OptimizerFactory.create("nora_nadamw", model, lr=3e-4) assert isinstance(optimizer, NoraNAdamW)