Training
- shown
- toy surface 2 weights
- AdamW
- β₁ 0.9, β₂ 0.95 λ 0.1 · LLaMA 2
- schedule
- warmup + cosine
All steps
A step against the gradient
A toy loss over two weights, with a valley 100 times steeper across than along. Plain gradient descent steps against the gradient. A step size that is safe across the valley is tiny along it, so it zigzags and crawls: loss 0.39 after 20 steps.
w ← w − η · gMomentum
Momentum keeps a running sum of past gradients. The zigzags across the valley cancel out and the steady push along it adds up, so it moves further per step.
v ← β v + g · w ← w − η vAdam: a step size per weight
Adam divides each weight’s step by a running size of its own gradients. Steep and shallow directions then move at a similar pace, set by η, whatever their scale. It is the default for Transformers.
w ← w − η · m̂ / (√v̂ + ε)AdamW: weight decay, apart
Weight decay pulls every weight a little toward zero each step, which keeps them from growing without need. AdamW applies it directly to the weight instead of through the gradient, so Adam’s scaling does not weaken it.
w ← w − η (m̂/(√v̂+ε) + λ w)Warmup, then decay
The step size changes over training. LLaMA 2 warmed up linearly for 2,000 steps (Adam’s running averages start noisy), then followed a cosine down to 10% of its peak of 3 × 10⁻⁴ over about half a million steps.
warmup 2,000 · cosine to 10%What the optimizer keeps
Adam keeps two numbers per weight, m and v, usually in 32 bits, beside a 32-bit master copy of the weights. Training a model needs several times the memory of running it.
≈ 16 bytes per weight vs 2 to serve
Code
opt = torch.optim.AdamW(model.parameters(), lr=3e-4, betas=(0.9, 0.95), weight_decay=0.1)
sched = get_cosine_schedule_with_warmup(opt, num_warmup_steps=2000, num_training_steps=500_000)
for batch in data:
loss = model(batch).loss; loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
opt.step(); sched.step(); opt.zero_grad()
# inside opt.step(), per weight: m = β1·m + (1−β1)·g; v = β2·v + (1−β2)·g²
# w -= lr · (m̂ / (√v̂ + ε) + weight_decay · w)