At a glance
Key takeaways
- 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.
- 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.
- Memory: Adam in mixed precision needs 16 bytes per parameter before activations (112 GB for 7B), so training needs many GPUs.
- 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)).
- 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.
- 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
flowchart LR W[Web crawl<br/>billions of pages] --> C[Curation<br/>language, quality,<br/>deduplication] C --> M[Mixture<br/>weights per source] M --> T[Tokenize<br/>trillions of tokens] T --> P[Parallel training<br/>data, tensor, pipeline] P --> MP[Mixed precision<br/>bf16 math, fp32 master] MP --> S[Stability<br/>clip, watch spikes,<br/>checkpoint] S --> B[Base model]
Chapter 1
Where the data comes from, and how it is cleaned
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
flowchart LR
H[HTML page] --> X[Extract<br/>main text]
X --> L{Language ID<br/>is it English?}
L -- no --> D1[drop, or route to<br/>that language]
L -- yes --> Q{Quality rules<br/>and classifier}
Q -- fail --> D2[drop]
Q -- pass --> E{Exact duplicate?<br/>hash of the text}
E -- yes --> D3[drop]
E -- no --> N{Near duplicate?<br/>MinHash + LSH}
N -- yes --> D4[drop]
N -- no --> K[Keep:<br/>goes to the mixture]
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
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
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
flowchart LR D[Page] --> S[5-word shingles] S --> H[k hash functions<br/>keep each minimum] H --> G[Signature<br/>k numbers] G --> B1[band 1: r numbers] --> K1[bucket] G --> B2[band 2] --> K2[bucket] G --> BB[band b] --> KB[bucket] K1 & K2 & KB --> C[Candidate pairs:<br/>pages sharing any bucket] C --> V[Check estimated Jaccard<br/>against the threshold]
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
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
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
flowchart LR R[Real data<br/>spread 1.0] --> F1[Fit model 1] F1 --> S1[Sample from model 1] S1 --> F2[Fit model 2] F2 --> S2[Sample from model 2] S2 --> FN[... model n] FN --> X[Spread shrinks:<br/>tails are forgotten first]
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
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
In code: recursive_gaussian_fit runs the fit-sample-refit loop and returns every generation's fitted spread.
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
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
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.
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
flowchart LR G0[GPU 0<br/>chunks A B C D] -->|one chunk per step| G1[GPU 1<br/>chunks A B C D] G1 -->|one chunk per step| G2[GPU 2<br/>chunks A B C D] G2 -->|one chunk per step| G3[GPU 3<br/>chunks A B C D] G3 -->|one chunk per step| G0
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
flowchart TB
subgraph S0["Stage 0: plain data parallel"]
A0["every GPU: weights + gradients + optimizer state"]
end
subgraph S1["Stage 1: shard the optimizer state"]
A1["every GPU: weights + gradients<br/>its 1/G of the optimizer state"]
end
subgraph S2["Stage 2: also shard the gradients"]
A2["every GPU: weights<br/>its 1/G of gradients and optimizer state"]
end
subgraph S3["Stage 3 = FSDP: shard everything"]
A3["every GPU: its 1/G of everything<br/>borrow each layer's weights just in time"]
end
S0 --> S1 --> S2 --> S3
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
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
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
flowchart LR X[X, full copy<br/>on both GPUs] --> A1["GPU 1: X · W1 left columns"] X --> A2["GPU 2: X · W1 right columns"] A1 --> R1[ReLU, locally] --> B1["· W2 top rows<br/>partial sum"] A2 --> R2[ReLU, locally] --> B2["· W2 bottom rows<br/>partial sum"] B1 & B2 --> AR[All-reduce:<br/>add the partial sums] --> Y[Y, full copy<br/>on both 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
flowchart LR MB[Micro-batches<br/>1, 2, 3, ...] --> S1[GPU 1<br/>layers 1-8] S1 -->|activations| S2[GPU 2<br/>layers 9-16] S2 -->|activations| S3[GPU 3<br/>layers 17-24] S3 -->|activations| S4[GPU 4<br/>layers 25-32] S4 -.->|gradients flow back| S1
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
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
Figure 15 · Chart
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
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
flowchart TB
subgraph DP["Data parallel: replicas see different data, all-reduce gradients"]
subgraph R1["Replica 1"]
direction LR
P1["Pipeline stage 1<br/>one server: 8 GPUs, tensor parallel"] --> P2["Pipeline stage 2<br/>one server: 8 GPUs, tensor parallel"]
end
subgraph R2["Replica 2"]
direction LR
Q1["Pipeline stage 1<br/>8 GPUs, tensor parallel"] --> Q2["Pipeline stage 2<br/>8 GPUs, tensor parallel"]
end
end
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
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
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
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
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
flowchart TB
M[fp32 master weights] -->|cast| W16[16-bit weights]
W16 --> F[Forward pass<br/>16-bit matmuls]
F --> L[Loss, in fp32]
L -->|times S| B[Backward pass<br/>16-bit gradients]
B --> CK{Any inf or NaN?}
CK -- yes --> SK[Skip the step<br/>halve S]
CK -- no --> U[Divide by S in fp32<br/>clip, then Adam update]
U --> M
SK --> M
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
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
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.
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
flowchart LR
T[Train step] --> W{Loss well above<br/>recent median?}
W -- no --> CP{Checkpoint<br/>interval reached?}
CP -- yes --> SV[Save weights,<br/>optimizer, data position]
CP -- no --> T
SV --> T
W -- "yes, and it persists" --> RB[Roll back to a<br/>checkpoint before the spike]
RB --> SKIP[Skip the batches<br/>around the spike]
SKIP --> T
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
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
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
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
In code: wasted_fraction is the waste formula and optimal_checkpoint_interval is Young's square-root rule.
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
Among much else, published the simple quality rules (length, symbols, stop words) that many open pipelines still apply.
The paper ↗Showed that exact and MinHash near-duplicate removal cuts verbatim memorization about tenfold without hurting quality.
The paper ↗Documented and ablated a full open curation pipeline over 96 Common Crawl snapshots, including the classifier-filtered FineWeb-Edu.
The paper ↗Found that parameters and training tokens should grow together, about 20 tokens per parameter.
Read the annotated companion →The paper ↗Showed that models trained recursively on their own outputs lose the tails of the original distribution: model collapse.
The paper ↗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 ↗Split transformer layers across GPUs by columns and rows, with one all-reduce per block.
Read the annotated companion →The paper ↗Split a model into stages fed by micro-batches, and analysed the resulting bubble.
The paper ↗Introduced the fp16 recipe: an fp32 master copy of the weights, loss scaling, and fp32 accumulation.
The paper ↗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.