The Weights Are the Smallest Thing You Have to Store
Training needs several copies of every parameter, not one. Counting them is how you find out whether a run fits on one device, and which budget runs out first.
A model that serves happily on one device will not train on one device, and
the gap is larger than most people expect. It is worth deriving the gap
carefully, because the number that comes out decides everything in the rest of
this course.
Four things, not one
Serving a model holds one thing in memory: the weights. Training holds four.
The weights, in whatever precision the arithmetic runs at. The gradients, which
are one number per weight and therefore exactly the same size. The optimiser
state, which is whatever the update rule needs to remember between steps. And
the activations, which are the intermediate values computed on the way forward
and kept because the backward pass needs them.
Only the first exists at serving time. The other three are the whole reason
this course exists.
- total state memory
- number of parameters
- bytes per weight in the working precision
- bytes per gradient, which matches the weights
- bytes for the full-precision master copy of each weight
- bytes of optimiser state per parameter
Sixteen bytes a parameter
The master copy deserves an explanation because it looks redundant. The
arithmetic runs in half precision for speed, but the updates applied each step
are tiny relative to the weights, and in half precision a small enough update
added to a large enough weight rounds to no change at all. The run stalls
silently. So the authoritative copy of each weight is kept at full precision,
updated there, and copied down to half precision for the next forward pass.
The optimiser state is the largest single item. The usual rule keeps two
running averages per parameter, one of the gradient and one of its square, both
at full precision, which is eight bytes.
The other budget
Activations scale differently. They do not depend on the number of parameters
at all: they depend on how many tokens are in flight and how many layers they
pass through.
Every layer, for every sequence in the batch, for every position in the
sequence, produces intermediate values that the backward pass will need. The
total is roughly proportional to batch size times sequence length times depth
times width, and for a large batch it can exceed the state memory.
This budget has a knob that the state budget does not. You can decline to keep
the activations, remember only the input to each layer, and recompute the
intermediate values during the backward pass when they are needed. That removes
almost all of the activation memory in exchange for running part of the forward
pass a second time, measured at around a third more compute.
Which one runs out
Before choosing any split, compute two numbers: sixteen bytes times the
parameter count, and the activation memory at the batch size you want per
device. Then compare both against what a device actually has.
| State memory, gigabytes | Activations at batch 8, | Devices needed for state | Fits on one device | |
|---|---|---|---|---|
| 1 billion parameters | 16 | 4 | 1 | 1 |
| 7 billion | 112 | 18 | 2 | 0 |
| 13 billion | 208 | 28 | 3 | 0 |
| 70 billion | 1120 | 95 | 14 | 0 |
| 405 billion | 6480 | 310 | 81 | 0 |
The three failures need different answers and it is worth naming them now,
because the rest of the course is organised around them.
If the state does not fit, splitting the data across devices does not help at
all, because that arrangement puts a full copy of the state on every device.
What helps is splitting the state itself, which is the subject of the sixth
lesson.
If the activations do not fit, reduce the batch per device, turn on
recomputation, or split the layers across devices so each one holds fewer. The
batch per device is the first thing to try and the cheapest.
And if everything fits but the run would take four months, nothing is broken:
you need more devices doing the same work on different data, which is the
simplest split of all and the subject of the next lesson.
One number is worth carrying out of this lesson. Divide the total compute of
your run by what one device delivers, and you have the single-device time. The
whole of the rest of this course is about how much of the speedup you keep when
you divide that by a hundred devices, and the answer is never a hundred.
Recap
- Training a model with the usual optimiser takes about sixteen bytes per parameter, of which the weights themselves are two. The optimiser state is the largest single item.
- Activations are the other budget and they scale with batch size and depth rather than with parameters, so they are the one you can trade compute against.
- Work out both numbers before choosing a split, because a run that fails on state needs a different arrangement from one that fails on activations.
This is the reading half
Starting the course gives you your own copy of it. Every idea on every page has problems standing under it, marked with a reason rather than a tick, and any sentence you do not believe can be opened and argued with. None of that can happen on a page nobody owns.
The contents