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

Inside LLM Fine-Tuning, part 5 of 15: Gradients and Optimizers: From SGD to Adam, and Why m and v Cost 8 Bytes

Gradient descent optimizers are the part of training people nod along to and then cannot explain. The gradient tells each weight which way to move. The optimizer decides how far it actually moves. Those are two separate jobs, and once they separate cleanly in your head, SGD, Adam, AdamW and 8-bit Adam stop being a list of names and become four points on one axis.

This part builds that axis. It also settles the arithmetic the memory parts kept referring to: why AdamW’s two running statistics cost 8 bytes per parameter, why the field happily pays that, and why the answer to “just use SGD and save the memory” is almost always no.

By the end you will be able to say what an optimizer stores, price any optimizer from that one fact, explain why per-weight step sizes matter so much specifically for transformers, measure a real optimizer’s state on your own machine, and tell apart the two different appearances of 32-bit precision in a mixed-precision run. The previous part covered the number formats that make the last one interesting.

Read a gradient as an instruction to one weight

Here is the version that sticks. Training is rolling a ball down into a valley, where height is the loss (how wrong the model is) and position is the weights. The gradient is the slope under the ball. It tells you which way is uphill and how steep. To reduce the loss you step the opposite way, downhill.

weight value → loss → gradient (uphill) step downhill = −gradient minimum (goal)

A steep slope means a big gradient, which means this weight matters a lot right now. A flat slope means barely nudge it.

Formally the gradient is the partial derivative of the loss with respect to a weight. That phrase is the only piece of calculus vocabulary in this part, so take it apart once.

A derivative is a rate: how fast one number changes when another number changes. The word partial means you nudge exactly one input and hold every other input perfectly still. That restriction is what makes the idea usable here. A model has 1.54 billion weights, and asking what happens when all of them move at once tells you nothing you can act on. Asking about one weight at a time gives you 1.54 billion separate, simple answers.

So a weight’s gradient reads as “if I nudge this one weight up a hair and change nothing else, does the loss rise or fall, and how fast”. Nothing more mysterious than that. Put numbers on it. The loss is 2.4000. You increase one weight by 0.001, and the loss becomes 2.4030. The loss rose by 0.0030 in response to a nudge of 0.001, so the rate is 0.0030 divided by 0.001, which is 3. That weight’s gradient is 3. It is positive, which says up on this weight is uphill on the loss, so the weight should go down.

The backward pass is simply the machinery that computes this instruction for every weight at once, running back from the loss through the layers, and it never has to nudge anything and remeasure to do it. The full derivation of every derivative involved is available if you want the calculus rather than the ball.

One consequence matters for memory. There is exactly one gradient number per weight, which is why the gradients tenant is the same size as the weights tenant, 2 bytes per parameter in a bf16 run.

Separate what the gradient decides from what the optimizer decides

The gradient gives a direction but not a policy. Do you step exactly along it? A fraction of it? Do you remember which way you were heading and carry some momentum? Do different weights get different step sizes? Those decisions belong to the optimizer.

gradientdirection to nudge OPTIMIZERhow big a step? use history?same rate for all, or adaptive? weight updatethe actual change

Every optimizer is the same shape: gradient in, weight update out. The differences are entirely in what it remembers between steps.

What it remembers between steps is the optimizer state, and the optimizer state is the memory tenant. Remembering is neither free nor temporary. Whatever an optimizer keeps is allocated once per trainable weight and then sits in VRAM from the first step of the run to the last. SGD remembers almost nothing. AdamW remembers two numbers per weight. That difference is the whole story of why the optimizer is 12 of the 16 bytes per parameter that the part on the four memory tenants derived.

Price SGD and momentum, and see where one learning rate runs out

Stochastic gradient descent is the plainest policy: take the gradient, multiply by a learning rate, subtract it from the weight. The learning rate is a single small number, often around 0.0001, that scales every step everywhere in the model. SGD remembers nothing between steps, so it costs zero extra memory per parameter.

Carry the earlier weight through it. Its gradient was 3 and the learning rate is 0.0001, so the step is 3 times 0.0001, which is 0.0003. The gradient was positive, meaning uphill, so the step is subtracted. A weight of 0.5000 becomes 0.4997. That is a complete SGD update, and there is nothing else in it.

