← Back to Table of Contents

Chapter 11 β€” Optimizers & Loss Functions for LLMs/VLMs

β€œThe loss function defines what you want; the optimizer determines how you get there. Both choices profoundly shape model behaviour.”

Loss Functions

Cross-Entropy Loss β€” The Foundation

Every LLM is trained to minimise cross-entropy between the predicted token distribution and the true next token. This is also called negative log-likelihood (NLL):

\[\mathcal{L}_{\text{CE}} = -\frac{1}{N} \sum_{t=1}^{N} \log p_\theta(x_t \mid x_{<t})\]

where N is the number of labelled tokens (non-masked positions), and p_ΞΈ is the model’s predicted probability for the true token.

In code:

1
2
3
4
5
6
7
8
9
10
import torch.nn.functional as F

# logits: [B, T, V]  β€” model output (unnormalised)
# labels: [B, T]    β€” target token IDs (-100 = ignore)
loss = F.cross_entropy(
    logits.view(-1, vocab_size),   # [B*T, V]
    labels.view(-1),               # [B*T]
    ignore_index=-100,
    reduction="mean",
)

Tensor shapes:

Tensor Shape Description
logits [B, T, V] Unnormalised scores over vocabulary
labels [B, T] True token IDs; -100 for masked positions
probs [B, T, V] After softmax (not needed for loss)
loss scalar Mean NLL over non-masked positions

Perplexity

Perplexity is the exponentiated cross-entropy loss β€” a human-interpretable measure of how β€œsurprised” the model is:

\[\text{PPL} = e^{\mathcal{L}_{\text{CE}}} = e^{-\frac{1}{N}\sum \log p(x_t \mid x_{<t})}\]

Lower perplexity = better model. Typical values: GPT-2 ~30, Llama-3-8B ~7–9 on Wikitext-103.

Z-Loss (Auxiliary Stability Loss)

A regularisation term added during training to prevent the softmax logits from growing very large (logit explosion), which causes instability:

\[\mathcal{L}_z = \alpha \cdot \frac{1}{B \cdot T} \sum_{b,t} \left(\log \sum_v e^{z_{b,t,v}}\right)^2\]

where z are the pre-softmax logits and Ξ± is a small coefficient (e.g., 1e-4). Used in PaLM, Gemini, and others.

Load-Balancing Loss (MoE)

For Mixture-of-Experts models, an auxiliary loss ensures experts are used roughly equally (preventing expert collapse):

\[\mathcal{L}_{\text{aux}} = \alpha \cdot N_E \sum_{e=1}^{N_E} f_e \cdot P_e\]

where f_e is the fraction of tokens routed to expert e, and P_e is the average routing probability to e. This is added to the main cross-entropy loss. See Chapter 33 for full MoE coverage.

Reward Model Loss

Reward models (used in RLHF) are trained with a ranking loss β€” the chosen response should score higher than the rejected one:

\[\mathcal{L}_{\text{RM}} = -\log \sigma(r_\theta(x, y_w) - r_\theta(x, y_l))\]

where y_w is the winning (chosen) response, y_l is the losing (rejected) response, and r_ΞΈ is the scalar reward. This is the Bradley-Terry preference model.

DPO Loss

Direct Preference Optimisation (DPO) reformulates preference learning as a supervised loss:

\[\mathcal{L}_{\text{DPO}} = -\log \sigma\!\left(\beta \log\frac{\pi_\theta(y_w \mid x)}{\pi_{\text{ref}}(y_w \mid x)} - \beta \log\frac{\pi_\theta(y_l \mid x)}{\pi_{\text{ref}}(y_l \mid x)}\right)\]

No reward model is needed β€” the reference policy Ο€_ref (frozen SFT model) implicitly regularises the policy. Ξ² controls how far the policy strays from the reference.

VLM-Specific Losses

Vision-Language Models combine image and text objectives:

Loss Component Objective When Used
Captioning loss (NTP) Predict text tokens conditioned on image Pre-training + SFT
Contrastive loss (CLIP-style) Align image and text embeddings VLM pre-training
Image-text matching Binary classification: does this text match this image? Pre-training auxiliary
Object detection loss Bounding box regression + classification Grounding VLMs

