Make Nora+NAdamW the default optimizer

This commit is contained in:
QueenAmish
2026-07-31 23:16:39 +08:00
parent 7aa5ed09d9
commit 04899a2b15
11 changed files with 1010 additions and 105 deletions
+2
View File
@@ -101,7 +101,9 @@ nohup python scripts/tools/train.py \
--batch_per_device=4 \
--grad_accum_steps=8 \
--warmup_ratio=0.05 \
--optimizer=nora_nadamw \
--max_lr=1e-4 \
--nora_lr=5e-3 \
--max_grad_norm=1.0 \
--weight_decay=0.1 \
--window_size=2048 \
+4
View File
@@ -30,6 +30,8 @@ class TrainConfig(BaseConfig):
strategy (str): Training strategy (seq, sft, dpo, grpo, online_*).
dataset (Dataset): Dataset for training.
optimizer_fn (Callable[[nn.Module], Optimizer]): Optimizer factory for training.
optimizer_name (Optional[str]): Serializable built-in optimizer identifier. Defaults to None.
optimizer_hyperparameters (Dict[str, Any]): Serializable optimizer settings. Defaults to {}.
scheduler_fn (Callable[[Optimizer], LRScheduler]): Scheduler factory for training.
n_epoch (int): Number of epochs for training. Defaults to 1.
batch_per_device (int): Batch size per device. Defaults to 4.
@@ -74,6 +76,8 @@ class TrainConfig(BaseConfig):
dataset: Dataset
optimizer_fn: Callable[[nn.Module], Optimizer]
scheduler_fn: Callable[[Optimizer], LRScheduler]
optimizer_name: Optional[str] = None
optimizer_hyperparameters: Dict[str, Any] = field(default_factory=dict)
n_epoch: int = 1
batch_per_device: int = 4
grad_accum_steps: int = 1
+36
View File
@@ -0,0 +1,36 @@
"""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.nora_nadamw import (
NAdamW,
Nora,
NoraNAdamW,
OptimizerParameterGroups,
nora_direction,
nora_lr_scale,
partition_optimizer_parameters,
)
OptimizerFactory.register("nora_nadamw")(NoraNAdamW)
OptimizerFactory.register("muon_adamw")(MuonMix)
__all__ = [
"MuonMix",
"NAdamW",
"Nora",
"NoraNAdamW",
"OptimizerFactory",
"OptimizerParameterGroups",
"nora_direction",
"nora_lr_scale",
"partition_optimizer_parameters",
]
+91
View File
@@ -0,0 +1,91 @@
"""Legacy Muon + AdamW combined optimizer."""
from typing import Any
import torch
from torch import Tensor, nn, optim
class MuonMix(optim.Optimizer):
"""Combined Muon (matrix) + AdamW (non-matrix) optimizer."""
optimizer_name = "muon_adamw"
def __init__(
self,
model: nn.Module,
lr: float = 3e-4,
weight_decay: float = 0.1,
momentum: float = 0.95,
nesterov: bool = True,
ns_steps: int = 5,
adjust_lr_fn: str = "match_rms_adamw",
):
defaults = {
"lr": lr,
"weight_decay": weight_decay,
"momentum": momentum,
"nesterov": nesterov,
"ns_steps": ns_steps,
"adjust_lr_fn": adjust_lr_fn,
}
params = [param for param in model.parameters() if param.requires_grad]
super().__init__(params, defaults)
matrix_params: list[Tensor] = []
other_params: list[Tensor] = []
for name, param in model.named_parameters():
if not param.requires_grad:
continue
if (
param.dim() >= 2
and "norm" not in name
and "bias" not in name
and "embed" not in name
and "lm_head" not in name
):
matrix_params.append(param)
else:
other_params.append(param)
self.muon = optim.Muon(
matrix_params,
lr=lr,
weight_decay=weight_decay,
momentum=momentum,
nesterov=nesterov,
ns_steps=ns_steps,
adjust_lr_fn=adjust_lr_fn,
)
self.adamw = optim.AdamW(
[{"params": other_params, "weight_decay": 0.0}],
lr=lr,
betas=(0.9, 0.95),
fused=True,
)
self.param_groups = [*self.muon.param_groups, *self.adamw.param_groups]
@torch.no_grad()
def step(self, closure=None):
self.muon.step(closure)
self.adamw.step(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)
def state_dict(self) -> dict[str, Any]:
return {
"muon": self.muon.state_dict(),
"adamw": self.adamw.state_dict(),
}
def load_state_dict(self, state_dict: dict[str, Any]):
if "muon" not in state_dict or "adamw" not in state_dict:
raise ValueError(
"Checkpoint optimizer state is not compatible with muon_adamw"
)
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]
+379
View File
@@ -0,0 +1,379 @@
"""Nora matrix optimizer combined with Nesterov AdamW."""
import math
from dataclasses import dataclass
from typing import Any
import torch
from torch import Tensor, nn
from torch.distributed.tensor import DTensor, Shard
from torch.optim import Optimizer
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
NORA_EPS = 1e-10
def _row_normalize(tensor: Tensor, eps: float) -> Tensor:
return tensor / tensor.norm(dim=-1, keepdim=True).clamp(min=eps)
def nora_direction(update: Tensor, param: Tensor, eps: float = NORA_EPS) -> Tensor:
"""Project an update onto each parameter row's tangent space and normalize."""
theta_hat = _row_normalize(param.to(torch.float32), eps)
update_fp32 = update.to(torch.float32)
radial = (update_fp32 * theta_hat).sum(dim=-1, keepdim=True) * theta_hat
direction = _row_normalize(update_fp32 - radial, eps)
return direction.to(update.dtype)
def nora_lr_scale(lr: float, shape: torch.Size) -> float:
"""Scale Nora's LR for tall ``[d_out, d_in]`` linear weights."""
return lr * math.sqrt(max(1.0, shape[-2] / shape[-1]))
def _validate_complete_rows(param: Tensor) -> None:
if not isinstance(param, DTensor):
return
last_dim = param.ndim - 1
for placement in param.placements:
if isinstance(placement, Shard) and placement.dim % param.ndim == last_dim:
raise ValueError(
"Nora requires complete parameter rows, but this DTensor is sharded "
"along its last dimension"
)
class Nora(Optimizer):
"""Normalized Orthogonal Row Alignment for two-dimensional matrices."""
def __init__(
self,
params,
lr: float = 5e-3,
weight_decay: float = 0.0,
momentum: float = 0.95,
beta: float = 0.95,
nesterov: bool = True,
eps: float = NORA_EPS,
):
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 not 0 <= beta < 1:
raise ValueError(f"Invalid beta: {beta}")
if eps <= 0:
raise ValueError(f"Invalid epsilon: {eps}")
defaults = {
"lr": lr,
"weight_decay": weight_decay,
"momentum": momentum,
"beta": beta,
"nesterov": nesterov,
"eps": eps,
}
super().__init__(params, defaults)
for group in self.param_groups:
for param in group["params"]:
if param.ndim != 2:
raise ValueError(
f"Nora only supports 2D matrices, got shape {tuple(param.shape)}"
)
_validate_complete_rows(param)
@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"]
beta = group["beta"]
nesterov = group["nesterov"]
eps = group["eps"]
for param in group["params"]:
if param.grad is None:
continue
if param.grad.is_sparse:
raise RuntimeError("Nora 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.lerp_(grad, 1 - beta)
update = (
grad.lerp(momentum_buffer, momentum)
if nesterov
else momentum_buffer
)
direction = nora_direction(update, param, eps)
if weight_decay != 0:
param.mul_(1 - lr * weight_decay)
param.add_(direction, alpha=-nora_lr_scale(lr, param.shape))
state["momentum_buffer"] = momentum_buffer
return loss
class NAdamW(Optimizer):
"""AdamW using the reference Nesterov first-moment update."""
def __init__(
self,
params,
lr: float = 3e-4,
betas: tuple[float, float] = (0.9, 0.999),
eps: float = 1e-8,
weight_decay: float = 0.1,
):
beta1, beta2 = betas
if lr < 0:
raise ValueError(f"Invalid learning rate: {lr}")
if not 0 <= beta1 < 1 or not 0 <= beta2 < 1:
raise ValueError(f"Invalid betas: {betas}")
if eps <= 0:
raise ValueError(f"Invalid epsilon: {eps}")
if weight_decay < 0:
raise ValueError(f"Invalid weight decay: {weight_decay}")
defaults = {
"lr": lr,
"betas": betas,
"eps": eps,
"weight_decay": weight_decay,
}
super().__init__(params, defaults)
@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:
beta1, beta2 = group["betas"]
eps = group["eps"]
lr = group["lr"]
weight_decay = group["weight_decay"]
for param in group["params"]:
if param.grad is None:
continue
if param.grad.is_sparse:
raise RuntimeError("NAdamW does not support sparse gradients")
grad = param.grad
state = self.state[param]
if not state:
state["step"] = 0
state["m"] = torch.zeros_like(param)
state["v"] = torch.zeros_like(param)
state["step"] += 1
first_moment = state["m"]
second_moment = state["v"]
first_moment.mul_(beta1).add_(grad, alpha=1 - beta1)
second_moment.mul_(beta2).addcmul_(grad, grad, value=1 - beta2)
bias_correction1 = 1 - beta1 ** state["step"]
bias_correction2 = 1 - beta2 ** state["step"]
nesterov_moment = (
beta1 * first_moment + (1 - beta1) * grad
) / bias_correction1
corrected_second_moment = second_moment / bias_correction2
if weight_decay != 0:
param.mul_(1 - lr * weight_decay)
param.addcdiv_(
nesterov_moment,
corrected_second_moment.sqrt().add_(eps),
value=-lr,
)
return loss
@dataclass
class OptimizerParameterGroups:
nora: list[Tensor]
nadamw_decay: list[Tensor]
nadamw_no_decay: list[Tensor]
def partition_optimizer_parameters(model: nn.Module) -> OptimizerParameterGroups:
"""Partition trainable parameters by module role and parameter identity."""
nora_ids: set[int] = set()
no_decay_ids: set[int] = set()
for module_name, module in model.named_modules():
if isinstance(module, LoRALinear):
for param in module.parameters(recurse=False):
if param.requires_grad:
no_decay_ids.add(id(param))
continue
if isinstance(module, (Embedding, RMSNorm)):
for param in module.parameters(recurse=False):
if param.requires_grad:
no_decay_ids.add(id(param))
continue
if not isinstance(module, Linear):
continue
if module.bias is not None and module.bias.requires_grad:
no_decay_ids.add(id(module.bias))
if not module.weight.requires_grad:
continue
if module_name.rsplit(".", 1)[-1] == "lm_head":
no_decay_ids.add(id(module.weight))
elif module.weight.ndim == 2:
nora_ids.add(id(module.weight))
nora: list[Tensor] = []
nadamw_decay: list[Tensor] = []
nadamw_no_decay: list[Tensor] = []
seen: set[int] = set()
for param in model.parameters():
param_id = id(param)
if not param.requires_grad or param_id in seen:
continue
seen.add(param_id)
if param_id in no_decay_ids or param.ndim <= 1:
nadamw_no_decay.append(param)
elif param_id in nora_ids:
nora.append(param)
else:
nadamw_decay.append(param)
trainable_ids = {id(param) for param in model.parameters() if param.requires_grad}
grouped_ids = {id(param) for param in [*nora, *nadamw_decay, *nadamw_no_decay]}
if grouped_ids != trainable_ids:
missing = len(trainable_ids - grouped_ids)
extra = len(grouped_ids - trainable_ids)
raise RuntimeError(
f"Optimizer parameter partition is incomplete: missing={missing}, extra={extra}"
)
return OptimizerParameterGroups(nora, nadamw_decay, nadamw_no_decay)
class NoraNAdamW(Optimizer):
"""Nora for internal linear weights and NAdamW for remaining parameters."""
optimizer_name = "nora_nadamw"
def __init__(
self,
model: nn.Module,
lr: float = 3e-4,
weight_decay: float = 0.1,
nora_lr: float = 5e-3,
nora_weight_decay: float = 0.0,
nora_beta: float = 0.95,
nora_momentum: float = 0.95,
):
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.nora = (
Nora(
groups.nora,
lr=nora_lr,
weight_decay=nora_weight_decay,
momentum=nora_momentum,
beta=nora_beta,
)
if groups.nora
else None
)
nadamw_groups = []
if groups.nadamw_decay:
nadamw_groups.append(
{"params": groups.nadamw_decay, "weight_decay": weight_decay}
)
if groups.nadamw_no_decay:
nadamw_groups.append(
{"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)
@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
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)
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,
}
def load_state_dict(self, state_dict: dict[str, Any]):
if "muon" in state_dict or "adamw" in state_dict:
raise ValueError(
"Checkpoint uses muon_adamw state; select optimizer='muon_adamw' "
"to resume it"
)
if "nora" not in state_dict or "nadamw" not in state_dict:
raise ValueError(
"Checkpoint optimizer state is not compatible with nora_nadamw"
)
saved_nora = state_dict["nora"]
saved_nadamw = state_dict["nadamw"]
if (self.nora is None) != (saved_nora is None):
raise ValueError("Checkpoint Nora parameter groups do not match the model")
if (self.nadamw is None) != (saved_nadamw is None):
raise ValueError(
"Checkpoint NAdamW parameter groups do not match the model"
)
if self.nora is not None:
self.nora.load_state_dict(saved_nora)
if self.nadamw is not None:
self.nadamw.load_state_dict(saved_nadamw)
self._refresh_param_groups()
+2
View File
@@ -107,7 +107,9 @@ nohup python scripts/tools/train.py \
--batch_per_device=4 \
--grad_accum_steps=8 \
--warmup_ratio=0.05 \
--optimizer=nora_nadamw \
--max_lr=1e-4 \
--nora_lr=5e-3 \
--max_grad_norm=1.0 \
--weight_decay=0.1 \
--window_size=2048 \
+24 -4
View File
@@ -25,21 +25,39 @@
| Parameter | Description | Default |
|-----------|-------------|---------|
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
| `--max_lr` | NAdamW learning rate; schedulers scale every optimizer group proportionally | 3e-4 |
| `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | 1.0 |
### Optimizer (MuonMix)
### Optimizer
Combined optimizer: matrix parameters via **Muon**, non-matrix via **AdamW** (`fused=True`).
The default `nora_nadamw` optimizer sends internal `Linear.weight` matrices to
**Nora** and embeddings, the LM head, norms, biases, LoRA factors, and fallback
parameters to **NAdamW**. Parameters are classified by module role and identity,
so tied embedding/head weights occur in exactly one group. Nora requires complete
rows under DTensor sharding and rejects layouts sharded along the last dimension.
| Parameter | Description | Default |
|-----------|-------------|---------|
| `--optimizer` | Built-in optimizer (`nora_nadamw`, `muon_adamw`) | `nora_nadamw` |
| `--weight_decay` | NAdamW decay for eligible fallback parameters; known embeddings, heads, norms, biases, and LoRA factors use 0 | 0.1 |
| `--nora_lr` | Nora learning rate | 5e-3 |
| `--nora_beta` | Nora momentum-buffer EMA factor | 0.95 |
| `--nora_momentum` | Nora Nesterov interpolation factor | 0.95 |
| `--nora_weight_decay` | Nora matrix weight decay | 0.0 |
`muon_adamw` preserves the previous MuonMix behavior and the following options:
| Parameter | Description | Default |
|-----------|-------------|---------|
| `--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 |
| `--muon_ns_steps` | Newton-Schulz iteration steps for Muon | 5 |
| `--muon_adjust_lr` | Muon LR adjustment strategy (`original`, `match_rms_adamw`) | `match_rms_adamw` |
Optimizer identity and hyperparameters are saved in checkpoint metadata. Optimizer
states are intentionally not interchangeable: resume older MuonMix checkpoints
with `--optimizer=muon_adamw`.
### Data Loading
| Parameter | Description | Default |
@@ -141,7 +159,9 @@ nohup python scripts/tools/train.py \
--batch_per_device=4 \
--grad_accum_steps=8 \
--warmup_ratio=0.05 \
--optimizer=nora_nadamw \
--max_lr=1e-4 \
--nora_lr=5e-3 \
--max_grad_norm=1.0 \
--weight_decay=0.1 \
--window_size=2048 \
+90 -99
View File
@@ -1,115 +1,43 @@
import os
from collections.abc import Callable
from functools import partial
from typing import Any
import click
import torch
from torch import Tensor, nn, optim
from click.core import ParameterSource
from torch import optim
from astrai import setup_logging
from astrai.config import AutoRegressiveLMConfig, TrainConfig
from astrai.dataset import DatasetFactory, dpo_collate_fn, grpo_collate_fn
from astrai.model import AutoRegressiveLM
from astrai.model.components.decoder_block import DecoderBlock
from astrai.optim import OptimizerFactory
from astrai.trainer import SchedulerFactory, Trainer
from astrai.trainer.rollout import BaseRewardModel
class MuonMix(optim.Optimizer):
"""Combined Muon (matrix) + AdamW (non-matrix) optimizer."""
def __init__(
self,
model: nn.Module,
lr: float = 3e-4,
weight_decay: float = 0.1,
momentum: float = 0.95,
nesterov: bool = True,
ns_steps: int = 5,
adjust_lr_fn: str = "match_rms_adamw",
):
defaults = {
"lr": lr,
"weight_decay": weight_decay,
"momentum": momentum,
"nesterov": nesterov,
"ns_steps": ns_steps,
"adjust_lr_fn": adjust_lr_fn,
}
params = [p for p in model.parameters() if p.requires_grad]
super().__init__(params, defaults)
matrix_params: list[Tensor] = []
other_params: list[Tensor] = []
for name, param in model.named_parameters():
if not param.requires_grad:
continue
if (
param.dim() >= 2
and "norm" not in name
and "bias" not in name
and "embed" not in name
and "lm_head" not in name
):
matrix_params.append(param)
else:
other_params.append(param)
self.muon = optim.Muon(
matrix_params,
lr=lr,
weight_decay=weight_decay,
momentum=momentum,
nesterov=nesterov,
ns_steps=ns_steps,
adjust_lr_fn=adjust_lr_fn,
)
self.adamw = optim.AdamW(
[{"params": other_params, "weight_decay": 0.0}],
lr=lr,
betas=(0.9, 0.95),
fused=True,
)
self.param_groups = [*self.muon.param_groups, *self.adamw.param_groups]
@torch.no_grad()
def step(self, closure=None):
self.muon.step(closure)
self.adamw.step(closure)
def zero_grad(self, set_to_none: bool = True):
self.muon.zero_grad(set_to_none)
self.adamw.zero_grad(set_to_none)
def state_dict(self) -> dict[str, Any]:
return {
"muon": self.muon.state_dict(),
"adamw": self.adamw.state_dict(),
}
def load_state_dict(self, state_dict: dict[str, Any]):
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]
def _merge_yaml_into_kwargs(config_path: str, passed_kwargs: dict) -> dict:
"""Load YAML config, then override with explicit CLI kwargs (None excluded)."""
def _merge_yaml_into_kwargs(
config_path: str,
passed_kwargs: dict,
explicit_keys: set[str] | None = None,
) -> dict:
"""Merge Click defaults, YAML values, then explicit CLI values."""
import yaml
with open(config_path) as f:
cfg = yaml.safe_load(f)
cfg = yaml.safe_load(f) or {}
merged = {}
merged = dict(passed_kwargs)
for section in ("model", "data", "parallel", "training", "ckpt", "log"):
if section in cfg:
merged.update(cfg[section])
for key, value in passed_kwargs.items():
if value is not None:
merged[key] = value
if explicit_keys is None:
explicit_keys = set(passed_kwargs)
for key in explicit_keys:
if key in passed_kwargs:
merged[key] = passed_kwargs[key]
return merged
@@ -117,6 +45,7 @@ def _merge_yaml_into_kwargs(config_path: str, passed_kwargs: dict) -> dict:
_TRAIN_TYPE = ["seq", "sft", "dpo", "grpo", "online_grpo", "online_dpo"]
_PARALLEL = ["none", "ddp", "fsdp"]
_SCHEDULES = ["cosine", "sgdr", "wsd"]
_OPTIMIZERS = OptimizerFactory.list_registered()
_BACKENDS = ["nccl", "gloo"]
_START_METHODS = ["spawn", "fork", "forkserver"]
@@ -162,10 +91,25 @@ _START_METHODS = ["spawn", "fork", "forkserver"]
help="Fraction of total steps for LR warmup.",
)
@click.option("--max_lr", type=float, default=3e-4, help="Max learning rate.")
@click.option(
"--optimizer",
type=click.Choice(_OPTIMIZERS),
default="nora_nadamw",
help="Built-in optimizer.",
)
@click.option(
"--max_grad_norm", type=float, default=1.0, help="Max gradient norm for clipping."
)
@click.option("--weight_decay", type=float, default=0.1, help="Weight decay.")
@click.option(
"--weight_decay",
type=float,
default=0.1,
help="Weight decay for eligible optimizer parameters.",
)
@click.option("--nora_lr", type=float, default=5e-3, help="Nora learning rate.")
@click.option("--nora_beta", type=float, default=0.95, help="Nora EMA factor.")
@click.option("--nora_momentum", type=float, default=0.95, help="Nora update momentum.")
@click.option("--nora_weight_decay", type=float, default=0.0, help="Nora weight decay.")
@click.option("--muon_momentum", type=float, default=0.95, help="Muon momentum factor.")
@click.option("--muon_nesterov/--no-muon_nesterov", default=True, help="Muon Nesterov.")
@click.option("--muon_ns_steps", type=int, default=5, help="Muon Newton-Schulz steps.")
@@ -283,8 +227,14 @@ _START_METHODS = ["spawn", "fork", "forkserver"]
@click.pass_context
def train_command(ctx, config_path, dry_run, metrics, **kwargs):
"""Start model training (pretrain / SFT / DPO / GRPO)."""
kwargs["metrics"] = metrics
if config_path:
kwargs = _merge_yaml_into_kwargs(config_path, kwargs)
explicit_keys = {
key
for key in kwargs
if ctx.get_parameter_source(key) is ParameterSource.COMMANDLINE
}
kwargs = _merge_yaml_into_kwargs(config_path, kwargs, explicit_keys)
required = ["train_type", "data_root_path", "param_path"]
missing = [k for k in required if kwargs.get(k) is None]
@@ -295,7 +245,7 @@ def train_command(ctx, config_path, dry_run, metrics, **kwargs):
)
# Convert tuple back to list
kwargs["metrics"] = list(metrics)
kwargs["metrics"] = list(kwargs["metrics"])
# Remove tp_size (not yet wired)
kwargs.pop("tp_size", None)
@@ -317,6 +267,7 @@ def _print_dry_run(kwargs: dict) -> None:
("Epochs", str(kwargs.get("n_epoch", 1))),
("Batch/device", str(kwargs.get("batch_per_device", 1))),
("Grad accum", str(kwargs.get("grad_accum_steps", 1))),
("Optimizer", str(kwargs.get("optimizer", "nora_nadamw"))),
("Max LR", str(kwargs.get("max_lr", "?"))),
("Schedule", str(kwargs.get("schedule_type", "cosine"))),
("Warmup ratio", str(kwargs.get("warmup_ratio", 0.05))),
@@ -336,8 +287,10 @@ def create_model(config):
return AutoRegressiveLM(config).to(dtype=torch.bfloat16)
def create_optimizer(model, **kwargs) -> MuonMix:
return MuonMix(model, **kwargs)
def create_optimizer(
model, optimizer_name: str = "nora_nadamw", **kwargs
) -> optim.Optimizer:
return OptimizerFactory.create(optimizer_name, model, **kwargs)
def create_scheduler(
@@ -459,15 +412,51 @@ def train(
tokenizer_path=param_path,
)
optimizer_name = kwargs.pop("optimizer", "nora_nadamw")
optimizer_kwargs = {
"lr": kwargs.pop("max_lr"),
"weight_decay": kwargs.pop("weight_decay"),
"nora_lr": kwargs.pop("nora_lr", 5e-3),
"nora_beta": kwargs.pop("nora_beta", 0.95),
"nora_momentum": kwargs.pop("nora_momentum", 0.95),
"nora_weight_decay": kwargs.pop("nora_weight_decay", 0.0),
"momentum": kwargs.pop("muon_momentum", 0.95),
"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"),
}
optimizer_fn = partial(
create_optimizer,
lr=kwargs.pop("max_lr"),
weight_decay=kwargs.pop("weight_decay"),
momentum=kwargs.pop("muon_momentum"),
nesterov=kwargs.pop("muon_nesterov"),
ns_steps=kwargs.pop("muon_ns_steps"),
adjust_lr_fn=kwargs.pop("muon_adjust_lr"),
optimizer_name=optimizer_name,
**optimizer_kwargs,
)
if optimizer_name == "nora_nadamw":
optimizer_hyperparameters = {
key: optimizer_kwargs[key]
for key in (
"lr",
"weight_decay",
"nora_lr",
"nora_beta",
"nora_momentum",
"nora_weight_decay",
)
}
optimizer_hyperparameters.update(
{"nadamw_betas": [0.9, 0.999], "nadamw_eps": 1e-8, "nora_eps": 1e-10}
)
else:
optimizer_hyperparameters = {
key: optimizer_kwargs[key]
for key in (
"lr",
"weight_decay",
"momentum",
"nesterov",
"ns_steps",
"adjust_lr_fn",
)
}
total_steps = compute_total_steps(
len(dataset), n_epoch, batch_per_device, nprocs, grad_accum_steps
@@ -516,6 +505,8 @@ def train(
dataset=dataset,
optimizer_fn=optimizer_fn,
scheduler_fn=scheduler_fn,
optimizer_name=optimizer_name,
optimizer_hyperparameters=optimizer_hyperparameters,
ckpt_dir=ckpt_dir,
n_epoch=n_epoch,
batch_per_device=batch_per_device,
+253
View File
@@ -0,0 +1,253 @@
import math
from copy import deepcopy
import pytest
import torch
from torch.utils.data import TensorDataset
from astrai.config import TrainConfig
from astrai.model import AutoRegressiveLM
from astrai.model.components.linear import Linear
from astrai.model.components.lora import LoRALinear, inject_lora
from astrai.optim import (
NAdamW,
Nora,
NoraNAdamW,
OptimizerFactory,
nora_lr_scale,
partition_optimizer_parameters,
)
from astrai.trainer.schedule import SchedulerFactory
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_nora_one_step_matches_row_geometry():
param = torch.nn.Parameter(torch.tensor([[3.0, 4.0], [0.0, 2.0]]))
grad = torch.tensor([[4.0, -3.0], [1.0, 1.0]])
param.grad = grad.clone()
optimizer = Nora([param], lr=0.1, beta=0.0, momentum=0.0)
optimizer.step()
theta_hat = torch.tensor([[0.6, 0.8], [0.0, 1.0]])
tangent = grad - (grad * theta_hat).sum(dim=-1, keepdim=True) * theta_hat
direction = tangent / tangent.norm(dim=-1, keepdim=True).clamp(min=1e-10)
expected = torch.tensor([[3.0, 4.0], [0.0, 2.0]]) - 0.1 * direction
torch.testing.assert_close(param, expected)
def test_nora_handles_zero_and_pure_radial_rows():
param = torch.nn.Parameter(torch.tensor([[0.0, 0.0], [3.0, 4.0]]))
param.grad = torch.tensor([[3.0, 4.0], [6.0, 8.0]])
optimizer = Nora([param], lr=0.1, beta=0.0, momentum=0.0)
optimizer.step()
torch.testing.assert_close(param[0], torch.tensor([-0.06, -0.08]))
torch.testing.assert_close(param[1], torch.tensor([3.0, 4.0]), atol=1e-6, rtol=0)
def test_nora_lr_scale_only_increases_tall_matrices():
assert nora_lr_scale(0.1, torch.Size([4, 2])) == pytest.approx(0.1 * math.sqrt(2.0))
assert nora_lr_scale(0.1, torch.Size([2, 4])) == pytest.approx(0.1)
def test_nadamw_one_step_matches_reference_formula():
param = torch.nn.Parameter(torch.tensor([1.0, -2.0]))
grad = torch.tensor([0.5, -0.25])
param.grad = grad.clone()
lr = 0.1
beta1, beta2 = 0.9, 0.999
eps = 1e-8
optimizer = NAdamW([param], lr=lr, betas=(beta1, beta2), eps=eps, weight_decay=0.2)
optimizer.step()
m = (1 - beta1) * grad
v = (1 - beta2) * grad.square()
m_hat = (beta1 * m + (1 - beta1) * grad) / (1 - beta1)
v_hat = v / (1 - beta2)
expected = torch.tensor([1.0, -2.0]) * (1 - lr * 0.2)
expected.add_(m_hat / (v_hat.sqrt() + eps), alpha=-lr)
torch.testing.assert_close(param, expected)
@pytest.mark.parametrize("tie_word_embeddings", [False, True])
def test_parameter_partition_is_complete_disjoint_and_role_based(
tie_word_embeddings,
):
model = AutoRegressiveLM(make_tiny_config(tie_word_embeddings=tie_word_embeddings))
inject_lora(model, r=2, alpha=4, target_modules={"q_proj"})
groups = partition_optimizer_parameters(model)
all_grouped = [*groups.nora, *groups.nadamw_decay, *groups.nadamw_no_decay]
trainable = [param for param in model.parameters() if param.requires_grad]
assert len({id(param) for param in all_grouped}) == len(all_grouped)
assert {id(param) for param in all_grouped} == {id(param) for param in trainable}
assert id(model.embed_tokens.weight) in {id(p) for p in groups.nadamw_no_decay}
assert id(model.lm_head.weight) in {id(p) for p in groups.nadamw_no_decay}
nora_ids = {id(param) for param in groups.nora}
no_decay_ids = {id(param) for param in groups.nadamw_no_decay}
for name, module in model.named_modules():
if isinstance(module, LoRALinear):
assert id(module.lora_A) in no_decay_ids
assert id(module.lora_B) in no_decay_ids
elif isinstance(module, Linear) and name != "lm_head":
if module.weight.requires_grad:
assert id(module.weight) in nora_ids
@pytest.mark.parametrize(
"model_overrides",
[
{"attn_type": "gqa", "ffn_type": "mlp"},
{
"attn_type": "gqa",
"ffn_type": "moe",
"n_routed_experts": 2,
"n_shared_experts": 1,
"n_activated_experts": 1,
"topk_method": "greedy",
},
{
"attn_type": "mla",
"ffn_type": "mlp",
"kv_lora_rank": 4,
"qk_nope_head_dim": 2,
"qk_rope_head_dim": 2,
},
{
"attn_type": "mla",
"ffn_type": "moe",
"kv_lora_rank": 4,
"qk_nope_head_dim": 2,
"qk_rope_head_dim": 2,
"n_routed_experts": 2,
"n_shared_experts": 1,
"n_activated_experts": 1,
"topk_method": "greedy",
},
],
)
def test_parameter_partition_covers_all_model_structures(model_overrides):
model = AutoRegressiveLM(make_tiny_config(**model_overrides))
groups = partition_optimizer_parameters(model)
grouped = [*groups.nora, *groups.nadamw_decay, *groups.nadamw_no_decay]
trainable = [param for param in model.parameters() if param.requires_grad]
assert {id(param) for param in grouped} == {id(param) for param in trainable}
assert len(grouped) == len({id(param) for param in grouped})
def test_factory_registers_nora_default_and_legacy_muon():
assert OptimizerFactory.list_registered() == ["muon_adamw", "nora_nadamw"]
model = AutoRegressiveLM(make_tiny_config())
optimizer = OptimizerFactory.create("nora_nadamw", model, lr=3e-4)
assert isinstance(optimizer, NoraNAdamW)
def test_scheduler_preserves_nora_to_nadamw_lr_ratio():
model = AutoRegressiveLM(make_tiny_config())
optimizer = NoraNAdamW(model, lr=3e-4, nora_lr=5e-3)
scheduler = SchedulerFactory.create(
"cosine", optimizer, warmup_steps=2, lr_decay_steps=2, min_rate=0.1
)
initial_ratio = optimizer.param_groups[0]["lr"] / optimizer.param_groups[-1]["lr"]
_set_constant_grads(model, 0.1)
optimizer.step()
scheduler.step()
stepped_ratio = optimizer.param_groups[0]["lr"] / optimizer.param_groups[-1]["lr"]
assert initial_ratio == pytest.approx(5e-3 / 3e-4)
assert stepped_ratio == pytest.approx(initial_ratio)
def test_optimizer_and_scheduler_resume_matches_uninterrupted_step():
torch.manual_seed(7)
model_a = AutoRegressiveLM(make_tiny_config())
optimizer_a = NoraNAdamW(model_a, lr=3e-4, nora_lr=5e-3)
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 = NoraNAdamW(model_b, lr=3e-4, nora_lr=5e-3)
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())
def test_train_config_serializes_optimizer_metadata():
config = TrainConfig(
model_fn=lambda: torch.nn.Linear(2, 2),
strategy="seq",
dataset=TensorDataset(torch.zeros(1, 2)),
optimizer_fn=lambda model: torch.optim.AdamW(model.parameters()),
scheduler_fn=lambda optimizer: torch.optim.lr_scheduler.LambdaLR(
optimizer, lambda _: 1.0
),
optimizer_name="nora_nadamw",
optimizer_hyperparameters={"lr": 3e-4, "nora_lr": 5e-3},
)
metadata = config.to_dict()
assert metadata["optimizer_name"] == "nora_nadamw"
assert metadata["optimizer_hyperparameters"] == {
"lr": 3e-4,
"nora_lr": 5e-3,
}
def test_nora_nadamw_rejects_legacy_muon_state():
model = AutoRegressiveLM(make_tiny_config())
optimizer = NoraNAdamW(model)
with pytest.raises(ValueError, match="muon_adamw"):
optimizer.load_state_dict({"muon": {}, "adamw": {}})
def test_combined_optimizer_runs_closure_once():
model = AutoRegressiveLM(make_tiny_config())
optimizer = NoraNAdamW(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
+65
View File
@@ -0,0 +1,65 @@
import pytest
import torch
import torch.distributed as dist
from torch.distributed.fsdp import fully_shard
from torch.distributed.tensor import DTensor, Shard
from torch.nn.parallel import DistributedDataParallel as DDP
from astrai.model import AutoRegressiveLM
from astrai.optim import NoraNAdamW
from astrai.parallel.setup import find_free_port
from tests.helpers import make_tiny_config
pytestmark = pytest.mark.skipif(
torch.cuda.device_count() < 1, reason="CUDA device required"
)
def _assign_grads_and_step(model):
optimizer = NoraNAdamW(model)
for param in model.parameters():
if param.requires_grad:
param.grad = torch.ones_like(param)
optimizer.step()
return optimizer
def test_nora_nadamw_steps_after_ddp_and_fsdp2_wrapping():
torch.cuda.set_device(0)
dist.init_process_group(
"nccl",
rank=0,
world_size=1,
init_method=f"tcp://127.0.0.1:{find_free_port()}",
)
try:
ddp_model = AutoRegressiveLM(make_tiny_config()).to(
device="cuda", dtype=torch.bfloat16
)
ddp_model = DDP(ddp_model, device_ids=[0], output_device=0)
ddp_optimizer = _assign_grads_and_step(ddp_model)
assert ddp_optimizer.state_dict()["nora"]["state"]
fsdp_model = AutoRegressiveLM(make_tiny_config()).to(
device="cuda", dtype=torch.bfloat16
)
for child in fsdp_model.children():
if isinstance(child, torch.nn.ModuleList):
for submodule in child:
fully_shard(submodule, reshard_after_forward=False)
else:
fully_shard(child, reshard_after_forward=False)
fsdp_optimizer = _assign_grads_and_step(fsdp_model)
nora_params = fsdp_optimizer.nora.param_groups[0]["params"]
assert nora_params
assert all(isinstance(param, DTensor) for param in nora_params)
assert all(
all(
not isinstance(placement, Shard) or placement.dim == 0
for placement in param.placements
)
for param in nora_params
)
finally:
dist.destroy_process_group()
+62
View File
@@ -0,0 +1,62 @@
import re
from click.testing import CliRunner
from scripts.tools.train import _merge_yaml_into_kwargs, train_command
def test_yaml_overrides_click_defaults_but_not_explicit_cli(tmp_path):
config_path = tmp_path / "train.yaml"
config_path.write_text(
"training:\n"
" optimizer: nora_nadamw\n"
" max_lr: 0.0002\n"
" nora_lr: 0.004\n"
" batch_per_device: 8\n",
encoding="utf-8",
)
click_values = {
"optimizer": "nora_nadamw",
"max_lr": 3e-4,
"nora_lr": 5e-3,
"batch_per_device": 16,
}
merged = _merge_yaml_into_kwargs(
str(config_path), click_values, explicit_keys={"batch_per_device"}
)
assert merged["max_lr"] == 2e-4
assert merged["nora_lr"] == 4e-3
assert merged["batch_per_device"] == 16
def test_train_dry_run_uses_yaml_then_explicit_cli(tmp_path):
data_path = tmp_path / "data"
model_path = tmp_path / "model"
data_path.mkdir()
model_path.mkdir()
config_path = tmp_path / "train.yaml"
config_path.write_text(
"data:\n"
f" data_root_path: {data_path}\n"
"model:\n"
f" param_path: {model_path}\n"
"training:\n"
" train_type: seq\n"
" optimizer: nora_nadamw\n"
" max_lr: 0.0002\n"
" nora_lr: 0.004\n"
" batch_per_device: 8\n",
encoding="utf-8",
)
result = CliRunner().invoke(
train_command,
["--config", str(config_path), "--dry-run", "--batch_per_device", "16"],
)
assert result.exit_code == 0, result.output
assert re.search(r"Optimizer\s+: nora_nadamw", result.output)
assert re.search(r"Batch/device\s+: 16", result.output)
assert re.search(r"Max LR\s+: 0.0002", result.output)