feat: add Mano manifold optimizer (mano_adamw)

- implement Mano (v2) with axis-rotating tangent projection and manifold normalization, replacing Newton-Schulz iteration
- composite ManoAdamW reuses partition_optimizer_parameters and composite helpers
- register mano_adamw in OptimizerFactory, export Mano and ManoAdamW
- add --mano_momentum and --mano_nesterov CLI options in Optimizer group
- add mano_adamw hyperparameters branch in train.py
- document mano_adamw in params.md
- add tests for single-step projection, axis alternation, factory registration, closure, and resume
This commit is contained in:
2026-08-01 08:51:08 +08:00
parent 6c76c16480
commit 6db276f37a
6 changed files with 367 additions and 2 deletions
+3
View File
@@ -7,6 +7,7 @@ from astrai.optim.composite import (
composite_zero_grad, composite_zero_grad,
refresh_param_groups, refresh_param_groups,
) )
from astrai.optim.mano_adamw import Mano, ManoAdamW
from astrai.optim.muon_adamw import MuonAdamW from astrai.optim.muon_adamw import MuonAdamW
from astrai.optim.nora_nadamw import ( from astrai.optim.nora_nadamw import (
NAdamW, NAdamW,
@@ -19,6 +20,8 @@ from astrai.optim.nora_nadamw import (
) )
__all__ = [ __all__ = [
"Mano",
"ManoAdamW",
"MuonAdamW", "MuonAdamW",
"NAdamW", "NAdamW",
"Nora", "Nora",
+205
View File
@@ -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])
+12 -1
View File
@@ -35,7 +35,7 @@ non-matrix parameters through **AdamW** (`fused=True`).
| Parameter | Description | Default | | 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 | | `--weight_decay` | Weight decay (applied to Muon matrix params; non-matrix use 0) | 0.1 |
| `--muon_momentum` | Muon momentum factor | 0.95 | | `--muon_momentum` | Muon momentum factor | 0.95 |
| `--muon_nesterov` | Enable Nesterov momentum for Muon | True | | `--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_momentum` | Nora Nesterov interpolation factor | 0.95 |
| `--nora_weight_decay` | Nora matrix weight decay | 0.0 | | `--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 Optimizer identity and hyperparameters are saved in checkpoint metadata. Optimizer
states are intentionally not interchangeable: resume older MuonAdamW checkpoints states are intentionally not interchangeable: resume older MuonAdamW checkpoints
with `--optimizer=muon_adamw`. with `--optimizer=muon_adamw`.
+23
View File
@@ -219,6 +219,19 @@ _START_METHODS = ["spawn", "fork", "forkserver"]
group="Optimizer", group="Optimizer",
help="Muon LR adjustment strategy.", 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( @opt(
"--random_seed", "--random_seed",
type=int, type=int,
@@ -674,6 +687,8 @@ def train(
"nesterov": kwargs.pop("muon_nesterov", True), "nesterov": kwargs.pop("muon_nesterov", True),
"ns_steps": kwargs.pop("muon_ns_steps", 5), "ns_steps": kwargs.pop("muon_ns_steps", 5),
"adjust_lr_fn": kwargs.pop("muon_adjust_lr", "match_rms_adamw"), "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( optimizer_fn = partial(
create_optimizer, create_optimizer,
@@ -695,6 +710,14 @@ def train(
optimizer_hyperparameters.update( optimizer_hyperparameters.update(
{"nadamw_betas": [0.9, 0.999], "nadamw_eps": 1e-8, "nora_eps": 1e-10} {"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: else:
optimizer_hyperparameters = { optimizer_hyperparameters = {
key: optimizer_kwargs[key] key: optimizer_kwargs[key]
+119
View File
@@ -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())
+5 -1
View File
@@ -148,7 +148,11 @@ def test_parameter_partition_covers_all_model_structures(model_overrides):
def test_factory_registers_nora_default_and_legacy_muon(): 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()) model = AutoRegressiveLM(make_tiny_config())
optimizer = OptimizerFactory.create("nora_nadamw", model, lr=3e-4) optimizer = OptimizerFactory.create("nora_nadamw", model, lr=3e-4)
assert isinstance(optimizer, NoraNAdamW) assert isinstance(optimizer, NoraNAdamW)