refactor: keep muon_adamw as default optimizer and drop nora docs
- revert CLI/create_optimizer/display defaults to muon_adamw - revert README, README-zh-CN, params.md to pre-merge state
This commit is contained in:
@@ -101,9 +101,7 @@ nohup python scripts/tools/train.py \
|
|||||||
--batch_per_device=4 \
|
--batch_per_device=4 \
|
||||||
--grad_accum_steps=8 \
|
--grad_accum_steps=8 \
|
||||||
--warmup_ratio=0.05 \
|
--warmup_ratio=0.05 \
|
||||||
--optimizer=nora_nadamw \
|
|
||||||
--max_lr=1e-4 \
|
--max_lr=1e-4 \
|
||||||
--nora_lr=5e-3 \
|
|
||||||
--max_grad_norm=1.0 \
|
--max_grad_norm=1.0 \
|
||||||
--weight_decay=0.1 \
|
--weight_decay=0.1 \
|
||||||
--window_size=2048 \
|
--window_size=2048 \
|
||||||
|
|||||||
@@ -107,9 +107,7 @@ nohup python scripts/tools/train.py \
|
|||||||
--batch_per_device=4 \
|
--batch_per_device=4 \
|
||||||
--grad_accum_steps=8 \
|
--grad_accum_steps=8 \
|
||||||
--warmup_ratio=0.05 \
|
--warmup_ratio=0.05 \
|
||||||
--optimizer=nora_nadamw \
|
|
||||||
--max_lr=1e-4 \
|
--max_lr=1e-4 \
|
||||||
--nora_lr=5e-3 \
|
|
||||||
--max_grad_norm=1.0 \
|
--max_grad_norm=1.0 \
|
||||||
--weight_decay=0.1 \
|
--weight_decay=0.1 \
|
||||||
--window_size=2048 \
|
--window_size=2048 \
|
||||||
|
|||||||
+18
-19
@@ -25,35 +25,36 @@
|
|||||||
| Parameter | Description | Default |
|
| Parameter | Description | Default |
|
||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
|
| `--warmup_ratio` | Fraction of total steps used for LR warmup | 0.05 |
|
||||||
| `--max_lr` | NAdamW learning rate; schedulers scale every optimizer group proportionally | 3e-4 |
|
| `--max_lr` | Maximum learning rate (cosine decay after warmup) | 3e-4 |
|
||||||
| `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | 1.0 |
|
| `--max_grad_norm` | Maximum gradient norm for clipping (None disables) | 1.0 |
|
||||||
|
|
||||||
### Optimizer
|
### Optimizer
|
||||||
|
|
||||||
The default `nora_nadamw` optimizer sends internal `Linear.weight` matrices to
|
The default `muon_adamw` optimizer sends matrix parameters through **Muon** and
|
||||||
**Nora** and embeddings, the LM head, norms, biases, LoRA factors, and fallback
|
non-matrix parameters through **AdamW** (`fused=True`).
|
||||||
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 |
|
| Parameter | Description | Default |
|
||||||
|-----------|-------------|---------|
|
|-----------|-------------|---------|
|
||||||
|
| `--optimizer` | Built-in optimizer (`muon_adamw`, `nora_nadamw`) | `muon_adamw` |
|
||||||
|
| `--weight_decay` | Weight decay (applied to Muon matrix params; non-matrix use 0) | 0.1 |
|
||||||
| `--muon_momentum` | Muon momentum factor | 0.95 |
|
| `--muon_momentum` | Muon momentum factor | 0.95 |
|
||||||
| `--muon_nesterov` | Enable Nesterov momentum for Muon | True |
|
| `--muon_nesterov` | Enable Nesterov momentum for Muon | True |
|
||||||
| `--muon_ns_steps` | Newton-Schulz iteration steps for Muon | 5 |
|
| `--muon_ns_steps` | Newton-Schulz iteration steps for Muon | 5 |
|
||||||
| `--muon_adjust_lr` | Muon LR adjustment strategy (`original`, `match_rms_adamw`) | `match_rms_adamw` |
|
| `--muon_adjust_lr` | Muon LR adjustment strategy (`original`, `match_rms_adamw`) | `match_rms_adamw` |
|
||||||
|
|
||||||
|
`nora_nadamw` routes 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 |
|
||||||
|
|-----------|-------------|---------|
|
||||||
|
| `--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 |
|
||||||
|
|
||||||
Optimizer identity and hyperparameters are saved in checkpoint metadata. Optimizer
|
Optimizer identity and hyperparameters are saved in checkpoint metadata. Optimizer
|
||||||
states are intentionally not interchangeable: resume older MuonMix checkpoints
|
states are intentionally not interchangeable: resume older MuonMix checkpoints
|
||||||
with `--optimizer=muon_adamw`.
|
with `--optimizer=muon_adamw`.
|
||||||
@@ -159,9 +160,7 @@ nohup python scripts/tools/train.py \
|
|||||||
--batch_per_device=4 \
|
--batch_per_device=4 \
|
||||||
--grad_accum_steps=8 \
|
--grad_accum_steps=8 \
|
||||||
--warmup_ratio=0.05 \
|
--warmup_ratio=0.05 \
|
||||||
--optimizer=nora_nadamw \
|
|
||||||
--max_lr=1e-4 \
|
--max_lr=1e-4 \
|
||||||
--nora_lr=5e-3 \
|
|
||||||
--max_grad_norm=1.0 \
|
--max_grad_norm=1.0 \
|
||||||
--weight_decay=0.1 \
|
--weight_decay=0.1 \
|
||||||
--window_size=2048 \
|
--window_size=2048 \
|
||||||
|
|||||||
@@ -94,7 +94,7 @@ _START_METHODS = ["spawn", "fork", "forkserver"]
|
|||||||
@click.option(
|
@click.option(
|
||||||
"--optimizer",
|
"--optimizer",
|
||||||
type=click.Choice(_OPTIMIZERS),
|
type=click.Choice(_OPTIMIZERS),
|
||||||
default="nora_nadamw",
|
default="muon_adamw",
|
||||||
help="Built-in optimizer.",
|
help="Built-in optimizer.",
|
||||||
)
|
)
|
||||||
@click.option(
|
@click.option(
|
||||||
@@ -267,7 +267,7 @@ def _print_dry_run(kwargs: dict) -> None:
|
|||||||
("Epochs", str(kwargs.get("n_epoch", 1))),
|
("Epochs", str(kwargs.get("n_epoch", 1))),
|
||||||
("Batch/device", str(kwargs.get("batch_per_device", 1))),
|
("Batch/device", str(kwargs.get("batch_per_device", 1))),
|
||||||
("Grad accum", str(kwargs.get("grad_accum_steps", 1))),
|
("Grad accum", str(kwargs.get("grad_accum_steps", 1))),
|
||||||
("Optimizer", str(kwargs.get("optimizer", "nora_nadamw"))),
|
("Optimizer", str(kwargs.get("optimizer", "muon_adamw"))),
|
||||||
("Max LR", str(kwargs.get("max_lr", "?"))),
|
("Max LR", str(kwargs.get("max_lr", "?"))),
|
||||||
("Schedule", str(kwargs.get("schedule_type", "cosine"))),
|
("Schedule", str(kwargs.get("schedule_type", "cosine"))),
|
||||||
("Warmup ratio", str(kwargs.get("warmup_ratio", 0.05))),
|
("Warmup ratio", str(kwargs.get("warmup_ratio", 0.05))),
|
||||||
@@ -288,7 +288,7 @@ def create_model(config):
|
|||||||
|
|
||||||
|
|
||||||
def create_optimizer(
|
def create_optimizer(
|
||||||
model, optimizer_name: str = "nora_nadamw", **kwargs
|
model, optimizer_name: str = "muon_adamw", **kwargs
|
||||||
) -> optim.Optimizer:
|
) -> optim.Optimizer:
|
||||||
return OptimizerFactory.create(optimizer_name, model, **kwargs)
|
return OptimizerFactory.create(optimizer_name, model, **kwargs)
|
||||||
|
|
||||||
@@ -412,7 +412,7 @@ def train(
|
|||||||
tokenizer_path=param_path,
|
tokenizer_path=param_path,
|
||||||
)
|
)
|
||||||
|
|
||||||
optimizer_name = kwargs.pop("optimizer", "nora_nadamw")
|
optimizer_name = kwargs.pop("optimizer", "muon_adamw")
|
||||||
optimizer_kwargs = {
|
optimizer_kwargs = {
|
||||||
"lr": kwargs.pop("max_lr"),
|
"lr": kwargs.pop("max_lr"),
|
||||||
"weight_decay": kwargs.pop("weight_decay"),
|
"weight_decay": kwargs.pop("weight_decay"),
|
||||||
|
|||||||
Reference in New Issue
Block a user