Cheap and simple. The catch is that it uses one global learning rate for every weight in the network.

plain SGD · 0 bytes jagged — bounces around, sensitive to lr SGD + momentum · 4 bytes smoother — a running average of past gradients

Momentum keeps one running average of recent gradients, so the path stops zig-zagging and builds speed in consistent directions. That remembered direction is the seed of Adam’s first moment.

Momentum costs one number per weight, 4 bytes in fp32, and genuinely helps. But it does not fix the core limitation. One learning rate still has to suit every weight in the model, and in a transformer that is a much harder ask than it sounds.

Explain m and v, and why they cost 8 bytes per parameter

Adam’s idea is to stop using one global learning rate and instead give every weight its own adaptive step, derived from that weight’s own gradient history. To do that it remembers two running statistics per weight, conventionally called m and v.

m · first moment avg of the gradient = “which way, smoothed” (momentum) 4 bytes / weight (fp32) v · second moment avg of the gradient² = “how big / noisy this weight’s grads are” 4 bytes / weight (fp32) step for this weight = lr × m / √v big/noisy grads (large v) → SMALLER step · small/steady grads → LARGER step

The first is direction, smoothed. The second is a personalised brake. Dividing the step by the square root of the second is the entire adaptive mechanism.

In words: m is a running average of the gradient, so it points the way with the noise smoothed out. v is a running average of the squared gradient, so it measures how large and erratic this particular weight’s gradients have been. Dividing by the square root of v means a weight with large, noisy gradients takes cautious small steps, while a weight with small steady gradients is allowed to move faster.

The braking is easier to believe with numbers in it. Take two weights in the same model. The first sees gradients of roughly 1.0 every step, so its running average of squared gradients settles near 1.0, and the square root of that is 1.0. The second sees gradients of roughly 0.01, so its v settles near 0.0001, and the square root is 0.01. Adam divides each weight’s smoothed gradient by its own square root. For the first weight that is 1.0 divided by 1.0, which is 1. For the second it is 0.01 divided by 0.01, which is also 1. Both weights end up taking a step of about the learning rate, even though their raw gradients differ by a factor of a hundred. Adam does not care how big a weight’s gradients are. It cares how they compare with that weight’s own recent history.

Every weight is tuned individually, and there is no single global learning rate to agonise over. The cost of that per-weight intelligence is storing m and v: 4 bytes plus 4 bytes, so 8 bytes per parameter.

Hold the running example against that figure. Qwen2.5-1.5B, the 1.54 billion parameter model this series sizes everything against, needs 1.54 billion times 8 bytes for m and v alone. That is 12.3 GB. Two numbers per weight, whose only job is to remember, take up more of the card than the weights and the gradients put together, which come to 6.2 GB between them at 2 bytes each. Those are static-state figures. The activations from the forward pass are counted separately and Part 3 covers them.

Why per-weight step sizes matter so much for transformers

This is the part that justifies the memory. A transformer’s gradients vary enormously across its parts: attention against feed-forward, early layers against late, layer norms against projections, embeddings against everything. With one global learning rate, a value that is right for one part is wrong for another. Too large and some weights diverge. Too small and others barely learn. Pascanu et al. described the first failure in its sharpest form, the exploding gradient, and proposed the guardrail still used against it: scale the whole set of gradients down whenever their combined size crosses a threshold. Part 6 covers that guardrail and why it is a separate setting from anything in this part.

Adam’s v term auto-scales each weight’s step to its own gradient magnitude, so one setting works across the whole varied network. You et al. took the same principle one level up, adding a per-layer rescaling that kept training stable at batch sizes in the tens of thousands, where a single global rate had stopped working. That robustness is why practically every large language model is trained with an Adam-family optimizer despite the memory cost. It removes a brutal tuning problem, and tuning problems cost far more than VRAM.

Say what the W in AdamW changes, and what it costs

AdamW is Adam with one fix, and the fix is about weight decay. Weight decay is a small pull applied to every weight on every step, toward zero, so that no weight grows larger than it needs to be. It is there to stop the model fitting its training data too tightly.

AdamW applies that decay directly to the weight, outside the adaptive step, instead of folding it into the gradient where Adam’s own per-weight scaling used to distort it. The distortion is worth one sentence. Fold the decay into the gradient and it gets divided by the square root of v along with everything else, so a weight with large noisy gradients quietly receives less decay than a quiet one, which is not a policy anybody chose. The fix costs nothing in memory, since it changes the update rule rather than the stored state.

