diff --git a/README.md b/README.md index e391e6e..b98fd50 100644 --- a/README.md +++ b/README.md @@ -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 \ @@ -256,4 +258,4 @@ This project is licensed under the [GPL-3.0 License](LICENSE).
A lightweight Transformer framework designed for both high performance and ease of use. -
\ No newline at end of file + diff --git a/astrai/config/train_config.py b/astrai/config/train_config.py index fadc973..d5b9403 100644 --- a/astrai/config/train_config.py +++ b/astrai/config/train_config.py @@ -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 diff --git a/astrai/optim/__init__.py b/astrai/optim/__init__.py new file mode 100644 index 0000000..7fcc7aa --- /dev/null +++ b/astrai/optim/__init__.py @@ -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", +] diff --git a/astrai/optim/muon_mix.py b/astrai/optim/muon_mix.py new file mode 100644 index 0000000..c48b287 --- /dev/null +++ b/astrai/optim/muon_mix.py @@ -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] diff --git a/astrai/optim/nora_nadamw.py b/astrai/optim/nora_nadamw.py new file mode 100644 index 0000000..2a4512f --- /dev/null +++ b/astrai/optim/nora_nadamw.py @@ -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() diff --git a/docs/README-zh-CN.md b/docs/README-zh-CN.md index df638dd..ad5a967 100644 --- a/docs/README-zh-CN.md +++ b/docs/README-zh-CN.md @@ -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 \ @@ -262,4 +264,4 @@ SSE 流式格式、错误码和统计端点详见[推理文档](guides/inference
专为高性能与易用性设计的轻量级 Transformer 框架。 -
\ No newline at end of file + diff --git a/docs/guides/params.md b/docs/guides/params.md index b4258d5..09bef2c 100644 --- a/docs/guides/params.md +++ b/docs/guides/params.md @@ -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 \ diff --git a/scripts/tools/train.py b/scripts/tools/train.py index 8a238f7..8010e63 100644 --- a/scripts/tools/train.py +++ b/scripts/tools/train.py @@ -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, diff --git a/tests/optim/test_nora_nadamw.py b/tests/optim/test_nora_nadamw.py new file mode 100644 index 0000000..10f0381 --- /dev/null +++ b/tests/optim/test_nora_nadamw.py @@ -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 diff --git a/tests/optim/test_optimizer_distributed.py b/tests/optim/test_optimizer_distributed.py new file mode 100644 index 0000000..d7ae652 --- /dev/null +++ b/tests/optim/test_optimizer_distributed.py @@ -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() diff --git a/tests/test_train_cli.py b/tests/test_train_cli.py new file mode 100644 index 0000000..2f7574f --- /dev/null +++ b/tests/test_train_cli.py @@ -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)