ContentsThe library

Training Across Many Machines

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.

FIG 1Memory for the model state
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
Every term is per parameter, so the whole thing scales linearly and the only question is the constant. For the standard arrangement, half-precision arithmetic with a full-precision master copy and an optimiser keeping two running averages, those four terms are 2, 2, 4 and 8, for a total of sixteen bytes 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.

FIG 2Where the memory goes for a seven billion parameter model
One hundred and twelve gigabytes before a single activation is stored, for a model whose weights are fourteen gigabytes and which serves comfortably on one device. The thing most people mean by the model is the smallest slice on this chart. Half of the memory belongs to the optimiser, which is why the first and cheapest split in this course is to stop keeping all of it everywhere.

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.

FIG 3Memory against model size, for three situations
0.00312.50625.00937.501250.000.017.535.052.570.0parameters, billions
training, mixed precision, two running averagestraining, momentum onlyserving, half precision
Three straight lines with very different slopes, and a device holds eighty gigabytes. The bottom line crosses that at forty billion parameters, so a forty billion parameter model serves on one device. The top line crosses it at five. Everything this course does is about the distance between those two crossings.

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.

FIG 4Five model sizes against one eighty gigabyte device
State memory, gigabytesActivations at batch 8, Devices needed for stateFits on one device
1 billion parameters16411
7 billion1121820
13 billion2082830
70 billion112095140
405 billion6480310810
The marked row is where it stops. At seven billion parameters, a size that serves on a single device with room to spare, the state alone needs two devices and the activations need a third of another. Every row below it is a distributed systems problem whether anybody wanted one or not. The last column has one entry and it is the small model.

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

NextSame Model, Different Data →

The rest of this course

  1. 01The Weights Are the Smallest Thing You Have to Storeyou are here
  2. 02Every Machine Does the Whole Thing, and Then They Agreeopening only
  3. 03Nobody Is in Charge and That Is Why It Scalesopening only
  4. 04Half a Layer Each, and a Phone Call in the Middle of Every Oneopening only
  5. 05Four Machines, and Three of Them Are Waitingopening only
  6. 06Sixteen Copies of the Same Numbers, in Sixteen Placesopening only
  7. 07You Do Not Choose One, You Multiply Fouropening only
  8. 08Three Hundred Devices Waiting on Oneopening only

Read alongside