Exercise Solutions: Gradients and Optimizers: From SGD to Adam, and Why m and v Cost 8 Bytes

Exercise solutions: Gradients and Optimizers: From SGD to Adam, and Why m and v Cost 8 Bytes

These are the worked solutions for the exercises in Part 5, Gradients and Optimizers: From SGD to Adam, and Why m and v Cost 8 Bytes. Read the exercise first; coming here before you have tried it defeats the point.

Exercise 1

SGD with momentum stores one running average per weight, 4 bytes: 2,000,000,000 times 4 is 8,000,000,000 bytes, 8 GB.

AdamW stores m and v plus the fp32 master, 12 bytes total: 2,000,000,000 times 12 is 24,000,000,000 bytes, 24 GB.

8-bit Adam stores quantised m and v plus the fp32 master, about 6 bytes total: 2,000,000,000 times 6 is 12,000,000,000 bytes, 12 GB.

Lined up: 8 GB, 24 GB, 12 GB. AdamW costs three times what momentum-only SGD costs and twice what 8-bit Adam costs, entirely for the privilege of a per-weight adaptive step size, which this part argues is worth it for transformers precisely because gradient scales vary so much across a network that one global learning rate cannot serve all of it.

Exercise 2

Real convergence and a rounding-driven plateau look different at the gradient level. If the model had actually converged, the gradients themselves would shrink toward zero, because there would be genuinely little left to correct. The scenario instead describes gradients that stay roughly the same order of magnitude as earlier in the run while the loss goes flat, which is the signature this part’s bf16 arithmetic predicts for a different failure: updates that are individually smaller than the storage format’s local resolution.

If weights are being updated directly in a 16-bit format with no fp32 master copy accumulating the running total, then as training approaches a good region, per-step updates naturally shrink, which is expected and healthy. Once those updates drop below the format’s local gap, roughly 0.4 percent of the current weight’s magnitude for bf16, they round to exactly zero on every single step. The optimizer keeps computing a nonzero gradient and believes it is applying a change; nothing is actually happening to the stored weight. That produces precisely the reported symptom: healthy-looking nonzero gradients in the logs, a learning rate that has not been decayed away, and a loss that simply will not move, which is easy to mistake for convergence but is not.

The one-line fix is to maintain a proper fp32 master copy of the weights and accumulate updates into that, exactly the mixed-precision scheme this part and Part 4 describe, so that updates far smaller than bf16’s local resolution still survive and compound correctly, with only the resulting value rounded to bf16 for the next forward and backward pass.

Exercise 3

SGD with momentum stores exactly one tensor per parameter, the momentum buffer, and that buffer is created in the same dtype as the parameter it belongs to. The layer in this part’s script is an untouched torch.nn.Linear, which defaults to fp32, so the momentum buffer is also fp32, 4 bytes per parameter.

state_bytes(opt) / n should report about 4.0, not 8 and not 12, because SGD with momentum has no second statistic and no separate master copy to add on top; there is exactly one 4-byte number per parameter. As with the AdamW version in this part, the momentum buffer does not exist until the first backward() and step() call, so the same ordering, produce gradients, then step, then measure, is required or the reported figure will be zero.

Exercise 4

PyTorch’s AdamW initialises its per-parameter state, exp_avg and exp_avg_sq, by creating a zero tensor that matches the corresponding parameter’s own dtype and device. Since layer here was cast to bf16 before the optimizer was created, both exp_avg and exp_avg_sq are allocated in bf16, not fp32. state_bytes(opt) / n would report about 4.0, 2 bytes for exp_avg plus 2 bytes for exp_avg_sq, and critically there is no fp32 master weight copy anywhere in this setup, because nothing in this snippet ever created one; layer.parameters() are themselves the only copy of the weights that exists, and they are bf16.

# torch 2.x, illustrative only
layer = torch.nn.Linear(4096, 4096, bias=False).to(torch.bfloat16)
opt = torch.optim.AdamW(layer.parameters(), lr=1e-4)
# exp_avg and exp_avg_sq inherit bf16 from layer's parameters here,
# and there is no separate fp32 master weight copy in this setup at all.

This looks like mixed-precision training because the model is technically running in a 16-bit format end to end, but it is not the recipe this series describes. Real mixed-precision training keeps both the optimizer’s running statistics and the master weight copy in fp32 specifically so that thousands of tiny per-step updates survive rounding across an entire run. A setup that is bf16 all the way down, including the optimizer’s own running averages, reintroduces the vanishing-update problem one level deeper than before: now even the smoothed gradient statistics, not just the raw weight, are too coarse to accumulate small changes reliably. The fix is not to avoid bf16 altogether. It is to keep the model’s stored, master parameters in fp32 and only cast to bf16 for the forward and backward compute, which is what a real mixed-precision path, an autocast context or a training framework such as the Hugging Face Trainer or Accelerate, does correctly, rather than converting the whole module to bf16 by hand and handing that directly to a plain optimizer.