Training

Optimizer

SGD, momentum, Adam, AdamW, and the schedule

shown
toy surface 2 weights
AdamW
β₁ 0.9, β₂ 0.95 λ 0.1 · LLaMA 2
schedule
warmup + cosine

All steps

  1. 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 − η · g

  2. Momentum

    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 − η v

  3. Adam: 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̂ + ε)

  4. 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)

  5. 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%

  6. 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)

Go deeper