w Adam step (m/√v) toward lower loss decay (−lr·λ·w) gently toward zero AdamW keeps these two SEPARATE (Adam tangled them together)

Two forces on each weight, applied separately. The adaptive step handles the gradient. The decay is applied straight to the weight, outside the adaptive machinery, so it acts uniformly.

The next part dissects the update equations line by line, including why bias correction exists.

Answer the argument for dropping back to SGD, and cut the right thing instead

AdamW’s power is also its cost. The m and v it stores are the biggest memory tenant in a training run. There are two ways to react to that, and they are not equally good.

The tempting one is to drop back to SGD and save the 8 bytes. The better one is to keep AdamW and store m and v in fewer bits.

AdamW m · 4 B (fp32) v · 4 B (fp32) = 8 B just for m + v 8-bit Adam m 1B v 1B ≈ 2 B — same behavior, ~4× smaller state m and v are quantized block-wise to 8-bit; the training result is nearly identical to full AdamW

Eight-bit Adam is still Adam. Same two running statistics, same adaptive behaviour, stored more compactly through block-wise quantisation. It shrinks the state, not the algorithm.

Take the SGD argument apart line by line, because it comes up in every memory-constrained conversation.

  • Yes, you would save memory. SGD stores nothing extra, or 4 bytes with momentum, against AdamW’s 12.
  • But transformer training with SGD is notoriously finicky. With one global learning rate you must hand-tune schedules, and it often still underperforms.
  • The v term is precisely what makes Adam robust to the wildly different gradient scales across a transformer’s layers, so one config works across the whole network.
  • So the field pays 8 bytes per parameter for not having to babysit training.
  • The right move under pressure is to shrink the state (8-bit Adam) or remove it for most weights (LoRA), rather than to abandon per-weight adaptivity.

What block-wise quantisation actually does

Shrinking the state deserves its own paragraphs, because “store the same two numbers in fewer bits” hides a real problem. Quantisation means storing a number by mapping it onto a small set of allowed values. Eight bits gives you 256 of them, so every value in a list has to be rounded to one of 256 rungs. Where you put those rungs decides how much you lose. Spread them evenly between the smallest and largest value in the whole list, and one outlier stretches the scale so far that every ordinary value near zero collapses onto the same few rungs.

Block-wise quantisation refuses to use one scale for the whole list. It chops the list into small blocks, a couple of thousand values each in Dettmers et al.’s original 8-bit optimizer work, and gives every block its own scale factor stored alongside it. Inside one small block the values tend to be of similar size, so that block’s 256 rungs land where its numbers actually are. An outlier now spoils only its own block instead of the entire tensor.

The bookkeeping is cheap. One extra 4-byte scale per couple of thousand values is a rounding error next to the values themselves, which is why m and v at 1 byte each come out at about 2 bytes per parameter rather than exactly 2. Run that through the example model. The optimizer tenant goes from 12 bytes per parameter to about 6, because the fp32 master copy of the weights is untouched and still costs its 4. On Qwen2.5-1.5B that is 18.5 GB of optimizer state down to about 9.2 GB, and static state as a whole from about 24.6 GB to about 15.4 GB. Roughly 8 GB of activations is still owed on top of both figures, so a full fine-tune goes from about 33 GB, which does not fit a 32 GB card, to about 23 GB, which does with room to spare.

Optimizer Stores per weight Extra bytes per parameter Behaviour on transformers
SGD nothing 0 finicky, needs careful tuning
SGD with momentum one running average 4 better, still one global learning rate
AdamW m and v, plus the fp32 master 8, plus 4 for the master robust default, per-weight steps
8-bit Adam m and v quantised to 8 bits about 2, plus 4 for the master close to AdamW, much smaller

Two related variants are worth naming now because they appear in later parts. Paged Adam offloads optimizer state to host RAM during a memory spike and brings it back, converting a hard crash into a brief slowdown. And LoRA removes the tenant entirely for every weight it freezes, which is the subject of the part on low-rank adaptation.

Measure the 8 bytes yourself rather than trusting them

You can see the 8 bytes directly. This walks a single linear layer through one step and sums the optimizer’s own tensors. Pin your PyTorch version before you run it: optimizer internals are not a stable public API and have changed between releases.