Optimizers

SGD with Momentum

The classical optimizer β€” rarely used for LLM training in practice:

\[v_t = \beta v_{t-1} + g_t \qquad \theta_t = \theta_{t-1} - \eta \cdot v_t\]

Memory: 1 momentum buffer per parameter β†’ 1Γ— params extra.

Adam β€” Adaptive Moment Estimation

Adam (Kingma & Ba, 2014) is the foundation of most LLM optimisers. It maintains running estimates of the first and second moments of gradients:

\(m_t = \beta_1 m_{t-1} + (1-\beta_1) g_t \qquad \text{(first moment β€” mean)}\) \(v_t = \beta_2 v_{t-1} + (1-\beta_2) g_t^2 \qquad \text{(second moment β€” variance)}\) \(\hat{m}_t = \frac{m_t}{1-\beta_1^t} \qquad \hat{v}_t = \frac{v_t}{1-\beta_2^t} \qquad \text{(bias-corrected)}\) \(\theta_t = \theta_{t-1} - \eta \cdot \frac{\hat{m}_t}{\sqrt{\hat{v}_t} + \epsilon}\)

Default hyperparameters: β₁=0.9, Ξ²β‚‚=0.95 (LLMs; note: 0.999 is the original paper default but 0.95 is standard for LLMs), Ξ΅=1e-8, Ξ·=1e-4 to 3e-4.

Memory cost: 2 optimizer states (m and v) per parameter, both in FP32:

Model Size Weights (BF16) Adam States (FP32) Gradients (FP32) Total
7B 14 GB 56 GB 28 GB ~98 GB
70B 140 GB 560 GB 280 GB ~980 GB

This is why distributed training (FSDP/ZeRO) is essential β€” see Chapter 26.

AdamW β€” Adam with Weight Decay

The standard choice for LLM training. AdamW (Loshchilov & Hutter, 2017) decouples weight decay from the gradient update:

Adam (incorrect weight decay): ΞΈ_t = ΞΈ_{t-1} - Ξ· (mΜ‚/√vΜ‚ + λθ_{t-1})
AdamW (correct weight decay): ΞΈ_t = (1 - Ξ·Ξ»)ΞΈ_{t-1} - Ξ· Β· mΜ‚/√vΜ‚

The difference: in AdamW, weight decay is applied directly to weights, not scaled by the adaptive learning rate. This gives more consistent regularisation.

Typical settings for LLM training:

  • lr = 1e-4 to 3e-4 (pre-training); lr = 2e-5 to 2e-4 (fine-tuning)
  • weight_decay = 0.1
  • β₁ = 0.9, Ξ²β‚‚ = 0.95
  • gradient_clip = 1.0

Adam-mini

A recent (2024) memory-efficient variant that reduces optimizer memory by using one learning rate per parameter group (e.g., per attention head) instead of per parameter:

  • Memory: ~45% of Adam memory
  • Performance: matches Adam/AdamW in practice
  • Particularly useful for large models where optimizer states dominate memory

Adafactor

Designed for extreme memory efficiency β€” used to train early T5 models:

  • Does not store full second moment vector; factorises it into rank-1 row/column factors
  • Memory: ~O(√params) instead of O(params) for second moment
  • Trade-off: can be less stable; often needs careful tuning
  • Used in: T5, mT5, Switch Transformer

SOAP

SOAP (Shampoo with Adam in the Preconditioner, 2024) applies second-order information via matrix preconditioning:

  • Precomputes Kronecker-factored curvature (similar to K-FAC/Shampoo)
  • Memory overhead: similar to Adam
  • Speed: can reach the same loss in ~40% fewer steps than AdamW
  • Increasingly used for smaller training runs where step count is the bottleneck

Muon (Momentum + Orthogonalisation)

A newer optimiser (2024) that applies Nesterov momentum and then orthogonalises the gradient update:

  • Applies to weight matrices only (not embeddings/biases)
  • Reaches lower loss than AdamW at same compute in recent experiments
  • Used in MicroGrad and some frontier training runs (reportedly used for DeepSeek V3’s internal layers)

