Skip to content
pykeioPublic

About

an optimizer for neural networks

Resources

Stars

9 stars

Watchers

1 watching

Forks

Latest commit

 

History

27 Commits

Folders and files

Repository files navigation

THORN 🌹

THORN is an optimizer for PyTorch.

THORN is an extension of Muon, which is quickly replacing AdamW in the language model space. THORN itself was used to train Earshot, a tiny voice activity detection model. THORN with minimal tuning provided a +2% validation accuracy boost over hand-tuned AdamW and made Earshot the most accurate VAD we tested in spite of its small size.

THORN works on any model, but it's most effective for models with lots of convolutions/linear layers. Transformer models will see the largest gains.

It won't give the best possible results, but you can often just reuse AdamW's same LR/betas/weight decay with THORN, making it effectively a free accuracy boost:

Results pretraining a ~300M Qwen3-based causal language model on FineWeb-Edu between AdamW, Muon, THORN, and THORN (decouple_md=True). $\gamma=10^{-3}$ (constant), $\beta_1=0.9$, $\beta_2=0.95$, $\lambda=0.1$ (on all but RMSNorm weights) for all optimizers.

⚙️ Setup details

We don't have the resources to do a proper sweep, so the AdamW parameters were chosen purely based on ~vibes~ and all other optimizers adopted them for fair comparison. The gap between AdamW and other optimizers would almost certainly be larger with more careful tuning. NorMuon was also tested, but it was within 3% of Muon the whole run, so it was excluded from the graphs as we felt something was wrong there.

Muon is the vanilla Muon from KellerJordan/Muon patched with Moonlight scaling; that is, the line update *= max(1, grad.size(-2) / grad.size(-1))**0.5 was replaced with update *= 0.2 * (max(1, *grad.shape[-2:]) ** 0.5), matching THORN's default scaling_mode='moonlight' behavior and allowing both to reuse AdamW's learning rate.

These are the exact configurations used for each optimizer. matrix is a list of all matrix parameters (excluding the embedding & LM head); embed_head is the embedding layer & LM head; vector is everything else (RMSNorm weights).

optim = AdamW([ # AdamW
	{'params': matrix + embed_head, 'weight_decay': 0.1},
	{'params': vector, 'weight_decay': 0.0}
], lr=1e-3, betas=(0.9, 0.95), eps=1e-8, fused=True)
optim = SingleDeviceMuonWithAuxAdam([ # Muon
	{'use_muon': True, 'params': matrix, 'lr': 1e-3, 'momentum': 0.9, 'weight_decay': 0.1},
	{'use_muon': False, 'params': embed_head, 'lr': 1e-3, 'betas': (0.9, 0.95), 'weight_decay': 0.1},
	{'use_muon': False, 'params': vector, 'lr': 1e-3, 'betas': (0.9, 0.95), 'weight_decay': 0.0},
])
optim = SingleDeviceNorMuonWithAuxAdam([ # NorMuon
	{'use_muon': True, 'params': matrix, 'lr': 1e-3, 'momentum': 0.9, 'beta2': 0.95, 'weight_decay': 0.1},
	{'use_muon': False, 'params': embed_head, 'lr': 1e-3, 'betas': (0.9, 0.95), 'weight_decay': 0.1},
	{'use_muon': False, 'params': vector, 'lr': 1e-3, 'betas': (0.9, 0.95), 'weight_decay': 0.0},
])
optim = THORN([ # THORN
	{'orthogonalize': True, 'params': matrix, 'lr': 1e-3, 'betas': (0.9, 0.95), 'weight_decay': 0.1},
	{'orthogonalize': False, 'params': embed_head, 'lr': 1e-3, 'betas': (0.9, 0.95), 'weight_decay': 0.1},
	{'orthogonalize': False, 'params': vector, 'lr': 1e-3, 'betas': (0.9, 0.95), 'weight_decay': 0.0},
])
optim = THORN([ # THORN-MD
	{'orthogonalize': True, 'params': matrix, 'lr': 8e-3, 'betas': (0.9, 0.95), 'weight_decay': 0.0, 'decouple_md': True},
	{'orthogonalize': False, 'params': embed_head, 'lr': 1e-3, 'betas': (0.9, 0.95), 'weight_decay': 0.1},
	{'orthogonalize': False, 'params': vector, 'lr': 1e-3, 'betas': (0.9, 0.95), 'weight_decay': 0.0},
])