# torch 2.x
import torch

layer = torch.nn.Linear(4096, 4096, bias=False)      # about 16.8M parameters, fp32
opt = torch.optim.AdamW(layer.parameters(), lr=1e-4)

layer(torch.randn(2, 4096)).sum().backward()          # produce gradients
opt.step()                                            # AdamW allocates m and v here

def state_bytes(optimizer):
    return sum(t.numel() * t.element_size()
               for slot in optimizer.state.values()
               for t in slot.values()
               if torch.is_tensor(t))

n = sum(p.numel() for p in layer.parameters())
print(f"{n:,} params, {state_bytes(opt) / n:.1f} bytes of optimizer state per param")

If you have never used PyTorch, that block is doing ten things. Here they are in order.

  1. torch.nn.Linear(4096, 4096, bias=False) builds one layer. Its weights are a 4,096 by 4,096 grid, which is 16,777,216 numbers, created in fp32 unless you say otherwise. Switching the bias off keeps the parameter count to exactly that grid, so the division at the end comes out clean.
  2. layer.parameters() hands over the list of tensors the optimizer is allowed to change, which here is that single grid. A tensor is a grid of numbers with any number of dimensions. torch.optim.AdamW(...) wraps that list and sets the learning rate. Nothing has been allocated for m and v yet.
  3. layer(torch.randn(2, 4096)) pushes two rows of random numbers through the layer, which is a forward pass. .sum() adds the whole output up into a single number, standing in for a loss. What that number means does not matter here. What matters is that there is exactly one of it, because a backward pass has to start from one number.
  4. .backward() is the backward pass. It computes a gradient for every parameter and attaches it to that parameter, so gradients now sit alongside the weights in memory.
  5. opt.step() is the update. The first time it runs, AdamW finds it has no history for this parameter and creates some: two new tensors the same shape as the weights, one holding m and one holding v.
  6. optimizer.state is where that history lives. It is a dictionary whose keys are the parameter tensors themselves and whose values are one small dictionary per parameter, holding whatever this optimizer needs to remember about it. .values() walks those small dictionaries and ignores the keys, which is all the function needs.
  7. The inner loop, for t in slot.values(), walks the remembered items inside one of those small dictionaries. AdamW puts three things in each: exp_avg, which is PyTorch’s name for m, exp_avg_sq, its name for v, and a step counter.
  8. torch.is_tensor(t) guards the arithmetic, because not every entry in that dictionary is guaranteed to be a grid of numbers. Current PyTorch stores the step counter as a one-element tensor, so it passes the guard and adds 4 bytes to a total of 134 million, which does not move the printed answer. Older versions and some third-party optimizers store the counter as a plain integer, which has no size to measure and would raise an error. The guard makes one function work across both, and it is the reason the version caution above is not boilerplate.
  9. t.numel() is how many numbers a tensor holds, and t.element_size() is how many bytes each of those numbers takes, which is 4 for fp32. Multiply the two for the bytes that tensor occupies, then sum over every remembered tensor for the whole optimizer state.
  10. The last two lines count the parameters and divide. The output is 16,777,216 params, 8.0 bytes of optimizer state per param. Two fp32 numbers per weight, exactly as advertised.

Note that the optimizer state does not exist until the first step(), which is a common source of confusion when a run survives the forward pass and then OOMs on what looks like nothing. The figure printed is about 8, not 12, and the reason is worth understanding: this toy layer is pure fp32, so its parameters already are the master copy. The extra 4 bytes appear only in a mixed-precision run, where a separate fp32 master is kept alongside the bf16 compute weights.

Tell the two appearances of fp32 apart, and know which one you pay for

The letters fp32 turn up twice in any description of a mixed-precision run, meaning two completely different things. One of them lands on your memory bill and the other never does. Separating them is the last piece of the optimizer’s 12 bytes, and it is the single most common confusion about mixed precision.

1 · inside a matmul (transient) multiply bf16 × bf16, sum many terms accumulate the sum in fp32, then cast out not stored — this is your intuition 2 · the master weights (stored) keep the running weight total in fp32 so tiny updates accumulate over steps THIS is the 4 B in the 16-byte bill why #2 is needed: bf16: 0.5000 + 0.0001 → 0.5000 (update vanishes) fp32: 0.5000 + 0.0001 → 0.5001 (update kept ✓)

