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:
@@ -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",
|
||||
|
||||
@@ -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
@@ -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`.
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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())
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user