Optimizers
Overview
Axolotl supports all optimizers supported by transformers OptimizerNames
Here is a list of optimizers supported by transformers as of v4.54.0:
adamw_torchadamw_torch_fusedadamw_torch_xlaadamw_torch_npu_fusedadamw_apex_fusedadafactoradamw_anyprecisionadamw_torch_4bitadamw_torch_8bitademamixsgdadagradadamw_bnb_8bitadamw_8bit# alias for adamw_bnb_8bitademamix_8bitlion_8bitlion_32bitpaged_adamw_32bitpaged_adamw_8bitpaged_ademamix_32bitpaged_ademamix_8bitpaged_lion_32bitpaged_lion_8bitrmsproprmsprop_bnbrmsprop_bnb_8bitrmsprop_bnb_32bitgalore_adamwgalore_adamw_8bitgalore_adafactorgalore_adamw_layerwisegalore_adamw_8bit_layerwisegalore_adafactor_layerwiselomoadalomogrokadamwschedule_free_radamschedule_free_adamwschedule_free_sgdapollo_adamwapollo_adamw_layerwisestable_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
optimizer: optimi_adamwao_adamw_4bit
Deprecated: Please use adamw_torch_4bit.
ao_adamw_8bit
Deprecated: Please use adamw_torch_8bit.
ao_adamw_fp8
optimizer: ao_adamw_fp8adopt_adamw
GitHub: https://github.com/iShohei220/adopt Paper: https://arxiv.org/abs/2411.02853
optimizer: adopt_adamwcame_pytorch
GitHub: https://github.com/yangluo7/CAME/tree/master Paper: https://arxiv.org/abs/2307.02047
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-16muon
Blog: https://kellerjordan.github.io/posts/muon/ Paper: https://arxiv.org/abs/2502.16982v1
optimizer: muondion
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 Paper: https://arxiv.org/pdf/2504.05295 Note: Implementation written for PyTorch 2.7+ for DTensor
optimizer: dion
dion_lr: 0.01
dion_momentum: 0.95
lr: 0.00001 # learning rate for embeddings and parameters that fallback to AdamWsinkgd
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
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 updatesWidth 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.
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 >10xSpectral 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.
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.
optim_args:
sinkgd_fused_kernel: trueWeight-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.
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 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 Paper: 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.
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