Two different appearances of fp32 in the same run. Only the second one is stored, and only the second one shows up in the memory bill.

The first appearance is transient. A bf16 matrix multiply accumulates its partial sums in fp32 inside the tensor core and hands back a result. Make that concrete. Multiplying one row of 1,536 numbers by one column of 1,536 numbers means 1,536 separate multiplications, all added into one running total. If that total were rounded back to bf16 after each of the 1,536 additions, the small rounding errors would pile up inside it. So the hardware keeps the running total in fp32 while it is being built and rounds once, at the very end. That total lives inside the chip for the duration of one multiply and is never written out as a stored tensor, so it costs no persistent memory. It is also not something you configure. It happens whether you know about it or not.

The second appearance is persistent, and it is the one in the memory accounting. It is a high-precision copy of the weights kept so that thousands of tiny per-step updates add up instead of each rounding away.

That sentence is easy to nod at, so check it with real numbers. A 16-bit format does not store every decimal number. It stores a fixed set of representable values, and near any given magnitude those values sit a fixed distance apart. For bf16 near a weight of 0.5, the gap between one representable value and the next is about 0.0039. Nothing in between exists. So take a weight of 0.5 and apply an update of 0.0001. The result, 0.5001, is not a value bf16 can hold, and the nearest value it can hold is 0.5 itself. The weight is written back unchanged. The update was not made less accurate. It was deleted.

Now do the same in fp32, whose representable values near 0.5 are about 0.00000006 apart. An update of 0.0001 is more than a thousand steps along that finer grid, so it lands cleanly and the weight becomes 0.5001. Run two thousand steps of that size and fp32 has accumulated 0.2 of real movement, while bf16 is still sitting on exactly 0.5, having thrown away every one of the two thousand updates for the same reason each time. This is why the fix has to be storage rather than arithmetic. No amount of care during the multiply rescues an update that the destination cannot record.

So the correct one-liner is: mixed precision does the arithmetic in bf16, which is fast and small, and keeps the weights’ running total in fp32, which is precise, so that small updates survive. Not “the multiplication result is stored in fp32”, even though the multiply does accumulate in fp32 internally. The stored fp32 is the weight master copy, and it exists to defeat the vanishing-update problem step after step after step.

That stored copy is 4 bytes per parameter, and it is the third of the three fp32 numbers AdamW keeps for every weight: the master copy, m, and v. Three numbers at 4 bytes each is where 12 of the 16 bytes per parameter come from. It is also the 4 bytes the measurement script could not show you, because a pure fp32 layer has no separate master to count.

Key takeaways

  • A gradient is a per-weight instruction: change me in this direction, by this much, to lower the loss. There is exactly one per weight, which is why the gradients tenant matches the weights tenant in size.
  • An optimizer converts that instruction into an actual weight change. What it remembers between steps is its state, and its state is the memory tenant.
  • SGD stores nothing and uses one global learning rate, which makes it finicky on transformers. Momentum adds one running average and helps, without fixing the global-rate problem.
  • Adam stores two running statistics per weight: a smoothed gradient for direction and a running average of the squared gradient as a per-weight brake. That is 8 bytes per parameter and it buys per-weight adaptive steps.
  • Per-weight adaptivity matters most for transformers, where gradient scales differ by orders of magnitude across layers and layer types, so no single learning rate suits all of them.
  • AdamW adds decoupled weight decay, applying the decay directly to the weight rather than through the adaptive scaling. It is free in memory and it is why the default everywhere is adamw.
  • Under memory pressure, quantise the optimizer state to 8 bits or freeze most of the model. Do not fall back to SGD to save the 8 bytes.

You can now

  • Say what one gradient number tells one weight, including what its sign and its size mean, from “Read a gradient as an instruction to one weight”.
  • Split any training step into the part the gradient decides and the part the optimizer decides, from “Separate what the gradient decides from what the optimizer decides”.
  • Explain m and v in plain words and derive the 8 bytes per parameter from them, from “Explain m and v, and why they cost 8 bytes per parameter”.
  • Price an optimizer’s memory from what it stores per weight, for SGD, momentum, AdamW and 8-bit Adam, from “Answer the argument for dropping back to SGD, and cut the right thing instead”.
  • Measure the optimizer state of a real run in a dozen lines of PyTorch and read the number it prints, from “Measure the 8 bytes yourself rather than trusting them”.
  • Tell the transient fp32 inside a matrix multiply apart from the stored fp32 master copy, and say which one you pay for, from “Tell the two appearances of fp32 apart, and know which one you pay for”.