Learning Rate Schedules

The learning rate schedule is as important as the optimiser choice for LLM training.

Cosine Decay with Warmup

The de facto standard for pre-training:

\[\eta_t = \begin{cases} \eta_{\max} \cdot \frac{t}{T_{\text{warmup}}} & t \leq T_{\text{warmup}} \\ \eta_{\min} + \frac{1}{2}(\eta_{\max} - \eta_{\min})\left(1 + \cos\frac{\pi(t - T_{\text{warmup}})}{T - T_{\text{warmup}}}\right) & t > T_{\text{warmup}} \end{cases}\]
Learning Rate Schedule β€” Cosine with Warmup
Warmup
~1000–2000 steps
linear ramp
Cosine Decay
majority of training
smooth decrease
Final LR
Ξ·_min β‰ˆ 0.1 Γ— Ξ·_max
prevents overshoot

Typical warmup: 1000–2000 steps. Starting cold (no warmup) causes early instability.

WSD β€” Warmup, Stable, Decay

Used in MiniCPM and other recent models. Enables continual training by decoupling the decay phase:

  1. Warmup: linear ramp from 0 to Ξ·_max
  2. Stable: constant Ξ·_max for the bulk of training
  3. Decay: cosine or linear decay to Ξ·_min

Key advantage: you can extend training (add more data) by extending the stable phase without restarting the decay. This enables the β€œannealing” trick β€” saving a checkpoint, then annealing on high-quality data.

Trapezoidal / Linear Decay

Used in models like OLMo 2. Simpler than cosine, easier to reason about:

  • Warmup β†’ flat β†’ linear decay to 0

Constant LR (Fine-Tuning)

For SFT/LoRA fine-tuning over few epochs, a constant LR with short warmup often works well:

  • lr = 2e-4 for LoRA, lr = 2e-5 for full fine-tune
  • 5–10% of total steps for warmup

Gradient Clipping

A near-universal component of LLM training. Clips the global gradient norm to prevent instability from gradient spikes:

\[\mathbf{g} \leftarrow \mathbf{g} \cdot \min\!\left(1, \frac{c}{\|\mathbf{g}\|_2}\right)\]

Standard value: max_norm = 1.0. Training loss spikes often coincide with large gradient norm events; clipping at 1.0 prevents most catastrophic updates.

1
2
# In PyTorch
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

Optimizer Comparison Summary

Optimizer Memory (relative) Speed LLM Use Notes
AdamW 3Γ— params β˜…β˜…β˜…β˜… Universal Standard; most widely used
Adam-mini ~1.5Γ— params β˜…β˜…β˜…β˜… Growing Drop-in replacement
Adafactor ~1.1Γ— params β˜…β˜…β˜… T5-era Memory-efficient; less stable
SOAP ~3Γ— params β˜…β˜…β˜…β˜…β˜… Research Faster convergence, same memory
Muon ~2Γ— params β˜…β˜…β˜…β˜…β˜… Frontier Orthogonalised momentum
SGD+momentum 2Γ— params β˜…β˜…β˜… Rarely No adaptivity; hard to tune

Hyperparameter Sensitivity

Key Training Hyperparameters
Learning Rate
Most sensitive. Too high β†’ divergence. Too low β†’ slow convergence. Typical: 1e-4 to 3e-4 for pre-training.
Batch Size
Linear scaling rule: double batch β†’ double LR. Token-level batch is more meaningful (B Γ— T). Typical: 4M–16M tokens/step.
Weight Decay
Regularises weights. Too high β†’ underfitting. Standard: 0.1 for pre-training, 0.01–0.1 for fine-tuning.
Ξ²β‚‚ in Adam
0.95 (LLMs) vs 0.999 (original). Lower Ξ²β‚‚ = faster adaptation to gradient changes. Important for stability.

What’s Next

With the training machinery understood β€” data, objectives, and optimisers β€” we’ll look at mid-training strategies that extend and refine pre-trained models.

← Previous: Chapter 10 β€” Pre-Training at Scale Β· Next: Chapter 12 β€” Mid-Training & Continued Pre-Training β†’


Last updated: April 2026