torch.manual_seed(3407), as that's all you need.

THORN works best when pretraining; it doesn't provide much benefit over Adam for non-LoRA fine-tuning, unless the base model was also trained with Muon/THORN.

Usage

Requires Python ≥ 3.12, PyTorch ≥ 2.6. Triton is optional but provides a decent speed boost. FSDP2 is supported & optimized for.

$ pip install git+https://github.com/pykeio/THORN.git
from thorn import THORN

THORN has two 'sub-optimizer's: one for matrix parameters (ConvXD kernels/Linear weights), specified with 'orthogonalize': True; and one for everything else, specified with 'orthogonalize': False. The former is similar to (Nor)Muon, and the latter is essentially (R)AdamW.

You could just give THORN your model and let it figure out which parameters to orthogonalize updates for and which to not:

-optim = AdamW(model.parameters(), lr=1e-3, betas=(0.9, 0.99), weight_decay=0.1)
+optim = THORN(model, lr=1e-3, betas=(0.9, 0.99), weight_decay=0.1)

Note the lack of .parameters() for THORN.

This will reuse the same parameters for orthogonalized-update & non-orthogonalized-update parameters, which isn't ideal. It also might orthogonalize your embedding/final output layer's updates, which isn't recommended. This method is still often better than Adam, but to get the most out of THORN, you should instead specify separate parameter groups:

optim = THORN([
	{
		'orthogonalize': True,
		'params': [p for k, p in model.named_parameters() if p.ndim >= 2 and k not in ['output', 'embed']],
		'lr': 0.001,
		'betas': (0.95, 0.95),
		'weight_decay': 0.1
	},
	{
		'orthogonalize': False,
		'params': [p for k, p in model.named_parameters() if p.ndim < 2 or k in ['output', 'embed']],
		'lr': 0.001,
		'betas': (0.9, 0.995),
		'weight_decay': 0.03
	}
])

$\beta_1$ and $\beta_2$ work differently from Adam for orthogonalized-update parameters: $\beta_1$ is the SGD momentum $\alpha$, like Muon's momentum parameter, and is often $0.9–0.95$. $\beta_2$ controls the second-order momentum $\alpha$ used in per-row normalization for non-tall matrices and is typically set to $0.95$. $\beta_2$ can also be set to $0$ to disable per-row normalization entirely.

weight_decay ($\lambda$) is actually cautious weight decay, so you should set it a bit higher than you normally would. $\approx0.1$ often works well.