Glossary

8-bit optimizer
AdamW with its two running averages stored in 1 byte each instead of 4, using block-wise quantisation. The algorithm and its behaviour are unchanged, and it reclaims most of 8 bytes per parameter. Dettmers et al., 8-bit Optimizers via Block-wise Quantization
Adam
The optimizer that gives every weight its own step size by tracking two running averages of that weight’s own gradients. AdamW is the corrected version everyone actually uses. Kingma and Ba, Adam
AdamW
Adam with the weight decay applied straight to the weight instead of folded into the gradient. It is the default optimizer for essentially every LLM fine-tune, and its stored state is 12 of the 16 bytes per parameter. AdamW explained, line by line
Attention
The step where each position in the sequence looks at other positions and mixes in whatever it finds useful. It is what lets a model use context instead of reading each token in isolation. Multi-head attention explained
Backpropagation
The procedure that computes a gradient for every weight in one sweep backwards through the model, from the loss at the end to the first layer. It works by applying the chain rule one step at a time. The calculus behind backpropagation
Backward pass
Running backwards from the loss through the model to produce a gradient for every weight. It is backpropagation in practice, and it needs the activations the forward pass stored. How a neural network learns
Batch
A group of examples processed together in one step, so the GPU stays busy and the gradient is averaged over several examples instead of one. Bigger batches give a steadier signal and cost more activation memory. Supervised fine-tuning end to end
bf16
A 16-bit number format with 8 exponent bits and 7 mantissa bits, so it reaches as far as fp32 with much coarser steps. It is the training default because it needs no loss scaling. Number formats for training
Block-wise quantisation
Quantising numbers in small blocks, each with its own scale factor, instead of using one scale for a whole tensor. It handles local variation, and it is how both 8-bit optimizer state and NF4 weights work. Dettmers et al., 8-bit Optimizers via Block-wise Quantization
Chain rule
The rule for finding the slope through a chain of steps: multiply the slope of each step together. It is what lets one loss at the end of the model tell every weight in every layer how to change. The calculus behind backpropagation
Cross-entropy
A score for how wrong a prediction was. It is small when the model gave high probability to the token that actually came next, and large when it did not. How a neural network learns
Decoupled weight decay
Applying weight decay straight to the weight rather than folding it into the gradient. It is the W in AdamW, it costs no extra memory, and it makes the decay land evenly across the model. Loshchilov and Hutter, Decoupled Weight Decay Regularization
Embedding
The lookup table that turns each token ID into a list of numbers the model can do arithmetic on. Its size is vocabulary times hidden size, so a large vocabulary spends a lot of a small model’s parameters here. The complete inference path
Feed-forward network
The part of a transformer block that processes each position on its own, widening it to a larger size and squeezing it back. It holds most of a transformer’s weights and produces its largest activation. The transformer feed-forward network
First moment
Adam’s running average of the gradient, written m. It gives a smoothed direction to move in, so the path stops zig-zagging on noisy batches.
Forward pass
Running data through the model from input to output to get a prediction and a loss. Along the way it produces the activations that the backward pass will need. The complete inference path
fp32
32-bit floating point, with 8 exponent bits and 23 mantissa bits. It is the precise reference format, used for the master copy of the weights and for the optimizer’s running averages. Number formats for training
fp32 master weights
The full-precision copy of the weights that the optimizer actually updates in a mixed-precision run. It exists because a 16-bit weight cannot record an update far smaller than itself, so without it the updates round away and training stalls. Micikevicius et al., Mixed Precision Training
Gradient
One number per weight saying which way to nudge that weight to make the loss smaller, and how steeply the loss responds. Picture the slope under a ball rolling into a valley. How a neural network learns
Layer
One processing stage inside the model, taking a list of numbers in and handing a transformed list out. The example model is 28 transformer layers deep. Inside one transformer block
Layer norm
A step that rescales the numbers flowing through a layer so they stay in a sensible range, which keeps training stable. Modern LLMs use a cheaper version of it called RMSNorm. Inside one transformer block
Learning rate
A single number that scales every weight change. Too high and the loss spikes or blows up, too low and the model barely moves off the base. Learning rate, the master dial
Learning-rate schedule
A rule that changes the learning rate over the course of a run, typically warming up and then decaying. The schedule is a separate thing from the optimizer, and both act on every step. Warmup and schedule
LoRA
Low-rank adaptation. Freeze the model and learn a small pair of skinny matrices beside each targeted weight matrix, so about 1 percent of parameters train and the static state for a 1.5B run drops from about 24.6 GB to about 4 GB. Hu et al., LoRA
Loss
One number saying how wrong the model was on this batch. Training is the whole business of making it smaller, and a falling loss on its own proves very little. How a neural network learns
Matrix multiply
The operation that dominates all the arithmetic in a transformer: multiply a block of inputs by a block of weights to get a block of outputs. Usually shortened to matmul. Inside one transformer block
Memory tenant
One of the four things sharing the card during training: weights, gradients, optimizer state and activations. All four are live at the same moment, so peak memory is their sum rather than the largest of them. The four tenants and the 16 bytes per parameter
Mixed precision
Doing the arithmetic in a 16-bit format for speed and memory while keeping a 32-bit copy of the weights so small updates are not lost to rounding. Essentially every modern training run works this way. Micikevicius et al., Mixed Precision Training
Momentum
Keeping a running average of recent gradients so the path stops zig-zagging and builds speed in a consistent direction. It costs one extra number per weight.
OOM
Out of memory, the error you get when a run needs more VRAM than the card has. In training it almost always strikes where the forward pass ends and the backward pass begins, which points straight at activations. Activation memory and gradient checkpointing
Optimizer
The part of training that turns gradients into actual weight changes. The gradient says which way to move, and the optimizer decides how far.
Optimizer state
The numbers an optimizer keeps between steps, such as running averages of past gradients. Under standard mixed-precision AdamW it is 12 of the 16 bytes per parameter, which makes it the largest memory tenant. Rajbhandari et al., ZeRO
Paged optimizer
An optimizer that can move its state out to ordinary system RAM when GPU memory spikes, then bring it back when the pressure passes. It turns a hard crash into a brief slowdown. Dettmers et al., QLoRA
Parameter
One of the numbers inside the model that training can change. Parameter and weight mean the same thing here, and a 1.5B model has about 1.5 billion of them. Training memory and the 16 bytes per parameter
Partial derivative
The rate at which one output changes when you nudge one input and hold everything else still. A gradient is the partial derivative of the loss with respect to a single weight. The calculus behind backpropagation
Quantisation
Storing numbers with fewer bits by mapping them onto a small set of allowed values. It saves memory and gives up some accuracy in return. bitsandbytes documentation
Regularisation
Any deliberate constraint that stops a model fitting its training data too closely, so it generalises better. Weight decay and dropout are the two you meet in this series. Loshchilov and Hutter, Decoupled Weight Decay Regularization
Running average
A number updated a little at each step to track recent history, where old values fade away instead of being stored. Adam keeps two of them per weight, which is where its memory cost comes from.
Second moment
Adam’s running average of the squared gradient, written v. It measures how large and erratic a weight’s gradients have been, and dividing the step by its square root is what gives each weight its own step size.
SGD
Stochastic gradient descent, the simplest optimizer: multiply the gradient by the learning rate and subtract. It stores nothing extra, and it is unreliable on transformers because one global learning rate has to suit every weight.
Tensor
A grid of numbers with any number of dimensions. One number is a scalar, a row of them is a vector, a table is a matrix, and anything past that is still a tensor with more dimensions. How a neural network learns
Tensor core
The part of an NVIDIA GPU built to do matrix multiplies in low precision very fast. It is why bf16 training beats fp32, and why a format tensor cores cannot multiply directly, such as NF4, costs throughput. NVIDIA, Accelerating AI training with TF32 tensor cores
Training step
One cycle of the loop: forward pass, loss, backward pass, optimizer step, then clear the gradients. Everything else in a training script is arrangements around those five moves. Supervised fine-tuning end to end
Transformer
The architecture behind every model in this series: a stack of blocks that alternate attention with a feed-forward network. Vaswani et al., Attention Is All You Need
VRAM
The memory on the GPU itself. Everything a training step touches has to fit inside it, which is what most of the arithmetic in this series is about. Training memory and the 16 bytes per parameter
Weight
A single learned number inside the model, used to multiply an input on its way through a layer. Weights are what fine-tuning changes, and the only thing it changes. What actually changes inside the model
Weight decay
A small pull on every weight toward zero on every step, so no weight grows larger than it needs to be. It is a form of regularisation, and it is separate from gradient clipping, which is a safety limit. Loshchilov and Hutter, Decoupled Weight Decay Regularization


