78ec5d9290
ci-cd / build-axolotl-uv (<nil>, 130, 13.0.0, linux/amd64,linux/arm64, 3.12, 2.11.0) (push) Has been cancelled
ci-cd / build-axolotl-uv (<nil>, 130, 13.0.0, true, linux/amd64,linux/arm64, 3.12, 2.12.0) (push) Has been cancelled
ci-cd / build-axolotl-cloud-uv (<nil>, 130, 13.0.0, linux/amd64,linux/arm64, 3.12, 2.11.0) (push) Has been cancelled
ci-cd / build-axolotl-cloud-uv (<nil>, 130, 13.0.0, true, linux/amd64,linux/arm64, 3.12, 2.12.0) (push) Has been cancelled
ci-cd / build-axolotl-cloud-no-tmux-uv (<nil>, 130, 13.0.0, linux/amd64,linux/arm64, 3.12, 2.11.0) (push) Has been cancelled
ci-cd / build-axolotl-cloud-no-tmux-uv (<nil>, 130, 13.0.0, true, linux/amd64,linux/arm64, 3.12, 2.12.0) (push) Has been cancelled
ci-cd-base / build-base-uv (130, 13.0.0, , Dockerfile-uv-base, linux/amd64,linux/arm64, 3.12, 2.12.0, 9.0 10.0 10.3 12.0+PTX) (push) Has been cancelled
ci-cd-base / build-base-uv (130, 13.0.0, , Dockerfile-uv-base, linux/amd64,linux/arm64, 3.12, 2.12.1, 9.0 10.0 10.3 12.0+PTX) (push) Has been cancelled
ci-cd-base / build-base-uv (130, 13.0.0, , Dockerfile-uv-base, linux/amd64,linux/arm64, 3.12, 2.13.0, 9.0 10.0 10.3 12.0+PTX) (push) Has been cancelled
ci-cd-base / build-base-uv (132, 13.2.1, , Dockerfile-uv-base, linux/amd64,linux/arm64, 3.12, 2.13.0, 9.0 10.0 10.3 12.0+PTX, https://download.pytorch.org/whl/cu132) (push) Has been cancelled
ci-cd-base / build-base-uv (130, 13.0.0, , Dockerfile-uv-base, linux/amd64,linux/arm64, 3.12, 2.11.0, 9.0 10.0 10.3 12.0+PTX) (push) Has been cancelled
Tests / PyTest (3.12, 2.12.1) (push) Has been cancelled
Tests / PyTest (3.12, 2.13.0) (push) Has been cancelled
docker-e2e-tests / gate-skip-e2e (push) Has been cancelled
docker-e2e-tests / docker-e2e-tests-1st (<nil>, 130, 13.0.0, 1, 3.12, 2.12.1) (push) Has been cancelled
docker-e2e-tests / docker-e2e-tests (<nil>, 130, 13.0.0, 1, 3.12, 2.11.0) (push) Has been cancelled
docker-e2e-tests / docker-e2e-kernel-tests (<nil>, 130, 13.0.0, 1, 3.12, 2.11.0) (push) Has been cancelled
docker-e2e-tests / docker-e2e-kernel-tests (<nil>, 130, 13.0.0, 1, 3.12, 2.12.1) (push) Has been cancelled
docker-e2e-tests / docker-e2e-cleanup (<nil>, 130, 13.0.0, 1, 3.12, 2.12.1) (push) Has been cancelled
Publish Docs / build-deploy (push) Has been cancelled
Tests / PyTest from Source Dist (3.12, 2.11.0) (push) Has been cancelled
Tests / PyTest from Source Dist (3.12, 2.12.1) (push) Has been cancelled
Tests / PyTest from Source Dist (3.12, 2.13.0) (push) Has been cancelled
Tests / pre-commit (push) Has been cancelled
Tests / Prefetch S3 once to prime the CDN cache (push) Has been cancelled
Tests / PyTest (3.12, 2.11.0) (push) Has been cancelled
281 lines
10 KiB
Plaintext
281 lines
10 KiB
Plaintext
---
|
||
title: Optimizers
|
||
description: Configuring optimizers
|
||
---
|
||
|
||
## Overview
|
||
|
||
Axolotl supports all optimizers supported by [transformers OptimizerNames](https://github.com/huggingface/transformers/blob/51f94ea06d19a6308c61bbb4dc97c40aabd12bad/src/transformers/training_args.py#L142-L187)
|
||
|
||
Here is a list of optimizers supported by transformers as of `v4.54.0`:
|
||
|
||
- `adamw_torch`
|
||
- `adamw_torch_fused`
|
||
- `adamw_torch_xla`
|
||
- `adamw_torch_npu_fused`
|
||
- `adamw_apex_fused`
|
||
- `adafactor`
|
||
- `adamw_anyprecision`
|
||
- `adamw_torch_4bit`
|
||
- `adamw_torch_8bit`
|
||
- `ademamix`
|
||
- `sgd`
|
||
- `adagrad`
|
||
- `adamw_bnb_8bit`
|
||
- `adamw_8bit` # alias for adamw_bnb_8bit
|
||
- `ademamix_8bit`
|
||
- `lion_8bit`
|
||
- `lion_32bit`
|
||
- `paged_adamw_32bit`
|
||
- `paged_adamw_8bit`
|
||
- `paged_ademamix_32bit`
|
||
- `paged_ademamix_8bit`
|
||
- `paged_lion_32bit`
|
||
- `paged_lion_8bit`
|
||
- `rmsprop`
|
||
- `rmsprop_bnb`
|
||
- `rmsprop_bnb_8bit`
|
||
- `rmsprop_bnb_32bit`
|
||
- `galore_adamw`
|
||
- `galore_adamw_8bit`
|
||
- `galore_adafactor`
|
||
- `galore_adamw_layerwise`
|
||
- `galore_adamw_8bit_layerwise`
|
||
- `galore_adafactor_layerwise`
|
||
- `lomo`
|
||
- `adalomo`
|
||
- `grokadamw`
|
||
- `schedule_free_radam`
|
||
- `schedule_free_adamw`
|
||
- `schedule_free_sgd`
|
||
- `apollo_adamw`
|
||
- `apollo_adamw_layerwise`
|
||
- `stable_adamw`
|
||
|
||
|
||
## Custom Optimizers
|
||
|
||
Enable custom optimizers by passing a string to the `optimizer` argument. Each optimizer will receive beta and epsilon args, however, some may accept additional args which are detailed below.
|
||
|
||
### optimi_adamw
|
||
|
||
```yaml
|
||
optimizer: optimi_adamw
|
||
```
|
||
|
||
### ao_adamw_4bit
|
||
|
||
Deprecated: Please use `adamw_torch_4bit`.
|
||
|
||
### ao_adamw_8bit
|
||
|
||
Deprecated: Please use `adamw_torch_8bit`.
|
||
|
||
### ao_adamw_fp8
|
||
|
||
|
||
```yaml
|
||
optimizer: ao_adamw_fp8
|
||
```
|
||
|
||
### adopt_adamw
|
||
|
||
GitHub: [https://github.com/iShohei220/adopt](https://github.com/iShohei220/adopt)
|
||
Paper: [https://arxiv.org/abs/2411.02853](https://arxiv.org/abs/2411.02853)
|
||
|
||
```yaml
|
||
optimizer: adopt_adamw
|
||
```
|
||
|
||
### came_pytorch
|
||
|
||
GitHub: [https://github.com/yangluo7/CAME/tree/master](https://github.com/yangluo7/CAME/tree/master)
|
||
Paper: [https://arxiv.org/abs/2307.02047](https://arxiv.org/abs/2307.02047)
|
||
|
||
```yaml
|
||
optimizer: came_pytorch
|
||
|
||
# optional args (defaults below)
|
||
adam_beta1: 0.9
|
||
adam_beta2: 0.999
|
||
adam_beta3: 0.9999
|
||
adam_epsilon: 1e-30
|
||
adam_epsilon2: 1e-16
|
||
```
|
||
|
||
### muon
|
||
|
||
Blog: [https://kellerjordan.github.io/posts/muon/](https://kellerjordan.github.io/posts/muon/)
|
||
Paper: [https://arxiv.org/abs/2502.16982v1](https://arxiv.org/abs/2502.16982v1)
|
||
|
||
```yaml
|
||
optimizer: muon
|
||
```
|
||
|
||
### dion
|
||
|
||
Microsoft's Dion (DIstributed OrthoNormalization) optimizer is a scalable and communication-efficient
|
||
orthonormalizing optimizer that uses low-rank approximations to reduce gradient communication.
|
||
|
||
GitHub: [https://github.com/microsoft/dion](https://github.com/microsoft/dion)
|
||
Paper: [https://arxiv.org/pdf/2504.05295](https://arxiv.org/pdf/2504.05295)
|
||
Note: Implementation written for PyTorch 2.7+ for DTensor
|
||
|
||
```yaml
|
||
optimizer: dion
|
||
dion_lr: 0.01
|
||
dion_momentum: 0.95
|
||
lr: 0.00001 # learning rate for embeddings and parameters that fallback to AdamW
|
||
```
|
||
|
||
### sinkgd
|
||
|
||
SinkGD (Gradient Multi-Normalization) is a **stateless** optimizer: 2D linear weight
|
||
matrices are updated via the SR-Sinkhorn procedure (alternating row/column L2
|
||
normalization of the raw gradient, no momentum or variance state), while embeddings,
|
||
the LM head, and 1D params (norms/biases) fall back to AdamW. The AdamW fallback uses
|
||
torchao's 8-bit optimizer base, so optimizer-state memory is very small (~87% less than
|
||
8-bit AdamW on an 8B full finetune, since the ~87% of params that are 2D linear carry
|
||
zero optimizer state).
|
||
|
||
Requires PyTorch >= 2.5.1 (relies on torchao's low-bit optimizer base and `torch.compile`).
|
||
|
||
Paper: [https://arxiv.org/abs/2502.06742](https://arxiv.org/abs/2502.06742)
|
||
|
||
```yaml
|
||
optimizer: sinkgd
|
||
learning_rate: 0.001
|
||
optim_args:
|
||
sinkhorn_iters: 5 # number of SR-Sinkhorn row/column normalization iterations
|
||
sinkgd_lr_scale: 0.05 # α scale applied to linear-layer updates
|
||
```
|
||
|
||
#### Width transfer (`sinkgd_base_width`) and spectral normalization
|
||
|
||
Two optional, default-off knobs help the learning rate transfer across model width and
|
||
condition the update. `sinkgd_base_width` and `sinkgd_spectral_target: muon` are **mutually
|
||
exclusive** — both correct for width, so enabling both double-counts (validated and rejected
|
||
at construction). Spectral norm at the default `unit` target combines fine with
|
||
`sinkgd_base_width`, since it conditions the update without touching width scaling.
|
||
|
||
**Width-aware scaling (`sinkgd_base_width`).** SinkGD's update is *Adam-class*: its per-layer
|
||
learning rate should scale as `1/d_in`. Set `sinkgd_base_width` to the hidden size you tuned
|
||
`sinkgd_lr_scale` on, and each 2D-linear update is scaled by
|
||
`alpha_eff = sinkgd_lr_scale * (base_width / d_in) ** sinkgd_lr_width_exponent`
|
||
(`d_in` is the layer's input dim — the shared input of a fused QKV / gate-up matrix). With
|
||
`sinkgd_base_width` unset, behavior is identical to before (plain scalar `sinkgd_lr_scale`),
|
||
so existing configs are unchanged.
|
||
|
||
```yaml
|
||
optim_args:
|
||
sinkgd_lr_scale: 0.05
|
||
sinkgd_base_width: 2048 # d_in you tuned sinkgd_lr_scale at; enables 1/d_in transfer
|
||
sinkgd_lr_width_exponent: 1.0 # 1.0 = pure 1/d_in (default); tune only if width varies >10x
|
||
```
|
||
|
||
**Spectral normalization (`sinkgd_spectral_norm`).** Rescales each update to a target operator
|
||
norm via a cheap warm-started power iteration (~2 matvecs/layer, +1–3% step time, one
|
||
`O(d_in)` state vector per matrix). With `sinkgd_spectral_target: muon` it pins the operator
|
||
norm to `sqrt(d_out/d_in)`, which makes SinkGD *Muon-class* and both transfers across width and
|
||
trained to lower loss than plain `1/d_in` in width sweeps — so use it **instead of**
|
||
`sinkgd_base_width`, not with it. `unit` (the default target) pins a width-independent norm and
|
||
acts as a pure conditioning stabilizer on top of `1/d_in`.
|
||
|
||
```yaml
|
||
optim_args:
|
||
sinkgd_lr_scale: 0.05
|
||
sinkgd_spectral_norm: true
|
||
sinkgd_spectral_target: muon # spectral norm owns width transfer (leave sinkgd_base_width unset)
|
||
sinkgd_spectral_norm_iters: 1 # power-iteration steps per update (warm-started; 1–2 is enough)
|
||
```
|
||
|
||
Distributed (`fsdp_version: 2`): width-aware scaling and spectral norm both work under FSDP2/TP
|
||
sharding. On replicated / expert-sharded weights everything runs locally with no extra
|
||
communication. On a matrix-dim-sharded weight (the common FSDP2 row-sharded case) the spectral
|
||
norm's power iteration adds **one small vector all-reduce** (length `d_in` or `d_out`) per
|
||
iteration over the `dp_shard` group — the same shape and group as SR-Sinkhorn's existing
|
||
norm-vector reduce, and the matrix itself is never gathered. The persisted power-iteration
|
||
vector round-trips through FSDP2 checkpoints.
|
||
|
||
#### Fused Triton kernels (`sinkgd_fused_kernel`)
|
||
|
||
`sinkgd_fused_kernel: true` replaces the compiled update with fused Triton kernels (one
|
||
kernel per SR-Sinkhorn iteration; the column scale, power iteration, and weight update fold
|
||
into the same passes). Works for all SinkGD variants (plain, spectral norm, MD sphere) on
|
||
single-device and FSDP2 rows-sharded weights — same all-reduce count as the compiled path,
|
||
with an automatic tall/wide grid layout so both full matrices and heavily-sharded wide-short
|
||
local shards stay fully occupied. Measured on B200: 1.4–1.75x on the optimizer step
|
||
single-GPU, 1.0–1.8x on 2-rank FSDP2, all regimes ≥1x with the wide layout. Numerics are
|
||
equivalent at bf16 rounding scale but not byte-identical to the compiled path, so the flag
|
||
is off by default. Falls back to the compiled path for cols-sharded (TP) weights,
|
||
`bf16_stochastic_round`, or when Triton is unavailable. Spectral-norm / MD-sphere updates
|
||
that run without a shard group to amortize against (single device, or replicated /
|
||
expert-sharded weights under FSDP2) also fall back when the matrix has fewer than
|
||
`sinkgd_fused_min_numel` elements (default `2^25`) — at those sizes the epilogue's extra
|
||
kernel launches make the compiled path faster.
|
||
|
||
```yaml
|
||
optim_args:
|
||
sinkgd_fused_kernel: true
|
||
```
|
||
|
||
#### Weight-sphere variant (`sinkgd_md_sphere`, experimental)
|
||
|
||
`sinkgd_md_sphere: true` switches to an experimental magnitude–direction variant: each
|
||
SinkGD-routed 2D weight is held on a fixed Frobenius sphere (`||W||_F` anchored at enable
|
||
time), the SR-Sinkhorn update is spectral-normalized to unit operator norm, applied, and the
|
||
weight is reprojected onto the sphere (no learnable gains — SinkGD's row/column balancing makes
|
||
them redundant). In width/depth sweeps this bounds the deepest-layer activation growth most
|
||
tightly of any variant, but its optimal LR is **width-dependent** (`lr_opt ∝ d_model**-0.6`)
|
||
and its optimum is narrow, so it needs per-scale LR tuning and is **off by default**. Prefer the
|
||
plain spectral-norm path (above) unless you specifically want the tightest activation bound.
|
||
It adds a per-matrix scalar all-reduce (the sphere's global Frobenius norm) on top of the
|
||
spectral all-reduce under FSDP2; the sphere radius and power-iteration vector round-trip.
|
||
|
||
```yaml
|
||
optim_args:
|
||
sinkgd_md_sphere: true
|
||
sinkgd_lr_scale: 0.05
|
||
# tune learning_rate per width: lr_opt ~ d_model**-0.6 (e.g. ~0.03 at d=2048)
|
||
```
|
||
|
||
### q_galore_adamw8bit
|
||
|
||
Q-GaLore extends [GaLore](https://arxiv.org/abs/2403.03507) with two extra ideas:
|
||
an INT4-quantized projection matrix and an adaptive SVD scheduler that skips
|
||
re-projection when a layer's gradient subspace stabilizes. Both are wired up in
|
||
axolotl. The third Q-GaLore trick — INT8 weight wrapping — is not yet
|
||
implemented and is tracked as a follow-up.
|
||
|
||
GitHub: [https://github.com/VITA-Group/Q-GaLore](https://github.com/VITA-Group/Q-GaLore)
|
||
Paper: [https://arxiv.org/abs/2407.08296](https://arxiv.org/abs/2407.08296)
|
||
|
||
Install: `pip install axolotl[qgalore]`
|
||
|
||
This optimizer is for **full fine-tuning**. It is incompatible with `adapter`
|
||
(LoRA/QLoRA), `load_in_8bit`, and `load_in_4bit`. DeepSpeed is currently gated
|
||
off; FSDP requires `fsdp_version: 2` with `use_orig_params: true`.
|
||
|
||
```yaml
|
||
optimizer: q_galore_adamw8bit
|
||
bf16: true
|
||
|
||
# which parameter substrings get the low-rank projection
|
||
# (defaults to ["attn", "mlp"] if unset — matches the reference impl)
|
||
optim_target_modules:
|
||
- attn
|
||
- mlp
|
||
|
||
# Q-GaLore hyperparameters (defaults shown)
|
||
qgalore_rank: 256
|
||
qgalore_update_proj_gap: 200 # max steps between SVD refreshes
|
||
qgalore_scale: 0.25
|
||
qgalore_proj_type: std
|
||
qgalore_proj_quant: true # INT-quantize the projection matrix P
|
||
qgalore_proj_bits: 4 # bitwidth for P
|
||
qgalore_proj_group_size: 256 # must divide P's last dim evenly
|
||
qgalore_cos_threshold: 0.4 # skip SVD if P_t is this similar to P_{t-1}
|
||
qgalore_gamma_proj: 2 # grow update_proj_gap by this factor when stable
|
||
qgalore_queue_size: 5
|
||
```
|