There are a few more knobs you can tune besides the usual:

  • none_grad (bool, ⚠️ default True) automatically sets gradients to None after the optimizer completes an update to save memory.
  • decouple_md (bool, default False) factors each parameter into a separately learned magnitude & fixed-norm direction. This works for any parameter group (but not recommended for embeddings/output heads) and can improve learning by 10-20%, though the LR will require tuning.
    • target_norm (Literal['init', 'min', 'max']) specifies the norm the direction is fixed to: init uses the norm at initialization and is the default for non-orthogonalized-update parameters; min targets $\sqrt{\min(M, N)}$; max targets $\sqrt{\max(M, N)}$ and is the default for orthogonalized-update parameters.
  • scaling_mode (Literal['md', 'moonlight', 'jordan']) selects which LR scaling method to use per orthogonalized-update parameter. md is the default when decouple_md=True since it performs best under magnitude-direction decoupling. moonlight is the default otherwise and effectively allows Adam's LR to be directly reused. jordan uses vanilla Muon's scaling.
    • target_rms (float, default $0.2$): When scaling_mode='moonlight', this controls the target RMS of the update to roughly match that of Adam. Adam's update RMS is often between $\approx0.2–0.4$. $0.2$ is a good baseline but increasing to $0.4$ may give better results in some cases.
  • iters (int, default $5$) controls the number of Newton-Schulz iterations performed when orthogonalizing updates. Bumping this up to $8$ improves precision for smaller singular values at the cost of extra compute; whether that matters for training is unclear. $5$ is the minimum number of iterations required to meaningfully converge.
  • momentum_align (bool, default False) scales updates based on their alignment with the momentum & randomly masks updates. Per Magma, this combination can improve learning by ~10% and is stable over a larger range of learning rates.
  • rectify (bool, default True) applies the variance rectification term from RAdam to stabilize the momentum for non-orthogonalized-update parameters during the early stages of training. This usually reduces the need for warmup steps.
  • force_per_neuron_norm (bool, default False) performs per-row normalization for all orthogonalized-update matrices. By default, per-row normalization is limited to tall matrices, which has been found to perform better than applying it to all matrices, at least for transformers. Non-transformers may benefit from enabling this.

For best results:

  • Exclude embedding layers from orthogonalize: True groups since embeddings are vectors, not matrices. For language models, you should also exclude the output head.
  • Keep separate matrices separate: for attention layers, don't merge the Q, K, and V projections into one single qkv_proj, and for GLU-style MLPs, don't merge up_proj and gate_proj into one.

Optional features

If Triton is installed, THORN will use custom kernels to speed up optimizer.step() by up to 40%. The THORN_DISABLE_TRITON environment variable can be set to 1 to disable the use of Triton if problems arise. The use of Triton also means the first optimizer.step() will be very slow as kernels are compiled.

THORN also tries to use torch.compile for additional performance; this may result in NaNs under specific (and uncommon) conditions, so the environment variable THORN_COMPILE can be set to 0 to disable it.

THORN batches parameters together to speed up computation. If you have memory to spare, you can increase optim.polar_decomp_batch_size above its default of 32 * 1024 * 1024 (elements per batch); if you're short on memory, you can set it to 0 to disable batching.

THORN also supports sparse gradients for embedding layers created with nn.Embedding(sparse=True) to save a little extra memory on single-GPU setups.

Gradient release mode

For further memory savings, you can limit gradients to one layer at a time with setup_gradient_release. This slows down FSDP significantly, so it is only recommended on single-GPU setups.

In gradient release mode, gradients are not kept around, so things like float16 mixed precision or gradient clipping will not work (but bfloat16 works).

model = MyModel().to(dtype=torch.bfloat16)

gradient_accumulation_steps = 16

optimizer = THORN(...)
# enable gradient release mode
optimizer.setup_gradient_release(model, update_rate=gradient_accumulation_steps)
# traditional gradient accumulation doesn't work; set update_rate to approximate it instead

scheduler = CosineAnnealingLR(optimizer, ...) # optional

for i, item in enumerate(dataset):
	loss = model(item)
	loss.backward() # <- optimization is done here...

	if (i + 1) % gradient_accumulation_steps == 0:
		# ...so no need to manually step THORN when gradient release is used!
		#optimizer.step()
		#optimizer.zero_grad()

		# but do step the scheduler if you're using one
		scheduler.step()

With the gradient accumulation approximation (update_rate $\gt 1$), the optimizer states are accumulated over microbatches, rather than the gradients themselves. To compensate, you might want to increase all $\beta_1$ to $0.95$ and/or use a lower learning rate, though note that you will always get worse results compared to real gradient accumulation without gradient release mode.

Based on

About

an optimizer for neural networks

Resources

Stars

9 stars

Watchers

1 watching

Forks

Releases

Sponsor this project

Packages

Contributors

Languages