Practical exercises

Price three optimizers for a 2B model

A 2 billion parameter model is a fine-tuning candidate. Using the bytes-stored-per-weight figures this part gives for SGD with momentum, AdamW, and 8-bit Adam, compute the optimizer-state memory alone, not the full four-tenant static state, for each of the three.

See the worked solution (opens in a new tab)

Diagnose a plateau that looks like convergence

A fine-tuning run drops loss steadily for a few hundred steps, then plateaus hard. Logged gradients are still nonzero and roughly the same order of magnitude as earlier in the run. The learning rate has not decayed to near zero. Someone on the team says the model has simply converged. Give this part’s alternative, numerically grounded explanation for exactly this pattern, and the one-line fix.

See the worked solution (opens in a new tab)

Swap AdamW for SGD in the measurement script

Take this part’s measurement script and replace torch.optim.AdamW(layer.parameters(), lr=1e-4) with torch.optim.SGD(layer.parameters(), lr=1e-4, momentum=0.9), keeping the fp32 layer exactly as given. Before running it, predict what state_bytes(opt) / n will print, and explain the number in terms of what SGD with momentum actually stores per parameter.

See the worked solution (opens in a new tab)

Cast a layer to bf16 and predict what AdamW does next

Instead of the fp32 layer in this part’s script, build the layer with layer = torch.nn.Linear(4096, 4096, bias=False).to(torch.bfloat16), then attach a plain torch.optim.AdamW(layer.parameters(), lr=1e-4) with no other changes, no separate master-weight handling, no framework mixed-precision wrapper. Predict what dtype PyTorch will use for the optimizer’s exp_avg and exp_avg_sq state, what state_bytes(opt) / n will report, and explain why this setup is not the mixed-precision recipe this series describes, even though the model is technically running in a 16-bit format.

