rumblr Work in progressWIP

● The AI Primer · Lesson 10 · Part 1: how the model works inside

Pretraining at scale

from a web crawl to a base model on thousands of GPUs

This lesson covers Data curation and deduplication, parallelism across GPUs, mixed precision

Members · open during launch 78 min23 figures and diagrams16 interactive
How it works builds the idea from scratch. Math & code adds the formulas and the Python.

At a glance

Key takeaways

  1. Data: pretraining data is mostly web crawl, pushed through language ID, quality rules, a quality classifier, and exact and near-duplicate removal (MinHash with LSH); most of the crawl is discarded. Duplicates cause memorization and waste compute.
  2. Mixture and budget: sources are mixed by weight, not size; a compute-optimal budget is about 20 tokens per parameter, compute is about 6·N·D, and models meant for heavy use are trained far longer.
  3. Memory: Adam in mixed precision needs 16 bytes per parameter before activations (112 GB for 7B), so training needs many GPUs.
  4. Parallelism: data parallel (all-reduce gradients), ZeRO/FSDP (shard the state), tensor parallel (split each matmul, inside a server) and pipeline parallel (split the layers, mind the bubble (p − 1)/(m + p − 1)).
  5. Precision: compute in bf16 (fp32's range, less precision) or fp8 with scaling; fp16 needs loss scaling; the master weights stay in fp32 so small updates are not rounded away.
  6. Stability: warmup, global-norm clipping, spike detection with rollback, and checkpoints every √(2·C·M).

Level 2

How it works, from scratch

Imagine writing an encyclopedia by reading everything ever printed. Two problems appear at once.

First, most of what is printed is junk: flyers, receipts, the same cookie notice on a million websites, spam. You need a sorting line that throws out the junk and the photocopies before anyone reads a page.

Second, no single reader can do the reading. You hire a thousand readers, and now you have a management problem: how do they split the work, share what they learned, and keep going when one of them gets sick?

Pretraining is exactly those two problems. The first is data curation. The second is distributed training.

Figure 1 · Diagram

Reading it: the left half of the chain (crawl, curation, mixture, tokenize) decides what the model learns; most of a model's knowledge and many of its quirks are settled here, before any GPU is switched on. The right half (parallel training, mixed precision, stability) decides whether the learning can happen at all: a 7-billion-parameter model does not fit on one GPU, and a months-long run on thousands of GPUs will see hardware fail many times. Every box gets its own section below, in this order.

The chain above, with a question you can ask it: which box decided that?

Chapter 1

Where the data comes from, and how it is cleaned

Run the belt yourself first; the diamonds below will then read as the stages you just switched off.

Everyday picture Think of a recycling plant. Trucks tip mixed rubbish onto a conveyor belt. Magnets pull out the steel, blowers lift out the paper, people pick out what the machines miss, and only a small fraction reaches the bale at the end. Pretraining data goes down a belt like this, and just as in a recycling plant, most of what goes in never comes out.

The raw material is usually Common Crawl, a nonprofit's public archive of the web: regular snapshots, each of billions of pages. A page arrives as HTML full of menus, adverts and scripts, so the first machine on the belt is text extraction, which keeps the main body text. After that come the filters this section builds: language identification, quality filters, deduplication and mixing. FineWeb, an open dataset built this way from 96 Common Crawl snapshots, ends with about 15 trillion tokens.

Figure 2 · Diagram

Reading it: each diamond is a filter and each "drop" box is a way a page leaves the belt. The order matters for cost: language ID and quality rules look at one page at a time, so they are cheap and run first. Deduplication compares pages with each other, which is the expensive part, so it runs last, on what survived. The rest of this section builds every diamond in turn.

1a. Language identification

Everyday picture Overhear two words of a phone call, "le" and "et", and you already guess French. Every language has a handful of little words that turn up in almost every sentence.

Tiny worked example "the cat is in the garden and it was happy" has 10 words; 7 of them are on the English list of little words (the, is, in, the, and, it, was) and none are on the French, German or Spanish lists, so the guess is English with a share of 0.7. "SKU-4431 X99 blk/wht 12pk" matches no list at all, so it is marked unknown and dropped.

Production pipelines use a trained classifier over character sequences (fastText's language ID model covers 176 languages) and keep a page only if the classifier is confident. The idea is the same: count the evidence for each language and pick the strongest.

In code: detect_language counts each language's stop words (STOP_WORDS) and returns the language with the largest share, or "unknown" below 10%.

1b. Heuristic quality filters

Everyday picture A librarian sorting donations does not read every book. A book with no pages, a pamphlet that is all hashtags, a sheet of keywords: each is rejected at a glance by a simple rule.

Tiny worked example The Gopher paper (Rae et al., 2021) published a set of such rules. Here is what they say about a navigation bar, "Home | About us | Contact | Privacy policy | Log in":

Rule Keep only if The navigation bar Verdict
length 50 to 100,000 words 12 (each "|" counts as a word) fail
mean word length 3 to 10 characters 3.3 pass
symbols (#, ...) per word at most 0.1 0 pass
words containing a letter at least 80% 8 of 12 = 67% fail
common English words at least 2 of: the, be, to, of, and, that, have, with none fail

The last rule is the clever one. Real prose cannot avoid words like "the" and "of"; keyword-stuffed pages and lists of product names avoid them completely. Other rules in the set catch pages that are mostly bullet points or mostly lines ending in "...".

In code: quality_failures applies each rule and returns the names of the ones a document breaks; an empty list means keep.

1c. Quality classifiers: scoring pages against a reference

Rules catch obvious junk. To prefer good text among the survivors, pipelines train a classifier: examples of the text you want (encyclopedia articles, books, pages that trusted sites link to) against random crawl, and then keep pages that score like the reference. GPT-3 filtered Common Crawl with a simple linear classifier of this kind; FineWeb-Edu asked a large model to rate pages for educational value and trained a small classifier on those ratings. We build the simplest version, naive Bayes: every word casts a vote, and the votes add up.

Tiny worked example Our reference examples contain the word "river" 3 times and the spam examples 0 times; "click" appears 0 times in the reference and 4 times in spam. So "river" votes for good and "click" votes for bad. A whole page's score is the sum of its words' votes.

Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
one word of the document "river"
add up the term for every word in the document
a class: good (reference) or bad (spam)
how often appears in class 's examples 3 in good, 0 in bad
total words in class 's examples 77 good, 45 bad
number of distinct words across both classes (the vocabulary) 80
add-one smoothing: pretend every word was seen once more, so an unseen word never gives probability 0 (and log 0 = −∞)
"probability of given ": how likely a word drawn from class is
natural logarithm: turns a ratio above 1 into a positive vote and below 1 into a negative vote

In words: "for each word, ask how much more likely it is in good text than in spam, take the log of that ratio as the word's vote, and add up the votes."

With the numbers: P(river | good) = (3 + 1) / (77 + 80) = 4/157 and P(river | bad) = (0 + 1) / (45 + 80) = 1/125. Their ratio is 3.185 and its log is +1.158. For "click": (0 + 1)/157 over (4 + 1)/125 is 0.159, a vote of −1.837. The sentence "the delta is formed when the river deposits sand over many years" scores +8.52; "click here buy now best price free free free" scores −16.49.

Level 3: in Python
import math
N_good, N_bad, V = 77, 45, 80
def vote(count_good, count_bad):
    # log P(w | good) / P(w | bad), with add-one smoothing
    p_good = (count_good + 1) / (N_good + V)
    p_bad = (count_bad + 1) / (N_bad + V)
    return math.log(p_good / p_bad)
# "river": 3 times in good, 0 in bad
round(vote(3, 0), 3)  # → 1.158
# "click": 0 times in good, 4 in bad
round(vote(0, 4), 3)  # → -1.837

Why "naive"? It treats every word as independent evidence, which is false (words come in phrases) but works well enough to sort billions of pages cheaply. The practical danger is that a classifier keeps whatever resembles its reference set, including its blind spots: pick only encyclopedia text as "good" and you quietly filter out dialects, forums and whole topics.

In code: QualityClassifier.trained_on_examples counts words in GOOD_EXAMPLES and BAD_EXAMPLES; QualityClassifier.word_vote is the formula for one word and QualityClassifier.score adds the votes.

1d. Exact duplicates: one fingerprint per page

Everyday picture A cloakroom attendant does not compare every coat with every other coat. Each coat gets a numbered ticket, and two coats with the same ticket are the same coat.

The web is full of copies: mirrored sites, syndicated news, the same terms of service on a million shops. For exact copies the trick is a hash: a function that turns any text into a short fingerprint, always the same for the same text and almost never the same for different texts. Lowercase the text and squash its spaces first, so that trivially reformatted copies get the same fingerprint, then keep the first page with each fingerprint. One pass, one lookup per page, no comparisons.

In code: normalize_for_exact lowercases and collapses whitespace, and exact_dedup keeps the first document with each SHA-1 fingerprint.

1e. Near duplicates: shingles and Jaccard similarity

A scraper site that copies an article and adds "Read more on our site" defeats the exact hash: one changed character gives a completely different fingerprint. We need a measure of how much two pages overlap.

Everyday picture Cut each page into overlapping strips of a few words, like roof shingles, and put each page's strips in a bag. Two pages are near duplicates when their bags hold mostly the same strips.

Tiny worked example With 2-word shingles:

Sentence Its shingles
"the cat sat on the mat" the cat, cat sat, sat on, on the, the mat
"the cat sat on a mat" the cat, cat sat, sat on, on a, a mat

Three shingles are shared (the cat, cat sat, sat on) and seven are distinct across both, so the overlap is 3/7 = 0.43.

Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
, the two sets of shingles 5 shingles each
the intersection: shingles in both sets {the cat, cat sat, sat on}
the union: shingles in either set, each counted once 7 shingles
the number of items in a set
the Jaccard similarity, from 0 (nothing shared) to 1 (identical) 0.43

In words: "the number of shingles the two pages share, divided by the number of different shingles they have between them."

With the numbers: J = 3 / 7 = 0.43. Real pipelines use 5-word shingles on whole pages; the scraper's copy of our 68-word river paragraph (one phrase changed, one sentence added) shares 78% of its shingles with the original.

Level 3: in Python
def shingles(text, k=2):
    ws = text.split()
    return {" ".join(ws[i:i + k]) for i in range(len(ws) - k + 1)}
A = shingles("the cat sat on the mat")
B = shingles("the cat sat on a mat")
# |A ∩ B| and |A ∪ B|
len(A & B), len(A | B)  # → (3, 7)
# J(A, B)
round(len(A & B) / len(A | B), 2)  # → 0.43

In code: shingles cuts a text into k-word windows (5 by default) and jaccard divides the intersection by the union.

1f. MinHash: estimating Jaccard from a few numbers

Jaccard needs both full shingle sets side by side. With billions of pages, comparing every pair is impossible (a billion pages make about 5 × 10¹⁷ pairs). MinHash compresses each page into a short list of numbers, its signature, such that comparing two signatures estimates their Jaccard.

Everyday picture Shuffle a deck containing every shingle in the world and deal from the top. Stop at the first card that belongs to page A, and separately at the first card that belongs to page B. Those two "first cards" are the same card exactly when the first card from A's-or-B's pile happens to be one they share. The more they share, the likelier that is.

Tiny worked example A = {a, b, c} and B = {b, c, d}, so J = 2/4 = 0.5. Shuffle the four letters in all 24 possible orders. In 12 of them the first letter from A ∪ B is b or c (a shared letter), and then A's first and B's first are the same letter. 12 / 24 = 0.5, exactly J.

Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
a random hash function: gives every shingle a random-looking number, which acts as a random shuffle of all shingles an order of a, b, c, d
the smallest number gives any shingle of : A's "first card"
the probability that the statement in brackets is true 12/24
how many independent hash functions (the signature length) 128
the -th hash function
the indicator: 1 if the statement is true, 0 if not
the estimate of (the hat means "estimated")

In words: "under a random shuffle, two sets have the same first element with probability equal to their Jaccard similarity; so repeat with k shuffles and report the fraction of times the first elements agreed."

With the numbers: all 24 orders of {a, b, c, d}: 12 agree, P = 0.5 = J. With k = 128 hashes, each pair of pages is compared with 128 numbers instead of hundreds of shingles, and the estimate's typical error is about √(J(1 − J)/k) = √(0.25/128) ≈ 0.044.

Level 3: in Python
import itertools, math
A, B = {"a", "b", "c"}, {"b", "c", "d"}
orders = list(itertools.permutations("abcd"))
def first(order, s):
    return next(x for x in order if x in s)
agree = sum(first(o, A) == first(o, B) for o in orders)
agree, len(orders)  # → (12, 24)
# P[min h(A) = min h(B)] equals J(A, B)
agree / len(orders), len(A & B) / len(A | B)  # → (0.5, 0.5)
# typical error of the estimate with k = 128 hashes
round(math.sqrt(0.5 * 0.5 / 128), 3)  # → 0.044

Figure 3 · Chart

2 0 2 1 2 2 2 3 2 4 2 5 2 6 2 7 2 8 number of hash functions k 0.0 0.2 0.4 0.6 0.8 1.0 MinHash estimate of J MinHash: more hashes, tighter estimate true J = 0.20 true J = 0.50 true J = 0.80

Three pairs of sets with true Jaccard 0.2, 0.5 and 0.8: with a handful of hashes the estimates jump around; by 256 hashes each sits close to its dashed true value

Reading it: the horizontal axis is the number of hash functions k (log scale) and the vertical axis the MinHash estimate. Each colour is one pair of sets and its dashed line is that pair's true Jaccard. On the left, with only a few hashes, the estimate can only be a coarse fraction and lurches around. Moving right, each line settles onto its dashed line. That is the whole promise of MinHash: a fixed, small signature per page, with error you choose by choosing k.

In code: MinHasher draws k hash functions of the form (a·x + b) mod p and MinHasher.signature keeps the minimum under each; estimate_jaccard counts agreeing slots.

1g. Locality-sensitive hashing: finding candidate pairs without comparing all of them

Signatures make each comparison cheap, but a billion pages still make too many pairs. Locality-sensitive hashing (LSH) avoids most comparisons: cut each signature into bands of numbers, and file every page into one bucket per band, keyed by that band's numbers. Only pages that share a bucket in at least one band are ever compared.

Figure 4 · Diagram

Reading it: a page flows left to right. Its shingles become a signature of k = b × r numbers, and the signature is sliced into b bands. Each band is a key into its own table of buckets. Two pages meet in the "candidate pairs" box only if all r numbers of some band match exactly; for everything else no comparison ever happens. The final box confirms each candidate with the full signature.
Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
the pair's Jaccard similarity 0.8 or 0.3
rows: numbers per band 5
number of bands 20
chance that all numbers of one band agree (each agrees with chance )
chance that one band does not fully agree 0.672
chance that no band agrees

In words: "a pair becomes a candidate unless every one of its b bands has at least one disagreeing number."

With the numbers: at s = 0.8: 1 − (1 − 0.328)²⁰ = 0.9996, so near duplicates are almost never missed. At s = 0.3: 0.3⁵ = 0.00243, and 1 − 0.99757²⁰ = 0.047, so dissimilar pages are almost never compared.

Level 3: in Python
def p_candidate(s, b=20, r=5):
    # 1 − (1 − s^r)^b
    return 1 - (1 - s ** r) ** b
round(p_candidate(0.8), 4)  # → 0.9996
round(p_candidate(0.3), 4)  # → 0.0475
# where the curve is steepest: about (1/b)^(1/r)
round((1 / 20) ** (1 / 5), 2)  # → 0.55

Figure 5 · Chart

0.0 0.2 0.4 0.6 0.8 1.0 true Jaccard similarity of a pair 0.0 0.2 0.4 0.6 0.8 1.0 chance the pair is compared LSH banding: an S-curve you can place 50 bands × 2 rows 20 bands × 5 rows 10 bands × 10 rows

Probability of becoming a candidate against true similarity, for three band layouts of about 100 hashes: each is an S-curve, steeper and further right as rows per band grow

Reading it: the horizontal axis is the true Jaccard of a pair and the vertical axis its chance of being compared. Every layout draws an S: pairs on the left are almost never compared (that is the saving) and pairs on the right almost always are (that is the recall). The dashed verticals mark (1/b)^(1/r), where each curve is steepest. Choosing b and r is choosing where the cliff sits: more rows per band push it right and make it sharper.

In code: lsh_candidate_probability is the formula, and near_duplicate_pairs runs the whole mechanism: signatures, band buckets, candidate pairs, and a final check against the threshold.

1h. Why duplicates hurt

Everyday picture A student who reads the same paragraph a hundred times can recite it, but has not learned a hundred paragraphs' worth. A model that sees a page thousands of times does the same: it spends capacity memorizing that page, and learns to recite it when prompted.

Tiny worked example A bigram model predicts each word from the word before it, by counting pairs. Train one on five short lines plus the boilerplate "click here to subscribe for free updates", and ask for the probability that it continues "click" into the whole boilerplate line:

Next-word step Kept once Kept 100 times
click → here 1/2 100/101
here → to 1/2 100/101
to → subscribe 1/3 100/102
subscribe → for 2/2 101/101
for → free 1/3 100/102
free → updates 1/2 100/101
whole line 1/72 = 0.014 0.93

Kept once, the model has many ways to continue "click". Repeated 100 times, the line becomes the only road, and the model regurgitates it 93% of the time. Lee et al. (2021) found a single 61-word sentence repeated more than 60,000 times in the C4 dataset, and that models trained on deduplicated data emit memorized training text about ten times less often. Duplicates also waste compute and leak into test sets: a benchmark question copied across the web ends up in the training data, and the model's score on it measures recall, not skill.

In code: verbatim_probability trains the bigram counts on OTHER_LINES plus the given number of copies of BOILERPLATE and multiplies the next-word probabilities along the line.

1i. The whole belt on a small crawl

CRAWL_SAMPLE holds seven pages. Running curate over it:

Page What it is Verdict
0 a clean English paragraph about river deltas kept
1 the same idea in French language: fr
2 a navigation bar quality: too_short, too_few_alphabetic_words, missing_stop_words
3 page 0 re-crawled with different capitals and spacing exact duplicate of 0
4 a hashtag spam page quality: too_many_symbols, missing_stop_words
5 page 0 with a phrase changed and a link added near duplicate of 0
6 a clean English paragraph about stars kept

Two of seven survive. That is not unusual: large public pipelines keep only a small fraction of the raw crawl they start from.

In code: curate runs language ID, the quality rules, the exact hash and a MinHash comparison in that order, and returns one verdict per page.

1j. Data mixtures: how much of each source

Everyday picture A diet is not "all the food in the shop". You choose proportions: mostly staples, some vegetables, a little of the rich stuff. Pretraining data is mixed the same way from sources that differ in size and value: web text, code, books, encyclopedias, scientific papers, maths.

Each source gets a mixture weight: its share of the tokens the model will see. Weights do not follow size. A small, valuable source (an encyclopedia) is often up-weighted, which means the model sees it more than once; a huge, noisy one (the web) is sampled less than one full pass.

Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
one data source the encyclopedia
its mixture weight: share of all training tokens drawn from it 0.05
the total training budget, in tokens 1,000 billion
the tokens that source actually has 20 billion
how many full passes over the source the mixture implies 2.5

In words: "the tokens you plan to draw from a source, divided by the tokens it has, is how many times the model reads it."

With the numbers: web 0.80 × 1,000B / 900B = 0.89 passes; wiki 0.05 × 1,000B / 20B = 2.5 passes; code 0.15 × 1,000B / 150B = 1.0. The LLaMA paper's mixture has the same shape: Common Crawl is two thirds of the tokens at about 1.1 passes, while Wikipedia and books are 4.5% each but are read more than twice (2.45 and 2.23 passes).

Level 3: in Python
weights = {"web": 0.80, "wiki": 0.05, "code": 0.15}
available = {"web": 900e9, "wiki": 20e9, "code": 150e9}
D = 1000e9
# epochs_i = w_i D / N_i
{k: round(weights[k] * D / available[k], 2) for k in weights}  # → {'web': 0.89, 'wiki': 2.5, 'code': 1.0}

Repeating a small source a few times is fine; many more passes and the model starts memorizing it (the duplicate problem again, on purpose). Mixture weights are usually chosen by training small models on candidate mixtures and comparing them, and many recipes change the mixture near the end of training, up-weighting the cleanest sources.

In code: epochs_per_source applies the formula to every source.

1k. Token budgets: how much data for how big a model

Everyday picture A bigger brain can learn more, but only if you give it more to read. Given a fixed amount of study time (compute), there is a best split between "bigger brain" and "more reading".

All budgets are counted in tokens, the pieces a tokenizer cuts text into (see primer.ml.tokenization); in English a token is roughly three quarters of a word. Two rules of thumb set the scale.

Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
the number of model parameters 7 billion
the number of training tokens 140 billion
the compute-optimal token count for size (Hoffmann et al., 2022)
training compute in FLOPs (floating-point operations)
2 FLOPs per parameter per token in the forward pass (one multiply, one add), 4 in the backward pass
"approximately equal": a rule of thumb, not an exact law

In words: "for the best model per unit of compute, train on about 20 tokens per parameter; and training costs about six operations per parameter per token."

With the numbers: a 7B model's compute-optimal budget is 20 × 7 × 10⁹ = 140 billion tokens, and training it costs 6 × 7 × 10⁹ × 1.4 × 10¹¹ = 5.88 × 10²¹ FLOPs. At a sustained 400 teraFLOP/s per GPU (about 40% of an H100's bf16 peak) that is 5.88 × 10²¹ / 4 × 10¹⁴ ≈ 1.5 × 10⁷ GPU-seconds, about 4,100 GPU-hours.

Level 3: in Python
N = 7e9
# D_opt ≈ 20 N
D = 20 * N
D  # → 140000000000.0
# C ≈ 6 N D
C = 6 * N * D
C  # → 5.88e+21
# GPU-hours at 400 teraFLOP/s each
round(C / 400e12 / 3600)  # → 4083

In practice, models that will serve billions of requests are trained far past this point, because a smaller model trained longer is cheaper to run forever after. Llama 2 7B saw 2 trillion tokens (about 286 per parameter); Llama 3 was trained on about 15 trillion tokens across the family (15.6 trillion for the 405B). That is why data, not compute, is now often the binding limit, and why the next topic exists.

In code: chinchilla_tokens and training_flops are the two rules of thumb.

1l. Synthetic data, and model collapse

When good human text runs short, models generate more: rewritten web pages, worked maths solutions checked by a program, question-and-answer pairs drawn out of textbooks. Used carefully this works well: it can turn a messy page into a clear one, and a checked answer is clean signal. Used carelessly it has a known failure.

Everyday picture Photocopy a photo, then photocopy the copy, and keep going. Each copy is nearly right, but fine detail disappears first, and after enough rounds you have a grey smudge. A model trained on its own outputs loses the rare, surprising parts of the data (the tails) first.

Tiny worked example Draw 20 numbers from a bell curve with spread 1. Fit a bell curve to those 20 numbers, draw 20 new numbers from the fit, fit again, and repeat. Each fit is a slightly narrower guess on average, and the narrowing compounds.

Figure 6 · Diagram

Reading it: follow the chain left to right. Only the first model ever sees real data; every later model learns from the previous model's samples. Nothing in the loop can put back a rare value once a sample happens to miss it, so information only leaks out. The fix in practice is to keep real data in every generation's mix, and to filter synthetic data with checks that do not come from the same model.
Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
the generation number 0, 1, 2, …
the fitted variance (spread squared) of generation 's model 1.0 at the start
samples each generation is fitted on 20
the expected value: the average over many repeats of the experiment
the shrink factor per generation of a maximum-likelihood fit 0.95

In words: "on average, each refit on n samples keeps only (n − 1)/n of the previous generation's variance."

With the numbers: 0.95 per generation; after 200 generations the expected variance is 0.95²⁰⁰ ≈ 0.000035 of the original, a spread of about 0.006. The single run below does not follow the average exactly, but it collapses all the same.

Level 3: in Python
import math
n, generations = 20, 200
# (n − 1)/n per generation, compounded
shrink = ((n - 1) / n) ** generations
f"{shrink:.1e}"  # → '3.5e-05'
# the spread is the square root of the variance
round(math.sqrt(shrink), 4)  # → 0.0059

Figure 7 · Chart

0 25 50 75 100 125 150 175 200 generation (each fitted to 20 samples of the last) 1 0 − 4 1 0 − 3 1 0 − 2 1 0 − 1 1 0 0 fitted spread σ Training on your own samples: the spread collapses seed 0 seed 1 seed 2 average shrink √(0.95ᵗ)

Fitted spread over 200 generations of refitting on 20 samples, for three random seeds on a log scale: every run falls by four to five orders of magnitude, even faster than the dashed line from the formula

Reading it: the horizontal axis is the generation and the vertical axis the fitted spread, on a log scale so that halving always looks the same size. Each coloured line is one run with a different seed and the dashed line is the formula's average. The runs wander, up some generations and down others, but the trend is relentlessly downward, and a typical run falls even faster than the dashed line: the average is held up by rare runs that happen to stay wide. After 200 generations the "model" can only produce values in a sliver of the original range. Shumailov et al. (2023) showed the same effect in language models trained recursively on their own text.

In code: recursive_gaussian_fit runs the fit-sample-refit loop and returns every generation's fitted spread.

The votes above, with the page in your hands.

Cut the two sentences yourself first; the formula will then read as a count.

The S-curve above follows three fixed layouts. Here b and r are yours.

Keep the line yourself first; the table below will then read as the roads closing.

The two rules of thumb, with N in your hands.

Chapter 2

Why one GPU is not enough: the memory bill

Everyday picture To bake one cake you need the recipe, but to learn to bake you also need your notes on every attempt: what went wrong, how much to change, how confident you are in each change. Training a model is the same. The weights are the recipe; training also keeps a gradient for every weight (what to change) and the optimizer's running notes (how it has been changing), and those notes are bigger than the recipe.

Tiny worked example Take one parameter trained with Adam (see primer.ml.optimizers) in mixed precision, the standard recipe (section 4 explains the number formats):

What is stored per parameter Format Bytes
the weight used in the forward and backward pass bf16 2
its gradient bf16 2
a full-precision master copy of the weight fp32 4
Adam's momentum (running average of gradients) fp32 4
Adam's variance (running average of squared gradients) fp32 4
total 16
Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
the number of parameters (Psi, the letter the ZeRO paper uses)
bytes for the bf16 weight and its bf16 gradient
bytes for the fp32 master weight and Adam's two fp32 averages
memory for the training state, before any activations 112 GB

In words: "training with Adam in mixed precision costs sixteen bytes for every parameter, before you store a single activation."

With the numbers: 16 × 7 × 10⁹ = 112 GB for a 7B model. A widely used data-centre GPU, the H100, holds 80 GB (see primer.ml.hardware), so the state alone does not fit. Serving the same model needs only the 2-byte weights, 14 GB: training costs eight times the memory of inference.

Level 3: in Python
params = 7e9
bytes_per_param = 2 + 2 + 4 + 4 + 4
bytes_per_param  # → 16
# M_state in GB
bytes_per_param * params / 1e9  # → 112.0
# inference needs only the bf16 weights
2 * params / 1e9  # → 14.0

Activations: the memory that grows with the batch

The backward pass needs the intermediate results of the forward pass (the activations) to compute gradients, so they are kept until it runs. Their size grows with the number of tokens in flight, not with the parameter count. Korthikanti et al. (2022) counted them for one transformer layer, with 16-bit activations:

Level 3: the formula and its symbols

Symbols

Symbol Meaning here For a 7B shape
number of layers 32
sequence length in tokens 4,096
sequences per GPU in the batch 1
hidden width 4,096
attention heads 32
the layer's ordinary tensors (inputs to each matrix multiply, norms, dropout masks)
attention's score and weight matrices, one per head (written as )

In words: "each layer keeps about 34 bytes per token per hidden unit, plus attention's square score matrices, which grow with the square of the sequence length."

With the numbers: per layer, 4096 × 4096 × (34 + 5 × 32 × 4096 / 4096) = 16.8 million × 194 bytes = 3.26 GB; over 32 layers, 104 GB. Kernels like FlashAttention never store the score matrices (they recompute them in the backward pass), which removes the 5as/h term: 34 × 16.8 million × 32 = 18.3 GB.

Level 3: in Python
L, s, b, h, a = 32, 4096, 1, 4096, 32
# A = L · s b h (34 + 5 a s / h)
round(L * s * b * h * (34 + 5 * a * s / h) / 1e9, 1)  # → 104.2
# without stored attention scores
round(L * s * b * h * 34 / 1e9, 1)  # → 18.3
# activation checkpointing keeps only each layer's 2-byte input
round(L * 2 * s * b * h / 1e9, 2)  # → 1.07

The last line is activation checkpointing (also called gradient checkpointing, Chen et al., 2016): keep only each layer's input, and during the backward pass re-run that layer's forward pass to rebuild what it needs. It trades about one extra forward pass (roughly a third more compute) for activation memory that no longer grows with depth.

Figure 8 · Chart

naive no stored scores checkpointed 0 50 100 150 200 GB on one GPU Training a 7B model on one GPU: it does not fit 216 GB 130 GB 113 GB one 80 GB GPU weights (bf16) gradients (bf16) master weights (fp32) Adam momentum (fp32) Adam variance (fp32) activations

Stacked bars for a 7B model at 4,096 tokens: the 112 GB of training state alone passes the 80 GB line, and naive activations add another 104 GB

Reading it: each bar is the memory one GPU would need to train the 7B model alone, split by what it holds; the dashed line is an 80 GB GPU. The three bars differ only in how activations are handled: stored naively, without attention scores, or checkpointed. Even the leanest bar is above the line, because the 112 GB of weights, gradients and optimizer state is there in every bar. Shrinking activations is not enough: the state itself has to be split across GPUs. That is the next section.

In code: training_memory itemizes the 16 bytes per parameter, and activation_bytes is the activation formula, with store_scores=False for FlashAttention-style kernels and checkpointed=True for activation checkpointing.

The three bars above, with the sequence length and the batch in your hands.

Chapter 3

Parallelism: many GPUs, one model

Everyday picture A restaurant can serve more diners in three ways. Open identical kitchens that each cook whole meals (data parallelism). Split one dish across cooks working side by side on the same step, one chopping the left half of the onions and one the right (tensor parallelism). Or build an assembly line, with each cook doing one stage and passing the plate on (pipeline parallelism). Real training runs use all three at once.

3a. Data parallelism: identical copies, averaged gradients

Every GPU holds a full copy of the model and processes a different slice of the batch. After the backward pass, the GPUs average their gradients so all copies take the same step and stay identical. This works because the gradient of an average loss is the average of the gradients.

Tiny worked example A one-weight model predicts y = w·x, with loss the mean of (w·x − y)². Batch: x = 1, 2, 3, 4 with y = 2, 4, 6, 8, and w = 1. On one GPU the gradient is (2/4)·Σ x(wx − y) = (2/4)·(−1 − 4 − 9 − 16) = −15. Split across two GPUs: GPU 1 gets x = 1, 2 and computes −5; GPU 2 gets x = 3, 4 and computes −25. Their average is (−5 − 25)/2 = −15, the same.

Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
the loss averaged over the whole batch
"the gradient of": the slope of the loss for every weight
number of GPUs, each with an equal share of the batch 2
the loss averaged over GPU 's share
add up over every GPU

In words: "the gradient for the whole batch is the average of the gradients each GPU computed on its own equal share."

With the numbers: (−5 + −25) / 2 = −15, matching the single-GPU −15.

Level 3: in Python
def grad(xs, ys, w=1.0):
    # d/dw of mean((w x − y)²)
    return 2 / len(xs) * sum(x * (w * x - y) for x, y in zip(xs, ys))
grad([1, 2, 3, 4], [2, 4, 6, 8])  # → -15.0
g1, g2 = grad([1, 2], [2, 4]), grad([3, 4], [6, 8])
g1, g2  # → (-5.0, -25.0)
(g1 + g2) / 2  # → -15.0

The averaging step is an all-reduce: every GPU contributes a vector and every GPU receives the sum. Sending every gradient to one GPU would jam its network link, so the standard algorithm is the ring all-reduce.

Figure 9 · Diagram

Reading it: the four GPUs sit in a ring and each only ever talks to its right-hand neighbour. Each cuts its gradient into four chunks. Phase one (reduce-scatter): for three steps, each GPU passes one chunk to the right, where it is added to the neighbour's copy; afterwards each GPU owns the complete sum for one chunk. Phase two (all-gather): for three more steps the finished chunks travel round the ring, so everyone ends with all four sums. Every link is busy at every step, and no GPU is a bottleneck.
Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
GPUs in the ring 4, or 1,000
size of the gradient being summed 14 GB (7B in bf16)
steps in each of the two phases 3
each step moves one chunk, of the gradient
two phases: reduce-scatter, then all-gather

In words: "each GPU sends a little under twice its gradient, however many GPUs are in the ring."

With the numbers: with 4 GPUs, 2 × 3/4 = 1.5 gradients' worth; for a 14 GB gradient on 8 GPUs, 2 × 7/8 × 14 = 24.5 GB per GPU per step; on 1,000 GPUs, 27.97 GB. The traffic per GPU barely grows, which is why data parallelism scales to thousands of GPUs.

Level 3: in Python
def sent(G, S):
    # 2 (G − 1) / G · S
    return 2 * (G - 1) / G * S
sent(4, 1.0)  # → 1.5
sent(8, 14e9) / 1e9  # → 24.5
round(sent(1000, 14e9) / 1e9, 2)  # → 27.97

Data parallelism alone does not solve the memory problem: every GPU still holds all 16 bytes per parameter.

In code: linear_regression_gradient computes a batch's gradient, so the averaging identity can be checked on shards; ring_all_reduce simulates both phases and counts what each GPU sends; all_reduce_traffic is the formula.

3b. Sharding the state: ZeRO and FSDP

Everyday picture A reading group with one expensive textbook does not buy a copy each. They tear it into chapters, each keeps one, and whoever needs a chapter borrows it for the evening and hands it back.

In plain data parallelism, every GPU stores an identical copy of the optimizer state, the gradients and the weights: pure waste. ZeRO (the Zero Redundancy Optimizer) keeps data parallelism's split of the batch but gives each GPU only a 1/G shard of that state, in three stages. PyTorch's FSDP (fully sharded data parallel) implements the third.

Figure 10 · Diagram

Reading it: read top to bottom; each stage moves one more kind of state from "every GPU keeps all of it" to "every GPU keeps a 1/G slice". Stage 1 works because each GPU only needs to update its own slice of the weights. Stage 2 works because a reduce-scatter (the first half of the ring) can deliver each GPU just the gradient slice it updates. Stage 3 goes furthest: before computing a layer, the GPUs all-gather that layer's weights, use them, and throw them away again.
Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the ZeRO paper's example
parameters
GPUs sharing the state 64
the fp32 master weights and Adam's two averages 90 GB
, the bf16 weights; weights plus gradients 15 GB; 30 GB
training state per GPU at that stage

In words: "whatever is sharded is divided by the number of GPUs; what is not sharded stays whole on every GPU."

With the numbers: 7.5B parameters on 64 GPUs (the example in Figure 1 of the ZeRO paper): stage 0, 120 GB; stage 1, 30 + 1.4 = 31.4 GB; stage 2, 15 + 1.6 = 16.6 GB; stage 3, 1.9 GB. The 7B model that could not fit on one GPU now needs under 2 GB of state per GPU.

Level 3: in Python
psi, G = 7.5e9, 64
stages = [16 * psi, 4 * psi + 12 * psi / G, 2 * psi + 14 * psi / G, 16 * psi / G]
[round(m / 1e9, 1) for m in stages]  # → [120.0, 31.4, 16.6, 1.9]

The price is communication. Stages 1 and 2 cost the same traffic as a plain all-reduce; stage 3 adds an all-gather of the weights in the forward pass and again in the backward pass, about 1.5 times the traffic. Frameworks hide it by fetching the next layer's weights while the current layer computes.

Figure 11 · Chart

2 1 2 3 2 5 2 7 2 9 GPUs sharing the state 1 0 − 1 1 0 0 1 0 1 1 0 2 training state per GPU (GB) ZeRO: sharding divides the 16 bytes per parameter 80 GB GPU stage 0: plain data parallel stage 1: + shard optimizer stage 2: + shard gradients stage 3 (FSDP): + shard weights

Training state per GPU for a 7B model as the number of GPUs grows from 1 to 1,024: stage 0 stays at 112 GB, stages 1 and 2 flatten at their unsharded floors, and stage 3 keeps falling

Reading it: the horizontal axis is the number of GPUs (doubling at each tick) and the vertical axis the state each GPU holds, both on log scales. The dashed line is an 80 GB GPU. Stage 0 is flat: adding GPUs never helps. Stages 1 and 2 fall at first and then level off at the part they do not shard (28 GB and 14 GB of weights and gradients). Only stage 3 keeps falling in a straight line, because nothing is left unsharded.

In code: zero_memory_per_gpu gives the state per GPU for any stage and GPU count.

3c. Tensor parallelism: splitting one matrix multiply

Everyday picture Two people fill in one large multiplication table: one does the left half of the columns, the other the right half. Neither needs the other's work until the end, when the halves are placed side by side.

The biggest operations in a transformer are matrix multiplies, X·W. They can be split across GPUs in two ways, and both give exactly the unsplit answer.

Tiny worked example X = [1, 2] and

W = [[1, 2, 3, 4], [5, 6, 7, 8]], so X·W = [11, 14, 17, 20].

  • By columns: GPU 1 holds W's first two columns and computes [11, 14]; GPU 2 holds the last two and computes [17, 20]. Place side by side: [11, 14, 17, 20].
  • By rows: GPU 1 holds W's first row and X's first entry: 1 × [1, 2, 3, 4] = [1, 2, 3, 4]. GPU 2 holds the second: 2 × [5, 6, 7, 8] = [10, 12, 14, 16]. Add: [11, 14, 17, 20].
Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
the input activations, one row per token [1, 2]
a weight matrix 2 × 4
place two matrices side by side (concatenate columns)
left: 's columns, split into two blocks 2 × 2 each
right: 's rows, split into two blocks 1 × 4 each
right: 's matching columns [1] and [2]

In words: "split the weight by columns and each GPU produces some of the output columns; split it by rows and each GPU produces a partial sum of every output, and the partial sums add up to the answer."

With the numbers: [11, 14] next to [17, 20], or [1, 2, 3, 4] + [10, 12, 14, 16]: both give [11, 14, 17, 20].

Level 3: in Python
X = [1, 2]
W = [[1, 2, 3, 4], [5, 6, 7, 8]]
def matmul(x, w):
    return [sum(x[i] * w[i][j] for i in range(len(x))) for j in range(len(w[0]))]
matmul(X, W)  # → [11, 14, 17, 20]
# by columns: each GPU computes half the output columns
matmul(X, [row[:2] for row in W]) + matmul(X, [row[2:] for row in W])  # → [11, 14, 17, 20]
# by rows: each GPU computes a partial sum of every output
p1, p2 = matmul(X[:1], W[:1]), matmul(X[1:], W[1:])
[a + b for a, b in zip(p1, p2)]  # → [11, 14, 17, 20]

Megatron-LM combines the two for a transformer's feed-forward layer, which is Y = activation(X·W₁)·W₂: split W₁ by columns and W₂ by rows.

Figure 12 · Diagram

Reading it: the input is copied to both GPUs. The column split of W₁ gives each GPU whole hidden units, so the activation function (applied number by number) runs locally with no communication. The row split of W₂ then consumes exactly those hidden units and produces a partial sum of the output. One all-reduce at the end adds the two partial sums. Needing only one all-reduce per block (the attention block is split the same way) is what makes this practical, but it still happens inside every layer, so tensor parallelism needs the fastest links available and usually stays within one server of 8 GPUs.

In code: column_parallel_matmul and row_parallel_matmul are the two splits, and tensor_parallel_mlp is the Megatron-LM feed-forward layer with its single all-reduce.

3d. Pipeline parallelism, and the bubble

Everyday picture A car assembly line with four stations. The first car takes four steps to roll off the end, and while it travels, stations further down stand idle. Only when many cars are on the line at once is every station busy.

Pipeline parallelism gives each GPU a consecutive block of layers (a stage). To keep the stages busy, the batch is cut into micro-batches that follow each other down the line. The idle time at the start and end is the pipeline bubble.

Figure 13 · Diagram

Reading it: each GPU owns a quarter of the layers. Micro-batches enter on the left one after another, and each GPU passes its output activations to the next. When the forward passes are done, gradients flow back along the dashed arrow in the reverse order. The only traffic is activations at stage boundaries, far less than tensor parallelism's per-layer exchange, which is why pipeline stages can sit on different servers.
Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
pipeline stages (GPUs in the line) 4
micro-batches per batch 1, 8 or 32
time steps to push micro-batches through stages: to fill the line, then one per extra micro-batch 11 when
steps each stage spends idle while the line fills (and again while it drains) 3

In words: "the fraction of time each GPU sits idle is the fill time divided by the total time; more micro-batches spread the same fill time over more work."

With the numbers: with 4 stages and 1 micro-batch, 3/4 = 75% of the time is bubble; with 8 micro-batches, 3/11 = 27%; with 32, 3/35 = 8.6%.

Level 3: in Python
def bubble(p, m):
    # (p − 1) / (m + p − 1)
    return (p - 1) / (m + p - 1)
[round(bubble(4, m), 3) for m in (1, 8, 32)]  # → [0.75, 0.273, 0.086]

Figure 14 · Chart

0 2 4 6 8 10 12 14 16 18 20 time step (blue: forward of micro-batch n, green: backward, white: idle) GPU 1 GPU 2 GPU 3 GPU 4 GPipe schedule, 4 stages, 8 micro-batches: bubble = 27% 1 2 3 4 5 6 7 8 1 2 3 4 5 6 7 8 1 2 3 4 5 6 7 8 1 2 3 4 5 6 7 8 1 2 3 4 5 6 7 8 1 2 3 4 5 6 7 8 1 2 3 4 5 6 7 8 1 2 3 4 5 6 7 8

The GPipe schedule for 4 stages and 8 micro-batches as a grid of GPUs by time step: forward passes form a staircase down, backward passes a staircase back up, and the empty triangles in the corners are the bubble

Reading it: each row is a GPU (stage) and each column one time step. Blue cells are forward passes and green cells backward passes, numbered by micro-batch. The forward staircase runs down and to the right as each micro-batch moves to the next stage; the backward staircase runs back up. The white triangles in the corners are the bubble: stage 4 waits three steps for the first micro-batch to arrive, and stage 1 waits three steps at the end. Count them: 24 of the 88 cells are empty, 3/11 of the grid.

Figure 15 · Chart

1 0 0 1 0 1 1 0 2 micro-batches per batch 0.0 0.2 0.4 0.6 0.8 share of time idle The pipeline bubble: (p − 1) / (m + p − 1) 2 stages 4 stages 8 stages 16 stages

Bubble fraction against number of micro-batches for 2, 4, 8 and 16 stages: every curve starts high and falls towards zero, and deeper pipelines need more micro-batches for the same efficiency

Reading it: the horizontal axis is the number of micro-batches (log scale) and the vertical axis the idle fraction. Each line is a pipeline depth. To keep the bubble under about 10%, you need roughly ten times as many micro-batches as stages, which pushes up the batch size. Smarter schedules exist for exactly this reason: "one forward, one backward" (1F1B) starts backward passes early to free activation memory, and interleaved stages (Narayanan et al., 2021) give each GPU several smaller stages to shrink the bubble further.

In code: pipeline_schedule builds the grid (positive numbers forward, negative backward, 0 idle) and bubble_fraction is the formula.

3e. All three at once

Figure 16 · Diagram

Reading it: the three kinds nest, and each is placed where its traffic fits. Tensor parallelism, which talks inside every layer, stays inside one server on its fastest links. Pipeline parallelism, which only passes activations between stages, spans servers. Data parallelism (often sharded with ZeRO or FSDP), which talks once per step, wraps the whole thing and multiplies it across the cluster. Llama 3 405B was trained this way on up to 16,384 GPUs, adding a fourth kind (context parallelism) that splits very long sequences.

The ring above, one step at a time.

The chart above, with G in your hands.

The schedule above, for any depth and any number of micro-batches.

Chapter 4

Mixed precision: doing the maths in fewer bits

Everyday picture A carpenter measures a room with a tape measure and a table leg with calipers. Using calipers for everything would be slow; using the tape measure for everything would give wobbly tables. Mixed precision does each job with the coarsest number format that is good enough: the heavy matrix multiplies in 16 (or 8) bits, and the few places where tiny differences matter in 32.

The payoff is large. Halving the bits halves the memory and the data moved, and GPU tensor cores run 16-bit maths many times faster than 32-bit (see primer.ml.hardware).

4a. What a floating-point number is

Everyday picture Scientific notation, in binary. "6.02 × 10²³" has a sign, a few significant digits and an exponent. A float is the same three parts in bits: the exponent sets the range (how large or small a number can be), and the fraction bits set the precision (how many significant digits it keeps).

Tiny worked example In bf16, 1/3 is stored as sign 0, exponent field 125 and fraction field 0101011 in binary (43). The value is 2^(125 − 127) × (1 + 43/128) = 0.25 × 1.3359375 = 0.333984375. Seven fraction bits keep only about three significant decimal digits, so bf16 cannot tell 1/3 from 0.33398.

Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
sign 1 bit: 0 for positive, 1 for negative 0
+1 or −1 +1
the exponent field, read as a whole number 125
bias a fixed offset, for exponent bits, so negative exponents can be stored 127 (bf16 and fp32)
the fraction field, read as a whole number 43
the number of fraction bits 7
the significant digits, between 1 and 2 (the leading 1 is implied, not stored) 1.3359375

In words: "a float is plus or minus a number between 1 and 2, scaled by a power of two; exponent bits choose the power, fraction bits choose the number between 1 and 2."

With the numbers: 2^(−2) × (1 + 43/128) = 0.333984375. The gap to the next bf16 number above 1 is 2⁻⁷ = 0.0078; in fp32, with 23 fraction bits, it is 2⁻²³ ≈ 0.00000012.

Level 3: in Python
sign, E, bias, F, m = 0, 125, 127, 43, 7
# (−1)^sign × 2^(E − bias) × (1 + F / 2^m)
(-1) ** sign * 2.0 ** (E - bias) * (1 + F / 2 ** m)  # → 0.333984375
# the gap after 1 (the "epsilon") in bf16 and fp32
2.0 ** -7, 2.0 ** -23  # → (0.0078125, 1.1920928955078125e-07)

The formats that matter for training:

Format Exponent bits Fraction bits Largest Smallest normal Gap after 1
fp32 8 23 3.4 × 10³⁸ 1.2 × 10⁻³⁸ 1.2 × 10⁻⁷
bf16 8 7 3.4 × 10³⁸ 1.2 × 10⁻³⁸ 0.0078
fp16 5 10 65,504 6.1 × 10⁻⁵ 0.00098
fp8 E5M2 5 2 57,344 6.1 × 10⁻⁵ 0.25
fp8 E4M3 4 3 448 0.016 0.125

Below the smallest normal number a format has a few subnormal values that trade precision for extra range (fp16 reaches down to 6.0 × 10⁻⁸), and below half of the smallest subnormal a number becomes 0: underflow. Above the largest value is overflow, which becomes infinity.

Figure 17 · Chart

−40 −30 −20 −10 0 10 20 30 40 log₁₀ of magnitude (light: subnormals, dark: normal range) fp32 bf16 fp16 fp8 E5M2 fp8 E4M3 Range of each format: bf16 keeps all of fp32's

Horizontal bars on a log scale showing the range of each format from its smallest subnormal to its largest value: fp32 and bf16 span the same huge range, fp16 and fp8 are far narrower

Reading it: each bar covers the magnitudes one format can represent, on a log scale where each tick is a factor of 10. The darker part is the normal range and the lighter tail on the left is the subnormals. bf16's bar is as long as fp32's: same 8 exponent bits, same range. fp16's is a small fraction of it, and fp8's smaller still. Range is what decides whether a gradient survives; precision (not shown) decides how finely it is recorded.

In code: FloatFormat describes a format by its bit counts (with FP32, BF16, FP16, FP8_E5M2 and FP8_E4M3 defined), and quantize rounds any number to the nearest value a format can hold, reproducing overflow, subnormals and underflow.

4b. Underflow, and loss scaling

Everyday picture A kitchen scale that reads to the nearest gram shows 0 for a pinch of saffron. Weigh the saffron together with a known 1 kg jar, subtract the jar afterwards, and the pinch shows up.

Gradients are often tiny. Many are smaller than fp16's smallest subnormal (about 6 × 10⁻⁸), so a backward pass in fp16 silently turns them to 0 and those weights stop learning. Loss scaling is the jar: multiply the loss by a large number S before the backward pass, so every gradient is S times larger (the chain rule passes the factor through unchanged); then divide by S in fp32, before the update.

Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
the true gradient of one weight 10⁻⁸
the loss scale, usually a power of two so multiplying is exact 65,536 = 2¹⁶
"rounded to the nearest fp16 value", which is where the backward pass stores it
the gradient the optimizer receives, after dividing in fp32 ≈ 10⁻⁸

In words: "scale the loss up so the gradients are big enough for fp16, compute them in fp16, and scale them back down in fp32."

With the numbers: unscaled, 10⁻⁸ is under half of fp16's smallest value (5.96 × 10⁻⁸) and rounds to 0. Scaled, 10⁻⁸ × 65,536 = 6.55 × 10⁻⁴, a normal fp16 number; it is stored as 6.5517 × 10⁻⁴ and divided back to 9.997 × 10⁻⁹, within 0.03% of the truth.

Level 3: in Python
import numpy as np
grad, S = 1e-8, 65536
# without scaling: fp16 flushes it to zero
float(np.float16(grad))  # → 0.0
# with scaling: fp16 holds S · grad, then divide in higher precision
f"{float(np.float16(S * grad)) / S:.4e}"  # → '9.9972e-09'

Figure 18 · Chart

−50 −40 −30 −20 −10 0 10 20 gradient size, log₂ 0 10000 20000 30000 40000 50000 60000 70000 number of gradients Loss scaling slides the gradients into fp16's range fp16 overflow flushed to 0 unscaled: 24% become 0 in fp16 × 2¹⁶: 0.00% become 0

Histogram of the size of a million simulated gradients on a log scale: unscaled, a large share falls left of fp16's smallest value and is lost; multiplied by 65,536 the whole histogram shifts right into fp16's range

Reading it: the horizontal axis is a gradient's size in powers of two and the height is how many gradients have that size. The shaded region on the left is below fp16's smallest subnormal: whatever lands there becomes 0. The grey histogram is the unscaled gradients, with the lost share printed; the blue one is the same gradients times 2¹⁶, shifted 16 steps to the right, clear of the shaded region and still far from the overflow wall on the right. Loss scaling does not change the shape, only where it sits.

A fixed scale is fragile: too small and gradients underflow, too large and they overflow to infinity. Dynamic loss scaling adapts it: if any gradient is infinite or not-a-number, skip the step and halve S; after a long run of clean steps (2,000 is common), double S.

Figure 19 · Diagram

Reading it: one training step goes round the loop once. The weights live in fp32 (top) but are cast to 16 bits for the expensive forward and backward passes. The loss is multiplied by S before the backward pass. The diamond is dynamic loss scaling's check: an overflow means S was too big, so the step is thrown away and S halved; otherwise the gradients are unscaled, clipped and applied to the fp32 master copy. With bf16 the scale can usually be dropped altogether, because bf16 has fp32's range: the same 10⁻⁸ gradient is stored as 1.0012 × 10⁻⁸ without any help. That is why bf16 became the default for training on hardware that supports it.

In code: scaled_gradient_roundtrip scales, stores and unscales one gradient, and DynamicLossScaler.update skips the step and halves the scale on overflow, or doubles it after a long enough run of clean steps.

4c. Why the master copy stays in fp32

Everyday picture Pour a teaspoon of water into a full bathtub and measure with a bucket: the level has not changed, as far as the bucket can tell. Do it a thousand times and you have added four litres, but every single measurement still reads "no change".

Tiny worked example A weight of 1.0 receives an update of +0.0001 per step. Next to 1.0, bf16 can only step in increments of 0.0078, so 1.0001 rounds straight back to 1.0. After 1,000 updates the bf16 weight is still exactly 1.0; an fp32 weight has moved to 1.1, as it should.

That is why the optimizer keeps an fp32 master copy of the weights and applies updates to it, even though the matrix multiplies use 16-bit copies. The 16-bit copy is re-made from the master after every step.

In Python:

import numpy as np
w16, w32 = np.float16(1.0), np.float32(1.0)
for _ in range(1000):
    w16 = np.float16(w16 + np.float16(1e-4))
    w32 = np.float32(w32 + np.float32(1e-4))
# every update lost in 16 bits, all kept in 32
float(w16), round(float(w32), 4)  # → (1.0, 1.1)

Figure 20 · Chart

0 200 400 600 800 1000 update step (each adds 0.0001) 1.00 1.02 1.04 1.06 1.08 1.10 weight value Tiny updates vanish in 16 bits; an fp32 master keeps them fp32 bf16 fp16

A weight receiving 1,000 updates of 0.0001: in fp32 it climbs in a straight line from 1.0 to 1.1; stored in bf16 or fp16 it never leaves 1.0

Reading it: the horizontal axis is the step and the vertical axis the weight's value. The fp32 line rises steadily by 0.0001 per step. The bf16 and fp16 lines are flat at 1.0 for the whole run: each update is smaller than half the gap to the next representable number, so it is rounded away every time. The error is not noise that averages out; it is a systematic loss of every small update, which is exactly what late training consists of.

In code: accumulate_updates adds the same update many times, rounding the weight to a chosen format after each step.

4d. fp8: the next halving

Everyday picture A ruler with only eight marks is useless for measuring a hair, unless you first slide it under a magnifying glass set to the right zoom. fp8 is that short ruler, and a per-tensor scale is the magnifying glass.

Recent GPUs multiply 8-bit floats at twice the 16-bit rate. With so few bits, one format cannot cover everything, so training uses two: E4M3 (more precision, range to 448) for weights and activations, and E5M2 (more range, to 57,344) for gradients. Because the range is so short, every tensor (or every small block of a tensor) gets its own scale factor, chosen from its recent largest value, the same idea as loss scaling applied everywhere. DeepSeek-V3 was trained largely in fp8 this way, keeping sensitive parts (normalizations, the optimizer, the master weights) in higher precision.

The table above, one number at a time.

The histogram above, for one gradient in your hands.

The flat lines above, with the update size in your hands.

Chapter 5

Stability at scale: loss spikes, clipping, warmup and checkpoints

Everyday picture A ship on a months-long voyage does not assume the sea stays calm. It trims the sails when gusts come (clipping), leaves port slowly (warmup), keeps a log of its position (checkpoints), and has a drill for when something goes wrong (rolling back).

A large run does see storms. The loss curve occasionally jumps upward, a loss spike, sometimes recovering by itself and sometimes diverging for good. Causes include a batch of bad data, a learning rate a little too high for the model's current state, attention logits growing without bound, and numerical overflow. And separately from the maths, hardware fails: in the Llama 3 405B run, 419 unexpected interruptions occurred over 54 days, about one every three hours.

Figure 21 · Diagram

Reading it: the main loop is train, check, maybe save, repeat. The upper diamond watches the loss. A brief blip is ignored, since clipping usually absorbs it; a spike that persists sends the run back to an earlier checkpoint, and the data batches that were in flight around the spike are skipped. PaLM's authors did exactly this: restart about 100 steps before the spike and skip roughly 200 to 500 batches, which removed the spikes. A checkpoint must include the optimizer state and the position in the data, or the resumed run is not the same run.

5a. Spotting a spike

Tiny worked example Losses 3.0, 2.9, 2.8, 2.8, 2.7, then 8.5. The median of the previous four is 2.8, and 8.5 is more than twice that, so step 5 is flagged. The next step, 4.0, is compared with the median of 2.8, 2.8, 2.7, 8.5, which is still 2.8: one spike does not raise the bar, which is why the rule uses the median and not the mean.

Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
the training loss at step 8.5 at
how many recent steps to look back over (the window) 4
the middle value after sorting; one outlier cannot move it far 2.8
how many times the median counts as a spike 2
"exactly when"

In words: "flag a step when its loss is more than f times the median of the last w losses."

With the numbers: median(2.9, 2.8, 2.8, 2.7) = 2.8, and 8.5 > 2 × 2.8 = 5.6, so step 5 is flagged; 4.0 < 5.6, so step 6 is not.

Level 3: in Python
import statistics
losses = [3.0, 2.9, 2.8, 2.8, 2.7, 8.5, 4.0, 2.7]
w, f = 4, 2.0
# L_t > f · median(L_{t−w}, …, L_{t−1})
[t for t in range(w, len(losses)) if losses[t] > f * statistics.median(losses[t - w:t])]  # → [5]

In code: detect_spikes applies the median rule to a whole loss history.

5b. Gradient clipping, when the gradient is spread across GPUs

Clipping by global norm (built in primer.ml.optimizers) caps the length of the whole gradient vector: if it is longer than c, shrink every entry by the same factor. At scale there is a twist. With ZeRO or tensor parallelism no GPU holds the whole gradient, so none can measure its length alone. Each GPU sums the squares of its own shard, one all-reduce adds those sums (a single number per GPU, so it is nearly free), and every GPU takes the square root and applies the same factor.

Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
the whole gradient, all parameters (1, 2, 2, 4)
entry of the shard on GPU GPU 1: (1, 2); GPU 2: (2, 4)
GPU 's local sum of squares 5 and 20
the length (norm) of the whole gradient 5
the clipping threshold 1
never scale up, only down 0.2
"is replaced by"

In words: "add up every GPU's sum of squares, take the square root to get the total length, and if it is over the limit, shrink every entry on every GPU by the same factor."

With the numbers: 5 + 20 = 25, √25 = 5; with c = 1 the factor is 1/5, so GPU 1's shard becomes (0.2, 0.4) and GPU 2's (0.4, 0.8).

Level 3: in Python
import math
shards = [[1.0, 2.0], [2.0, 4.0]]
# each GPU's local sum of squares
local = [sum(x * x for x in s) for s in shards]
local  # → [5.0, 20.0]
# all-reduce the sums, then the square root
norm = math.sqrt(sum(local))
norm  # → 5.0
c = 1.0
factor = min(1.0, c / norm)
[[x * factor for x in s] for s in shards]  # → [[0.2, 0.4], [0.4, 0.8]]

Figure 22 · Chart

0 10 20 30 40 50 60 training step 1 0 − 7 1 0 − 5 1 0 − 3 1 0 − 1 1 0 1 1 0 3 clean held-out loss One bad batch: clipping limits the damage corrupted batch no clipping clip global norm at 1

Clean held-out loss over 60 steps of training on y = 3x with one corrupted batch at step 30: without clipping the loss leaps above 6,000 and takes dozens of steps to come back; with clipping it barely moves

Reading it: the horizontal axis is the training step and the vertical axis the loss on clean data, on a log scale. Both runs have converged by step 30, when one batch arrives with corrupted labels. Without clipping (red), that single gradient is hundreds of times too large, the weight is thrown far off, and the loss jumps to over 6,000 before slowly recovering. With clipping at 1 (blue), the same batch can only move the weight by one learning-rate step, and the loss rises to about 0.01. Clipping cannot tell a bad batch from a good one; it just limits how much damage any one batch can do.

In code: sharded_global_norm computes the norm from per-GPU sums of squares, and train_through_a_bad_batch trains through a corrupted batch with or without primer.ml.optimizers.clip_by_global_norm.

5c. Warmup

Everyday picture Nobody floors the accelerator in a car they have never driven; they ease on until they know how it responds.

At the start of training the weights are random and Adam's running averages have seen only a handful of gradients, so its step sizes are unreliable. A full learning rate at step 1 can throw the model into a region it never recovers from. Warmup ramps the learning rate linearly from 0 to its peak over the first few hundred to few thousand steps; primer.ml.optimizers derives the warmup-then-cosine schedule with worked numbers. At scale it matters more, not less: bigger models and bigger batches tolerate smaller peak learning rates, and a spike early in a months-long run wastes the most.

5d. Checkpoints: how often to save

Everyday picture Saving a long document every few seconds wastes time on saving; saving once an hour risks losing an hour of work to a crash. Somewhere in between is the least total waste.

Tiny worked example Suppose writing a checkpoint pauses training for 1 minute, and the cluster suffers a failure every 180 minutes on average. Saving every 19 minutes spends 1/19 of the time saving, and each failure loses on average half an interval, 9.5 minutes, every 180 minutes. Both costs come to about 5.3%, and their total, 10.5%, is the smallest possible.

Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
time between checkpoints 19 minutes
time to write one checkpoint 1 minute
mean time between failures for the whole cluster 180 minutes
share of time spent saving 1/19
share of time redoing lost work: on average half an interval per failure 9.5/180
the interval with the least total waste (Young, 1974) 18.97 minutes

In words: "saving often wastes time saving and saving rarely wastes time redoing; the best interval is the square root of twice the save time times the time between failures."

With the numbers: √(2 × 1 × 180) = √360 = 18.97 minutes, and the waste is 1/18.97 + 18.97/360 = 0.053 + 0.053 = 10.5%. Cut the save time to 10 seconds (by writing asynchronously, in the background, from every GPU's shard at once) and the best interval falls to 7.7 minutes with only 4.3% waste.

Level 3: in Python
import math
C, M = 1.0, 180.0
# T* = √(2 C M)
T = math.sqrt(2 * C * M)
round(T, 2)  # → 18.97
# waste(T) = C/T + T/(2M)
round(C / T + T / (2 * M), 3)  # → 0.105
# a 10-second save
C = 1 / 6
round(math.sqrt(2 * C * M), 1), round(C / math.sqrt(2 * C * M) + math.sqrt(2 * C * M) / (2 * M), 3)  # → (7.7, 0.043)

Figure 23 · Chart

1 0 0 1 0 1 1 0 2 minutes between checkpoints (failure every 180 min) 0.0 0.1 0.2 0.3 0.4 0.5 share of time wasted Checkpoint interval: a valley at √(2·C·M) 19.0 min, 10.5% 7.7 min, 4.3% 1-minute save 10-second save

Share of time wasted against checkpoint interval for a 1-minute and a 10-second save with a failure every 3 hours: each curve is a valley whose floor is marked at the square-root interval

Reading it: the horizontal axis is the checkpoint interval in minutes (log scale) and the vertical axis the share of time lost. Each curve is a valley: on the left, saving too often; on the right, losing too much work per failure. The dots mark √(2CM), the bottom of each valley. A faster save moves the whole valley down and to the left, which is why large training systems invest heavily in fast, asynchronous checkpointing: at thousands of GPUs, failures are not an exception but the weather.

In code: wasted_fraction is the waste formula and optimal_checkpoint_interval is Young's square-root rule.

The rule above, with the spike in your hands.

The valley above, with C, M and T in your hands.

Test yourself

10 questions

Answer each one out loud or on paper before you open it. If you can explain it, you know it.

Question 1Why does a pretraining pipeline run deduplication after the quality filters, not before?Think it through, then reveal

Language ID and quality rules look at one page at a time, so they are cheap per page. Near-duplicate detection compares pages with one another, which is the expensive step. Running the cheap filters first means the expensive one sees far fewer pages.

Question 2How can MinHash estimate the overlap of two documents without comparing their contents?Think it through, then reveal

Under a random ordering of all shingles, two sets have the same first member with probability equal to their Jaccard similarity. A signature records each document's first member under k random hash functions, so the fraction of matching slots estimates the Jaccard. LSH then groups signatures by bands so only likely pairs are ever compared.

Question 3Why do duplicated documents hurt a model, when more data usually helps?Think it through, then reveal

A repeated document gets many times the training signal of any other, so the model memorizes it and tends to regurgitate it; the repeats also spend compute that would have taught something new, and copies of benchmark questions contaminate evaluations.

Question 4Where do the 16 bytes per parameter come from, and what do they mean for a 7B model?Think it through, then reveal

2 bytes for the bf16 weight, 2 for its gradient, and 12 for fp32 state: the master weight and Adam's two running averages. 16 × 7 × 10⁹ = 112 GB, more than one 80 GB GPU holds, before any activations.

Question 5What does each ZeRO stage shard, and what does it cost?Think it through, then reveal

Stage 1 shards the optimizer state, stage 2 also the gradients, stage 3 (FSDP) also the weights, dividing each by the number of GPUs. Stages 1 and 2 cost no more communication than plain data parallelism; stage 3 adds all-gathers of each layer's weights in both passes, about 1.5 times the traffic.

Question 6Why is tensor parallelism kept inside one server while pipeline parallelism spans servers?Think it through, then reveal

Tensor parallelism exchanges partial results inside every layer, so it needs the fastest links, which exist only between GPUs in the same server. Pipeline parallelism only passes activations at stage boundaries, a small and infrequent exchange that slower links between servers can carry.

Question 7What is the pipeline bubble, and how do you shrink it?Think it through, then reveal

The time stages sit idle while the pipeline fills and drains: (p − 1)/(m + p − 1) of the schedule for p stages and m micro-batches. More micro-batches shrink it (4 stages: 75% with 1, 8.6% with 32), as do schedules that interleave forward and backward passes.

Question 8Why does fp16 training need loss scaling while bf16 usually does not?Think it through, then reveal

fp16 has 5 exponent bits, so its smallest value is about 6 × 10⁻⁸ and many gradients underflow to zero; multiplying the loss by a large scale lifts them into range. bf16 keeps fp32's 8 exponent bits and therefore its range, giving up precision instead.

Question 9Why keep an fp32 copy of the weights if the maths runs in 16 bits?Think it through, then reveal

Late in training, updates are tiny compared with the weights. Next to 1.0 the bf16 grid spacing is about 0.008, so an update of 0.0001 rounds away completely, every step. Applying updates to an fp32 master copy keeps them.

Question 10How often should a large run write checkpoints?Think it through, then reveal

Roughly every √(2·C·M), where C is the time to save and M the mean time between failures: saving more often wastes time saving, less often wastes work redone after failures. Faster, asynchronous saves allow more frequent checkpoints and less waste.

Primary sources

The papers behind this lesson

Rae et al., Scaling Language Models: Methods, Analysis & Insights from Training Gopher (2021)

Among much else, published the simple quality rules (length, symbols, stop words) that many open pipelines still apply.

The paper ↗
Lee et al., Deduplicating Training Data Makes Language Models Better (2021)

Showed that exact and MinHash near-duplicate removal cuts verbatim memorization about tenfold without hurting quality.

The paper ↗
Penedo et al., The FineWeb Datasets: Decanting the Web for the Finest Text Data at Scale (2024)

Documented and ablated a full open curation pipeline over 96 Common Crawl snapshots, including the classifier-filtered FineWeb-Edu.

The paper ↗
Hoffmann et al., Training Compute-Optimal Large Language Models (2022)

Found that parameters and training tokens should grow together, about 20 tokens per parameter.

Read the annotated companion →The paper ↗
Shumailov et al., The Curse of Recursion: Training on Generated Data Makes Models Forget (2023)

Showed that models trained recursively on their own outputs lose the tails of the original distribution: model collapse.

The paper ↗
Rajbhandari et al., ZeRO: Memory Optimizations Toward Training Trillion Parameter Models (2019)

Introduced the 16-bytes-per-parameter accounting and the three stages of sharding the training state across data-parallel GPUs.

Read the annotated companion →The paper ↗
Shoeybi et al., Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism (2019)

Split transformer layers across GPUs by columns and rows, with one all-reduce per block.

Read the annotated companion →The paper ↗
Huang et al., GPipe: Efficient Training of Giant Neural Networks using Pipeline Parallelism (2018)

Split a model into stages fed by micro-batches, and analysed the resulting bubble.

The paper ↗
Micikevicius et al., Mixed Precision Training (2017)

Introduced the fp16 recipe: an fp32 master copy of the weights, loss scaling, and fp32 accumulation.

The paper ↗
Korthikanti et al., Reducing Activation Recomputation in Large Transformer Models (2022)

Counted the activation memory of a transformer layer and showed how to recompute only the parts that are cheap to recompute.

The paper ↗

Researcher's shelf

Further reading

  • Micikevicius et al., FP8 Formats for Deep Learning (2022): https://arxiv.org/abs/2209.05433
  • Kalamkar et al., A Study of BFLOAT16 for Deep Learning Training (2019): https://arxiv.org/abs/1905.12322
  • Zhao et al., PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel (2023): https://arxiv.org/abs/2304.11277
  • Narayanan et al., Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM (2021): https://arxiv.org/abs/2104.04473
  • Chen et al., Training Deep Nets with Sublinear Memory Cost (activation checkpointing, 2016): https://arxiv.org/abs/1604.06174
  • Chowdhery et al., PaLM: Scaling Language Modeling with Pathways (loss spikes and rollback, 2022): https://arxiv.org/abs/2204.02311
  • Touvron et al., LLaMA: Open and Efficient Foundation Language Models (data mixture, 2023): https://arxiv.org/abs/2302.13971
  • Llama Team, The Llama 3 Herd of Models (4D parallelism and failures at 16K GPUs, 2024): https://arxiv.org/abs/2407.21783
  • DeepSeek-AI, DeepSeek-V3 Technical Report (fp8 training, 2024): https://arxiv.org/abs/2412.19437
  • PyTorch automatic mixed precision: https://pytorch.org/docs/stable/amp.html
  • PyTorch FullyShardedDataParallel: https://pytorch.org/docs/stable/fsdp.html

About this lesson. This is the illustrated edition of a lesson from the open-source AI Primer. Its text, figures and numbers are generated from the Primer's source at commit c8d5c21, so the two always agree: the explanation, the code that builds it and the tests that prove it.