From 04899a2b159ff293e9eacebf2f2608fb78d28fd0 Mon Sep 17 00:00:00 2001
From: QueenAmish <2591836946@qq.com>
Date: Fri, 31 Jul 2026 23:16:39 +0800
Subject: [PATCH] Make Nora+NAdamW the default optimizer
---
README.md | 4 +-
astrai/config/train_config.py | 4 +
astrai/optim/__init__.py | 36 ++
astrai/optim/muon_mix.py | 91 ++++++
astrai/optim/nora_nadamw.py | 379 ++++++++++++++++++++++
docs/README-zh-CN.md | 4 +-
docs/guides/params.md | 28 +-
scripts/tools/train.py | 189 +++++------
tests/optim/test_nora_nadamw.py | 253 +++++++++++++++
tests/optim/test_optimizer_distributed.py | 65 ++++
tests/test_train_cli.py | 62 ++++
11 files changed, 1010 insertions(+), 105 deletions(-)
create mode 100644 astrai/optim/__init__.py
create mode 100644 astrai/optim/muon_mix.py
create mode 100644 astrai/optim/nora_nadamw.py
create mode 100644 tests/optim/test_nora_nadamw.py
create mode 100644 tests/optim/test_optimizer_distributed.py
create mode 100644 tests/test_train_cli.py
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)