See the worked solution (opens in a new tab)

Frequently asked questions

What is the difference between a gradient and an optimizer?

The gradient is data: one number per weight saying which direction lowers the loss and how steeply. The optimizer is policy: it decides how far to actually move each weight given that number and whatever history it has kept. The backward pass produces the gradient, and the optimizer step applies the change.

What are m and v in Adam?

They are two running averages kept per weight. The first, m, averages the gradient itself and gives a smoothed direction. The second, v, averages the squared gradient and measures how large and erratic that weight’s gradients have been. The step divides by the square root of v, so noisy weights are braked and steady ones move faster.

Why does AdamW use 8 bytes per parameter?

Because m and v are each stored as one 32-bit number per weight, which is 4 plus 4 bytes. Adding the fp32 master copy of the weight that mixed-precision training requires brings the optimizer tenant to 12 bytes per parameter, which is three quarters of the standard 16-byte figure.

Should I use SGD instead of AdamW to save memory?

Almost never for language models. SGD’s single global learning rate cannot suit gradients that vary by orders of magnitude across a transformer, so it needs careful hand-tuned schedules and often still trains worse. Quantising AdamW’s state to 8 bits or freezing most of the model with LoRA saves more memory and keeps the reliable behaviour.

What does 8-bit Adam actually change?

Only the storage of the two running statistics, which are quantised block-wise to 8 bits instead of 32. The algorithm is unchanged, so the per-weight adaptive behaviour is the same. It reclaims most of the 8 bytes per parameter and is the standard first move when a run is slightly too large.

Why is there an fp32 copy of the weights if we compute in bf16?

Because bf16 has only about two to three significant decimal digits, so an update far smaller than the weight’s own magnitude rounds away to nothing and training stalls. The fp32 master accumulates those small increments faithfully. Separately, and confusingly, bf16 matrix multiplies also accumulate in fp32 inside the hardware, but that copy is transient and costs no memory.

Sources and further reading

Previous