primer.ml.pretraining
Pretraining at scale: from a web crawl to a base model on thousands of GPUs
Run: python -m primer.ml.pretraining
New to the notation? primer.notation explains every symbol used here
(Σ, log, subscripts, ‖x‖ and so on) from zero. This lesson deepens section 1
of primer.ml.training_stages, which shows what pretraining optimizes: the
average next-token loss. Here we open the two boxes that make it hard in
practice: the data (where trillions of tokens come from and how they are
cleaned) and the machine (how one training run is spread across
thousands of GPUs without running out of memory, precision or patience).
Level 1: The practitioner's guide
In one sentence. Pretraining is the stage that turns trillions of tokens of curated text into a base model on thousands of GPUs; you will almost never run it, but every choice made there (which data, how much, how clean) reaches you as what a model knows, which languages and code it handles, how well a given size performs, and the date its knowledge stops.
When you need it. You need this lesson the day you choose a model, and
again the day you build a corpus of your own. Choosing a model is mostly
reading the consequences of someone else's pretraining: the knowledge
cutoff, the languages in the mixture, how many tokens a model of that size
saw, and whether code was in the diet. Building a corpus (for retrieval,
for fine-tuning, or for continued pretraining) means running the same belt
the labs run: language identification, quality rules, exact and
near-duplicate removal. The tell that you are in pretraining territory:
the question is "does the model know X?" rather than "does it behave
well?", or a model is reciting a page of the web word for word. What you
do not need is to run pretraining yourself: the compute-optimal run for a
7-billion-parameter model is about 140 billion tokens and 4,100 GPU-hours
at a realistic 400 teraFLOP/s per GPU (this lesson's training_flops),
and the 405-billion-parameter Llama 3 run used up to 16,384 GPUs for 54
days with 419 unexpected interruptions (the Llama 3 paper, quoted in
section 5).
Your options. From the cheapest to the most committed:
| Option | What it does | What it guarantees | What it costs | Where it lives |
|---|---|---|---|---|
| Use a hosted model as-is | Someone else's pretraining, tuning and serving, behind an API | The best general knowledge you can buy; a cutoff and a mixture you did not choose | Per-token prices, no control over the data | The vendor |
| Use an open-weights model as-is | Pick a checkpoint whose card states its tokens, mixture and cutoff | The same, with the model on your hardware and the card to read | Serving, and a licence to check | Your servers |
| Curate your own corpus with the pretraining toolkit | Language ID, quality rules, hashing and MinHash over your documents before they reach retrieval or a fine-tune | No duplicates, no junk, no benchmark leaks in your data | A pipeline run; the filters are cheap, deduplication is the expensive step | Your data pipeline |
| Continued pretraining on domain text | Keep training an existing base model on billions of tokens of your field, with the next-token loss | Domain vocabulary and facts learned the way general knowledge was; the model still needs its assistant tuning afterwards | GPUs for days or weeks, a data mixture that keeps some general text, evaluations for what it forgot | Your training stack |
| Train a small model from scratch | Curate, mix, tokenize, then run the parallel training loop | Full control of the data and the cutoff | About 20 tokens per parameter for a compute-optimal model, far more if it will serve billions of requests; the whole engineering of section 2 onwards | A cluster, and a team |
How to choose. Start from what the model has to know, then from what it will cost to serve.
- Missing knowledge that changes or must be cited: retrieval, not any form
of training (
primer.ml.training_stagesmakes the case). - Choosing between two models of the same size: prefer the one trained on more tokens of cleaner data. Past the compute-optimal 20 tokens per parameter, a smaller model trained longer is cheaper to run forever after, which is why Llama 2 7B saw about 286 tokens per parameter and Llama 3 (the whole family) about 15 trillion tokens (section 1k).
- A field with its own vocabulary that prompting and a small fine-tune cannot cover (a language of contracts, a scientific literature): continued pretraining, once you hold billions of domain tokens. Gururangan et al. (2020) found a second phase of pretraining on domain text improves the tasks in that domain, across four domains and eight tasks.
- Any corpus of your own: deduplicate and filter before you train or index. The lesson's toy belt keeps 2 of 7 crawled pages, and the large public pipelines keep only a small fraction of the crawl they start from.
- Whatever you pick, the rule holds: data decides what the model learns, and compute cannot put back what the data never held.
What it costs. Data comes first. FineWeb, an open pipeline over 96 Common Crawl snapshots, ends at about 15 trillion tokens after filtering and deduplication (section 1). Compute follows a rule of thumb from this lesson: about six floating-point operations per parameter per token, so the 7B model at 140 billion tokens costs 5.88 × 10²¹ operations. Memory is the reason it takes a cluster: Adam in mixed precision holds 16 bytes per parameter, 112 GB for 7B, before any activation, and one 80 GB GPU cannot hold it (section 2). Sharding the state across 64 GPUs brings it under 2 GB each (the ZeRO paper's example), at the price of communication that the frameworks hide. Failures are the weather at that scale: with a one-minute checkpoint and a failure every three hours, the least you can waste is 10.5% of the run, and a ten-second asynchronous save cuts that to 4.3% (section 5d). None of this is your bill, but all of it is why a frontier model's price per token is what it is, and why open checkpoints are released at the sizes they are.
What breaks.
- Duplicates. A page kept 100 times is recited: in the lesson's bigram toy, the chance of regurgitating a boilerplate line goes from 0.014 to 0.93, and Lee et al. (2021) measured about ten times less memorized text after deduplication. Deduplicate any corpus you train on.
- Contamination. A benchmark question copied across the web ends up in the training data, and the score on it measures recall. Check your evaluation set against your training set with the same near-duplicate tools.
- Filters with blind spots. A quality classifier keeps what resembles its reference set; pick encyclopedia text alone as "good" and you filter out dialects, forums and whole topics (section 1c).
- Model collapse. Training on a model's own outputs loses the rare values first: refitting on 20 samples keeps 95% of the spread per generation, and 200 generations leave 0.0035% of the variance. Keep real data in every mix and filter synthetic data with checks that do not come from the same model.
- The cutoff. A base model's knowledge is frozen at its crawl date. Retrieve what changes.
- Numerics, if you do train. fp16 underflows gradients below about 6 × 10⁻⁸ without loss scaling; bf16 keeps fp32's range and became the default. A tiny update rounds away in 16 bits, so the master weights stay in fp32 (section 4).
In the wild. Common Crawl is the raw material for nearly every open pretraining corpus; FineWeb and FineWeb-Edu (Penedo et al., 2024) are open curations of it, built with the datatrove library, whose pipeline blocks include filters and MinHash deduplication. The quality rules are the ones the Gopher paper published, and fastText's language identifier covers 176 languages. On the machine side, DeepSpeed implements the ZeRO stages, PyTorch FSDP is stage 3, Megatron-LM is the tensor-parallel split, GPipe introduced the micro-batch pipeline, and PyTorch's automatic mixed precision runs the bf16 loop with its master copy. Llama 3 nests all of these with a fourth kind, context parallelism, across 16,384 GPUs; DeepSeek-V3 trained largely in fp8. The papers are linked at the end of the lesson.
Go deeper. Level 2 builds the whole belt on seven crawled pages: a stop-word language detector, the Gopher rules, a naive Bayes quality classifier, MinHash and locality-sensitive hashing with their S-curve, then the memory bill of a 7B model, ring all-reduce, ZeRO's stages, tensor and pipeline parallelism with the bubble counted cell by cell, floating-point formats bit by bit, and the checkpoint interval as a square root. If you only needed to choose a model or clean a corpus, you are done.
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.
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]
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.
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.
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]
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
$$ \text{score}(\text{doc}) = \sum_{w \in \text{doc}} \log \frac{P(w \mid \text{good})}{P(w \mid \text{bad})}, \qquad P(w \mid c) = \frac{\text{count}_c(w) + 1}{N_c + V} $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $w$ | one word of the document | "river" |
| $\sum_{w \in \text{doc}}$ | add up the term for every word in the document | |
| $c$ | a class: good (reference) or bad (spam) | |
| $\text{count}_c(w)$ | how often $w$ appears in class $c$'s examples | 3 in good, 0 in bad |
| $N_c$ | total words in class $c$'s examples | 77 good, 45 bad |
| $V$ | number of distinct words across both classes (the vocabulary) | 80 |
| $+1$ | add-one smoothing: pretend every word was seen once more, so an unseen word never gives probability 0 (and log 0 = −∞) | |
| $P(w \mid c)$ | "probability of $w$ given $c$": how likely a word drawn from class $c$ is $w$ | $P(\text{river} \mid \text{good}) = 4/157$ |
| $\log$ | natural logarithm: turns a ratio above 1 into a positive vote and below 1 into a negative vote | $\log 3.185 = 1.158$ |
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
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
$$ J(A, B) = \frac{|A \cap B|}{|A \cup B|} $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $A$, $B$ | the two sets of shingles | 5 shingles each |
| $A \cap B$ | the intersection: shingles in both sets | {the cat, cat sat, sat on} |
| $A \cup B$ | the union: shingles in either set, each counted once | 7 shingles |
| $\lvert \cdot \rvert$ | the number of items in a set | $\lvert A \cap B \rvert = 3$ |
| $J(A, B)$ | 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
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
$$ P\big[\min h(A) = \min h(B)\big] = J(A, B), \qquad \hat{J} = \frac{1}{k} \sum_{i=1}^{k} \mathbf{1}\big[\min h_i(A) = \min h_i(B)\big] $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $h$ | 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 |
| $\min h(A)$ | the smallest number $h$ gives any shingle of $A$: A's "first card" | |
| $P[\ldots]$ | the probability that the statement in brackets is true | 12/24 |
| $k$ | how many independent hash functions (the signature length) | 128 |
| $h_i$ | the $i$-th hash function | |
| $\mathbf{1}[\ldots]$ | the indicator: 1 if the statement is true, 0 if not | |
| $\hat{J}$ | the estimate of $J$ (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
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
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 $b$ bands of $r$ 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.
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]
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
$$ P(\text{candidate}) = 1 - \left(1 - s^{r}\right)^{b} $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $s$ | the pair's Jaccard similarity | 0.8 or 0.3 |
| $r$ | rows: numbers per band | 5 |
| $b$ | number of bands | 20 |
| $s^r$ | chance that all $r$ numbers of one band agree (each agrees with chance $s$) | $0.8^5 = 0.328$ |
| $1 - s^r$ | chance that one band does not fully agree | 0.672 |
| $(1 - s^r)^b$ | chance that no band agrees | $0.672^{20} = 0.00036$ |
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
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
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
$$ \text{epochs}_i = \frac{w_i \, D}{N_i} $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $i$ | one data source | the encyclopedia |
| $w_i$ | its mixture weight: share of all training tokens drawn from it | 0.05 |
| $D$ | the total training budget, in tokens | 1,000 billion |
| $N_i$ | the tokens that source actually has | 20 billion |
| $\text{epochs}_i$ | 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
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
$$ D_{\text{opt}} \approx 20\,N, \qquad C \approx 6\,N\,D $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $N$ | the number of model parameters | 7 billion |
| $D$ | the number of training tokens | 140 billion |
| $D_{\text{opt}}$ | the compute-optimal token count for size $N$ (Hoffmann et al., 2022) | $20 \times 7 \times 10^9$ |
| $C$ | training compute in FLOPs (floating-point operations) | $5.88 \times 10^{21}$ |
| $6$ | 2 FLOPs per parameter per token in the forward pass (one multiply, one add), 4 in the backward pass | |
| $\approx$ | "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
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.
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]
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
$$ \mathbb{E}\big[\hat{\sigma}^2_{t+1}\big] = \frac{n - 1}{n}\,\hat{\sigma}^2_t $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $t$ | the generation number | 0, 1, 2, … |
| $\hat{\sigma}^2_t$ | the fitted variance (spread squared) of generation $t$'s model | 1.0 at the start |
| $n$ | samples each generation is fitted on | 20 |
| $\mathbb{E}[\ldots]$ | the expected value: the average over many repeats of the experiment | |
| $\frac{n-1}{n}$ | 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
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
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.
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
$$ M_{\text{state}} = (2 + 2 + 4 + 4 + 4)\,\Psi = 16\,\Psi \ \text{bytes} $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $\Psi$ | the number of parameters (Psi, the letter the ZeRO paper uses) | $7 \times 10^9$ |
| $2 + 2$ | bytes for the bf16 weight and its bf16 gradient | |
| $4 + 4 + 4$ | bytes for the fp32 master weight and Adam's two fp32 averages | |
| $M_{\text{state}}$ | 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
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
$$ A = L \cdot s\,b\,h \left(34 + \frac{5\,a\,s}{h}\right) \ \text{bytes} $$
Symbols
| Symbol | Meaning here | For a 7B shape |
|---|---|---|
| $L$ | number of layers | 32 |
| $s$ | sequence length in tokens | 4,096 |
| $b$ | sequences per GPU in the batch | 1 |
| $h$ | hidden width | 4,096 |
| $a$ | attention heads | 32 |
| $34\,s\,b\,h$ | the layer's ordinary tensors (inputs to each matrix multiply, norms, dropout masks) | |
| $5\,a\,s^2 b$ | attention's $s \times s$ score and weight matrices, one per head (written as $s b h \cdot 5as/h$) |
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
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.
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.
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
$$ \nabla \mathcal{L} = \frac{1}{G} \sum_{k=1}^{G} \nabla \mathcal{L}_k $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $\mathcal{L}$ | the loss averaged over the whole batch | |
| $\nabla$ | "the gradient of": the slope of the loss for every weight | |
| $G$ | number of GPUs, each with an equal share of the batch | 2 |
| $\mathcal{L}_k$ | the loss averaged over GPU $k$'s share | |
| $\sum_{k=1}^{G}$ | 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
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.
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
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
$$ \text{sent per GPU} = \frac{2\,(G - 1)}{G}\,S $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $G$ | GPUs in the ring | 4, or 1,000 |
| $S$ | size of the gradient being summed | 14 GB (7B in bf16) |
| $G - 1$ | steps in each of the two phases | 3 |
| $\frac{1}{G}$ | each step moves one chunk, $1/G$ of the gradient | |
| $2$ | 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
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.
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
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
$$ M_{\text{stage 0}} = 16\Psi, \quad M_{1} = 4\Psi + \frac{12\Psi}{G}, \quad M_{2} = 2\Psi + \frac{14\Psi}{G}, \quad M_{3} = \frac{16\Psi}{G} $$
Symbols
| Symbol | Meaning here | In the ZeRO paper's example |
|---|---|---|
| $\Psi$ | parameters | $7.5 \times 10^9$ |
| $G$ | GPUs sharing the state | 64 |
| $12\Psi$ | the fp32 master weights and Adam's two averages | 90 GB |
| $2\Psi$, $4\Psi$ | the bf16 weights; weights plus gradients | 15 GB; 30 GB |
| $M_{\text{stage}}$ | 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
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.
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
$$ X W = \big[\,X W_1 \;\; X W_2\,\big] \qquad\text{and}\qquad X W = X_1 W_1 + X_2 W_2 $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $X$ | the input activations, one row per token | [1, 2] |
| $W$ | a weight matrix | 2 × 4 |
| $[\,A \;\; B\,]$ | place two matrices side by side (concatenate columns) | |
| left: $W_1, W_2$ | $W$'s columns, split into two blocks | 2 × 2 each |
| right: $W_1, W_2$ | $W$'s rows, split into two blocks | 1 × 4 each |
| right: $X_1, X_2$ | $X$'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
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.
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]
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.
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
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
$$ \text{bubble} = \frac{p - 1}{m + p - 1} $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $p$ | pipeline stages (GPUs in the line) | 4 |
| $m$ | micro-batches per batch | 1, 8 or 32 |
| $m + p - 1$ | time steps to push $m$ micro-batches through $p$ stages: $p$ to fill the line, then one per extra micro-batch | 11 when $m = 8$ |
| $p - 1$ | 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
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]
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.
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
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
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.
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
$$ x = (-1)^{\text{sign}} \times 2^{\,E - \text{bias}} \times \left(1 + \frac{F}{2^{m}}\right) $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| sign | 1 bit: 0 for positive, 1 for negative | 0 |
| $(-1)^{\text{sign}}$ | +1 or −1 | +1 |
| $E$ | the exponent field, read as a whole number | 125 |
| bias | a fixed offset, $2^{e-1} - 1$ for $e$ exponent bits, so negative exponents can be stored | 127 (bf16 and fp32) |
| $F$ | the fraction field, read as a whole number | 43 |
| $m$ | the number of fraction bits | 7 |
| $1 + F/2^m$ | 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
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.
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
$$ g = \frac{\operatorname{fp16}!\left(S \cdot \nabla \mathcal{L}\right)}{S} $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $\nabla \mathcal{L}$ | the true gradient of one weight | 10⁻⁸ |
| $S$ | the loss scale, usually a power of two so multiplying is exact | 65,536 = 2¹⁶ |
| $\operatorname{fp16}(\ldots)$ | "rounded to the nearest fp16 value", which is where the backward pass stores it | |
| $g$ | 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
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'
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
- 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.
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
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.
Level 3: in Python
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)
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.
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.
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
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
$$ \text{spike at } t \iff \mathcal{L}_t > f \cdot \operatorname{median}\big(\mathcal{L}_{t-w}, \ldots, \mathcal{L}_{t-1}\big) $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $\mathcal{L}_t$ | the training loss at step $t$ | 8.5 at $t = 5$ |
| $w$ | how many recent steps to look back over (the window) | 4 |
| $\operatorname{median}$ | the middle value after sorting; one outlier cannot move it far | 2.8 |
| $f$ | how many times the median counts as a spike | 2 |
| $\iff$ | "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
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
$$ \lVert g \rVert = \sqrt{\sum_{k=1}^{G} \sum_{j} g_{kj}^2}, \qquad g \leftarrow g \cdot \min!\left(1, \frac{c}{\lVert g \rVert}\right) $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $g$ | the whole gradient, all parameters | (1, 2, 2, 4) |
| $g_{kj}$ | entry $j$ of the shard on GPU $k$ | GPU 1: (1, 2); GPU 2: (2, 4) |
| $\sum_j g_{kj}^2$ | GPU $k$'s local sum of squares | 5 and 20 |
| $\lVert g \rVert$ | the length (norm) of the whole gradient | 5 |
| $c$ | the clipping threshold | 1 |
| $\min(1, \ldots)$ | never scale up, only down | 0.2 |
| $\leftarrow$ | "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
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]]
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
$$ \text{waste}(T) = \frac{C}{T} + \frac{T}{2M}, \qquad T^{*} = \sqrt{2\,C\,M} $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $T$ | time between checkpoints | 19 minutes |
| $C$ | time to write one checkpoint | 1 minute |
| $M$ | mean time between failures for the whole cluster | 180 minutes |
| $C/T$ | share of time spent saving | 1/19 |
| $T/(2M)$ | share of time redoing lost work: on average half an interval per failure | 9.5/180 |
| $T^{*}$ | 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
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)
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.
In 20 seconds
- 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).
Self-test questions
Why does a pretraining pipeline run deduplication after the quality filters, not before? 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.
How can MinHash estimate the overlap of two documents without comparing their contents? 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.
Why do duplicated documents hurt a model, when more data usually helps? 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.
Where do the 16 bytes per parameter come from, and what do they mean for a 7B model? 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.
What does each ZeRO stage shard, and what does it cost? 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.
Why is tensor parallelism kept inside one server while pipeline parallelism spans servers? 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.
What is the pipeline bubble, and how do you shrink it? 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.
Why does fp16 training need loss scaling while bf16 usually does not? 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.
Why keep an fp32 copy of the weights if the maths runs in 16 bits? 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.
How often should a large run write checkpoints? 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.
The papers behind this lesson
- Rae et al., Scaling Language Models: Methods, Analysis & Insights from Training Gopher (2021): https://arxiv.org/abs/2112.11446. Among much else, published the simple quality rules (length, symbols, stop words) that many open pipelines still apply.
- Lee et al., Deduplicating Training Data Makes Language Models Better (2021): https://arxiv.org/abs/2107.06499. Showed that exact and MinHash near-duplicate removal cuts verbatim memorization about tenfold without hurting quality.
- Penedo et al., The FineWeb Datasets: Decanting the Web for the Finest Text Data at Scale (2024): https://arxiv.org/abs/2406.17557. Documented and ablated a full open curation pipeline over 96 Common Crawl snapshots, including the classifier-filtered FineWeb-Edu.
- Hoffmann et al., Training Compute-Optimal Large Language Models (2022): https://arxiv.org/abs/2203.15556. Found that parameters and training tokens should grow together, about 20 tokens per parameter. Annotated companion
- Shumailov et al., The Curse of Recursion: Training on Generated Data Makes Models Forget (2023): https://arxiv.org/abs/2305.17493. Showed that models trained recursively on their own outputs lose the tails of the original distribution: model collapse.
- Rajbhandari et al., ZeRO: Memory Optimizations Toward Training Trillion Parameter Models (2019): https://arxiv.org/abs/1910.02054. Introduced the 16-bytes-per-parameter accounting and the three stages of sharding the training state across data-parallel GPUs. Annotated companion
- Shoeybi et al., Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism (2019): https://arxiv.org/abs/1909.08053. Split transformer layers across GPUs by columns and rows, with one all-reduce per block. Annotated companion
- Huang et al., GPipe: Efficient Training of Giant Neural Networks using Pipeline Parallelism (2018): https://arxiv.org/abs/1811.06965. Split a model into stages fed by micro-batches, and analysed the resulting bubble.
- Micikevicius et al., Mixed Precision Training (2017): https://arxiv.org/abs/1710.03740. Introduced the fp16 recipe: an fp32 master copy of the weights, loss scaling, and fp32 accumulation.
- Korthikanti et al., Reducing Activation Recomputation in Large Transformer Models (2022): https://arxiv.org/abs/2205.05198. Counted the activation memory of a transformer layer and showed how to recompute only the parts that are cheap to recompute.
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
1r""" 2# Pretraining at scale: from a web crawl to a base model on thousands of GPUs 3 4Run: `python -m primer.ml.pretraining` 5 6New to the notation? `primer.notation` explains every symbol used here 7(Σ, log, subscripts, ‖x‖ and so on) from zero. This lesson deepens section 1 8of `primer.ml.training_stages`, which shows *what* pretraining optimizes: the 9average next-token loss. Here we open the two boxes that make it hard in 10practice: **the data** (where trillions of tokens come from and how they are 11cleaned) and **the machine** (how one training run is spread across 12thousands of GPUs without running out of memory, precision or patience). 13 14## Level 1: The practitioner's guide 15 16**In one sentence.** Pretraining is the stage that turns trillions of tokens 17of curated text into a base model on thousands of GPUs; you will almost never 18run it, but every choice made there (which data, how much, how clean) reaches 19you as what a model knows, which languages and code it handles, how well a 20given size performs, and the date its knowledge stops. 21 22**When you need it.** You need this lesson the day you choose a model, and 23again the day you build a corpus of your own. Choosing a model is mostly 24reading the consequences of someone else's pretraining: the knowledge 25cutoff, the languages in the mixture, how many tokens a model of that size 26saw, and whether code was in the diet. Building a corpus (for retrieval, 27for fine-tuning, or for continued pretraining) means running the same belt 28the labs run: language identification, quality rules, exact and 29near-duplicate removal. The tell that you are in pretraining territory: 30the question is "does the model know X?" rather than "does it behave 31well?", or a model is reciting a page of the web word for word. What you 32do not need is to run pretraining yourself: the compute-optimal run for a 337-billion-parameter model is about 140 billion tokens and 4,100 GPU-hours 34at a realistic 400 teraFLOP/s per GPU (this lesson's `training_flops`), 35and the 405-billion-parameter Llama 3 run used up to 16,384 GPUs for 54 36days with 419 unexpected interruptions (the Llama 3 paper, quoted in 37section 5). 38 39**Your options.** From the cheapest to the most committed: 40 41| Option | What it does | What it guarantees | What it costs | Where it lives | 42|---|---|---|---|---| 43| Use a hosted model as-is | Someone else's pretraining, tuning and serving, behind an API | The best general knowledge you can buy; a cutoff and a mixture you did not choose | Per-token prices, no control over the data | The vendor | 44| Use an open-weights model as-is | Pick a checkpoint whose card states its tokens, mixture and cutoff | The same, with the model on your hardware and the card to read | Serving, and a licence to check | Your servers | 45| Curate your own corpus with the pretraining toolkit | Language ID, quality rules, hashing and MinHash over your documents before they reach retrieval or a fine-tune | No duplicates, no junk, no benchmark leaks in your data | A pipeline run; the filters are cheap, deduplication is the expensive step | Your data pipeline | 46| Continued pretraining on domain text | Keep training an existing base model on billions of tokens of your field, with the next-token loss | Domain vocabulary and facts learned the way general knowledge was; the model still needs its assistant tuning afterwards | GPUs for days or weeks, a data mixture that keeps some general text, evaluations for what it forgot | Your training stack | 47| Train a small model from scratch | Curate, mix, tokenize, then run the parallel training loop | Full control of the data and the cutoff | About 20 tokens per parameter for a compute-optimal model, far more if it will serve billions of requests; the whole engineering of section 2 onwards | A cluster, and a team | 48 49**How to choose.** Start from what the model has to know, then from what it 50will cost to serve. 51 52- Missing knowledge that changes or must be cited: retrieval, not any form 53 of training (`primer.ml.training_stages` makes the case). 54- Choosing between two models of the same size: prefer the one trained on 55 more tokens of cleaner data. Past the compute-optimal 20 tokens per 56 parameter, a smaller model trained longer is cheaper to run forever after, 57 which is why Llama 2 7B saw about 286 tokens per parameter and Llama 3 (the whole family) about 15 trillion tokens 58 (section 1k). 59- A field with its own vocabulary that prompting and a small fine-tune 60 cannot cover (a language of contracts, a scientific literature): continued 61 pretraining, once you hold billions of domain tokens. Gururangan et al. 62 (2020) found a second phase of pretraining on domain text improves the 63 tasks in that domain, across four domains and eight tasks. 64- Any corpus of your own: deduplicate and filter before you train or index. 65 The lesson's toy belt keeps 2 of 7 crawled pages, and the large public 66 pipelines keep only a small fraction of the crawl they start from. 67- Whatever you pick, the rule holds: data decides what the model learns, and 68 compute cannot put back what the data never held. 69 70**What it costs.** Data comes first. FineWeb, an open pipeline over 96 71Common Crawl snapshots, ends at about 15 trillion tokens after filtering and 72deduplication (section 1). Compute follows a rule of thumb from this lesson: 73about six floating-point operations per parameter per token, so the 7B model 74at 140 billion tokens costs 5.88 × 10²¹ operations. Memory is the reason it 75takes a cluster: Adam in mixed precision holds 16 bytes per parameter, 112 76GB for 7B, before any activation, and one 80 GB GPU cannot hold it (section 772). Sharding the state across 64 GPUs brings it under 2 GB each (the ZeRO 78paper's example), at the price of communication that the frameworks hide. 79Failures are the weather at that scale: with a one-minute checkpoint and a 80failure every three hours, the least you can waste is 10.5% of the run, and 81a ten-second asynchronous save cuts that to 4.3% (section 5d). None of this 82is your bill, but all of it is why a frontier model's price per token is 83what it is, and why open checkpoints are released at the sizes they are. 84 85**What breaks.** 86 87- **Duplicates.** A page kept 100 times is recited: in the lesson's bigram 88 toy, the chance of regurgitating a boilerplate line goes from 0.014 to 89 0.93, and Lee et al. (2021) measured about ten times less memorized text 90 after deduplication. Deduplicate any corpus you train on. 91- **Contamination.** A benchmark question copied across the web ends up in 92 the training data, and the score on it measures recall. Check your 93 evaluation set against your training set with the same near-duplicate 94 tools. 95- **Filters with blind spots.** A quality classifier keeps what resembles 96 its reference set; pick encyclopedia text alone as "good" and you filter 97 out dialects, forums and whole topics (section 1c). 98- **Model collapse.** Training on a model's own outputs loses the rare 99 values first: refitting on 20 samples keeps 95% of the spread per 100 generation, and 200 generations leave 0.0035% of the variance. Keep real 101 data in every mix and filter synthetic data with checks that do not come 102 from the same model. 103- **The cutoff.** A base model's knowledge is frozen at its crawl date. 104 Retrieve what changes. 105- **Numerics, if you do train.** fp16 underflows gradients below about 106 6 × 10⁻⁸ without loss scaling; bf16 keeps fp32's range and became the 107 default. A tiny update rounds away in 16 bits, so the master weights stay 108 in fp32 (section 4). 109 110**In the wild.** Common Crawl is the raw material for nearly every open 111pretraining corpus; FineWeb and FineWeb-Edu (Penedo et al., 2024) are open 112curations of it, built with the datatrove library, whose pipeline blocks 113include filters and MinHash deduplication. The quality rules are the ones 114the Gopher paper published, and fastText's language identifier covers 176 115languages. On the machine side, DeepSpeed implements the ZeRO stages, 116PyTorch FSDP is stage 3, Megatron-LM is the tensor-parallel split, GPipe 117introduced the micro-batch pipeline, and PyTorch's automatic mixed 118precision runs the bf16 loop with its master copy. Llama 3 nests all of 119these with a fourth kind, context parallelism, across 16,384 GPUs; 120DeepSeek-V3 trained largely in fp8. The papers are linked at the end of the 121lesson. 122 123**Go deeper.** Level 2 builds the whole belt on seven crawled pages: a 124stop-word language detector, the Gopher rules, a naive Bayes quality 125classifier, MinHash and locality-sensitive hashing with their S-curve, then 126the memory bill of a 7B model, ring all-reduce, ZeRO's stages, tensor and 127pipeline parallelism with the bubble counted cell by cell, floating-point 128formats bit by bit, and the checkpoint interval as a square root. If you 129only needed to choose a model or clean a corpus, you are done. 130 131## Level 2: How it works, from scratch 132 133Imagine writing an encyclopedia by reading everything ever printed. Two 134problems appear at once. 135 136First, most of what is printed is junk: flyers, receipts, the same cookie 137notice on a million websites, spam. You need a sorting line that throws out 138the junk and the photocopies before anyone reads a page. 139 140Second, no single reader can do the reading. You hire a thousand readers, 141and now you have a management problem: how do they split the work, share 142what they learned, and keep going when one of them gets sick? 143 144Pretraining is exactly those two problems. The first is data curation. The 145second is distributed training. 146 147```mermaid 148flowchart LR 149 W[Web crawl<br/>billions of pages] --> C[Curation<br/>language, quality,<br/>deduplication] 150 C --> M[Mixture<br/>weights per source] 151 M --> T[Tokenize<br/>trillions of tokens] 152 T --> P[Parallel training<br/>data, tensor, pipeline] 153 P --> MP[Mixed precision<br/>bf16 math, fp32 master] 154 MP --> S[Stability<br/>clip, watch spikes,<br/>checkpoint] 155 S --> B[Base model] 156``` 157 158**Reading it:** the left half of the chain (crawl, curation, mixture, 159tokenize) decides *what* the model learns; most of a model's knowledge and 160many of its quirks are settled here, before any GPU is switched on. The right 161half (parallel training, mixed precision, stability) decides whether the 162learning can happen at all: a 7-billion-parameter model does not fit on one 163GPU, and a months-long run on thousands of GPUs will see hardware fail many 164times. Every box gets its own section below, in this order. 165 166## 1. Where the data comes from, and how it is cleaned 167 168**Everyday picture.** Think of a recycling plant. Trucks tip mixed rubbish 169onto a conveyor belt. Magnets pull out the steel, blowers lift out the paper, 170people pick out what the machines miss, and only a small fraction reaches the 171bale at the end. Pretraining data goes down a belt like this, and just as in 172a recycling plant, most of what goes in never comes out. 173 174The raw material is usually **Common Crawl**, a nonprofit's public archive of 175the web: regular snapshots, each of billions of pages. A page arrives as 176HTML full of menus, adverts and scripts, so the first machine on the belt is 177**text extraction**, which keeps the main body text. After that come the 178filters this section builds: language identification, quality filters, 179deduplication and mixing. FineWeb, an open dataset built this way from 96 180Common Crawl snapshots, ends with about 15 trillion tokens. 181 182```mermaid 183flowchart LR 184 H[HTML page] --> X[Extract<br/>main text] 185 X --> L{Language ID<br/>is it English?} 186 L -- no --> D1[drop, or route to<br/>that language] 187 L -- yes --> Q{Quality rules<br/>and classifier} 188 Q -- fail --> D2[drop] 189 Q -- pass --> E{Exact duplicate?<br/>hash of the text} 190 E -- yes --> D3[drop] 191 E -- no --> N{Near duplicate?<br/>MinHash + LSH} 192 N -- yes --> D4[drop] 193 N -- no --> K[Keep:<br/>goes to the mixture] 194``` 195 196**Reading it:** each diamond is a filter and each "drop" box is a way a page 197leaves the belt. The order matters for cost: language ID and quality rules 198look at one page at a time, so they are cheap and run first. Deduplication 199compares pages *with each other*, which is the expensive part, so it runs 200last, on what survived. The rest of this section builds every diamond in turn. 201 202### 1a. Language identification 203 204**Everyday picture.** Overhear two words of a phone call, "le" and "et", and 205you already guess French. Every language has a handful of little words that 206turn up in almost every sentence. 207 208**Tiny worked example.** "the cat is in the garden and it was happy" has 10 209words; 7 of them are on the English list of little words (the, is, in, the, 210and, it, was) and none are on the French, German or Spanish lists, so the 211guess is English with a share of 0.7. "SKU-4431 X99 blk/wht 12pk" matches no 212list at all, so it is marked unknown and dropped. 213 214Production pipelines use a trained classifier over character sequences 215(fastText's language ID model covers 176 languages) and keep a page only if 216the classifier is confident. The idea is the same: count the evidence for 217each language and pick the strongest. 218 219**In code:** `detect_language` counts each language's stop words (`STOP_WORDS`) and returns the language with the largest share, or "unknown" below 10%. 220 221### 1b. Heuristic quality filters 222 223**Everyday picture.** A librarian sorting donations does not read every 224book. A book with no pages, a pamphlet that is all hashtags, a sheet of 225keywords: each is rejected at a glance by a simple rule. 226 227**Tiny worked example.** The Gopher paper (Rae et al., 2021) published a 228set of such rules. Here is what they say about a navigation bar, "Home | 229About us | Contact | Privacy policy | Log in": 230 231| Rule | Keep only if | The navigation bar | Verdict | 232|---|---|---|---| 233| length | 50 to 100,000 words | 12 (each "\|" counts as a word) | fail | 234| mean word length | 3 to 10 characters | 3.3 | pass | 235| symbols (#, ...) per word | at most 0.1 | 0 | pass | 236| words containing a letter | at least 80% | 8 of 12 = 67% | fail | 237| common English words | at least 2 of: the, be, to, of, and, that, have, with | none | fail | 238 239The last rule is the clever one. Real prose cannot avoid words like "the" 240and "of"; keyword-stuffed pages and lists of product names avoid them 241completely. Other rules in the set catch pages that are mostly bullet points 242or mostly lines ending in "...". 243 244**In code:** `quality_failures` applies each rule and returns the names of the ones a document breaks; an empty list means keep. 245 246### 1c. Quality classifiers: scoring pages against a reference 247 248Rules catch obvious junk. To prefer *good* text among the survivors, 249pipelines train a classifier: examples of the text you want (encyclopedia 250articles, books, pages that trusted sites link to) against random crawl, and 251then keep pages that score like the reference. GPT-3 filtered Common Crawl 252with a simple linear classifier of this kind; FineWeb-Edu asked a large 253model to rate pages for educational value and trained a small classifier on 254those ratings. We build the simplest version, **naive Bayes**: every word 255casts a vote, and the votes add up. 256 257**Tiny worked example.** Our reference examples contain the word "river" 3 258times and the spam examples 0 times; "click" appears 0 times in the 259reference and 4 times in spam. So "river" votes for good and "click" votes 260for bad. A whole page's score is the sum of its words' votes. 261 262$$ 263\text{score}(\text{doc}) = \sum_{w \in \text{doc}} \log \frac{P(w \mid \text{good})}{P(w \mid \text{bad})}, 264\qquad 265P(w \mid c) = \frac{\text{count}_c(w) + 1}{N_c + V} 266$$ 267 268**Symbols** 269 270| Symbol | Meaning here | In the example | 271|---|---|---| 272| $w$ | one word of the document | "river" | 273| $\sum_{w \in \text{doc}}$ | add up the term for every word in the document | | 274| $c$ | a class: good (reference) or bad (spam) | | 275| $\text{count}_c(w)$ | how often $w$ appears in class $c$'s examples | 3 in good, 0 in bad | 276| $N_c$ | total words in class $c$'s examples | 77 good, 45 bad | 277| $V$ | number of distinct words across both classes (the vocabulary) | 80 | 278| $+1$ | add-one smoothing: pretend every word was seen once more, so an unseen word never gives probability 0 (and log 0 = −∞) | | 279| $P(w \mid c)$ | "probability of $w$ given $c$": how likely a word drawn from class $c$ is $w$ | $P(\text{river} \mid \text{good}) = 4/157$ | 280| $\log$ | natural logarithm: turns a ratio above 1 into a positive vote and below 1 into a negative vote | $\log 3.185 = 1.158$ | 281 282**In words:** "for each word, ask how much more likely it is in good text 283than in spam, take the log of that ratio as the word's vote, and add up the 284votes." 285 286**With the numbers:** P(river | good) = (3 + 1) / (77 + 80) = 4/157 and 287P(river | bad) = (0 + 1) / (45 + 80) = 1/125. Their ratio is 3.185 and its 288log is **+1.158**. For "click": (0 + 1)/157 over (4 + 1)/125 is 0.159, a vote 289of **−1.837**. The sentence "the delta is formed when the river deposits 290sand over many years" scores +8.52; "click here buy now best price free free 291free" scores −16.49. 292 293**In Python:** 294 295```python 296import math 297N_good, N_bad, V = 77, 45, 80 298def vote(count_good, count_bad): 299 # log P(w | good) / P(w | bad), with add-one smoothing 300 p_good = (count_good + 1) / (N_good + V) 301 p_bad = (count_bad + 1) / (N_bad + V) 302 return math.log(p_good / p_bad) 303# "river": 3 times in good, 0 in bad 304round(vote(3, 0), 3) # → 1.158 305# "click": 0 times in good, 4 in bad 306round(vote(0, 4), 3) # → -1.837 307``` 308 309Why "naive"? It treats every word as independent evidence, which is false 310(words come in phrases) but works well enough to sort billions of pages 311cheaply. The practical danger is that a classifier keeps whatever resembles 312its reference set, including its blind spots: pick only encyclopedia text 313as "good" and you quietly filter out dialects, forums and whole topics. 314 315**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. 316 317### 1d. Exact duplicates: one fingerprint per page 318 319**Everyday picture.** A cloakroom attendant does not compare every coat 320with every other coat. Each coat gets a numbered ticket, and two coats with 321the same ticket are the same coat. 322 323The web is full of copies: mirrored sites, syndicated news, the same terms 324of service on a million shops. For exact copies the trick is a **hash**: a 325function that turns any text into a short fingerprint, always the same for 326the same text and almost never the same for different texts. Lowercase the 327text and squash its spaces first, so that trivially reformatted copies get 328the same fingerprint, then keep the first page with each fingerprint. One 329pass, one lookup per page, no comparisons. 330 331**In code:** `normalize_for_exact` lowercases and collapses whitespace, and `exact_dedup` keeps the first document with each SHA-1 fingerprint. 332 333### 1e. Near duplicates: shingles and Jaccard similarity 334 335A scraper site that copies an article and adds "Read more on our site" 336defeats the exact hash: one changed character gives a completely different 337fingerprint. We need a measure of *how much* two pages overlap. 338 339**Everyday picture.** Cut each page into overlapping strips of a few words, 340like roof shingles, and put each page's strips in a bag. Two pages are near 341duplicates when their bags hold mostly the same strips. 342 343**Tiny worked example.** With 2-word shingles: 344 345| Sentence | Its shingles | 346|---|---| 347| "the cat sat on the mat" | the cat, cat sat, sat on, on the, the mat | 348| "the cat sat on a mat" | the cat, cat sat, sat on, on a, a mat | 349 350Three shingles are shared (the cat, cat sat, sat on) and seven are distinct 351across both, so the overlap is 3/7 = 0.43. 352 353$$ 354J(A, B) = \frac{|A \cap B|}{|A \cup B|} 355$$ 356 357**Symbols** 358 359| Symbol | Meaning here | In the example | 360|---|---|---| 361| $A$, $B$ | the two sets of shingles | 5 shingles each | 362| $A \cap B$ | the **intersection**: shingles in both sets | {the cat, cat sat, sat on} | 363| $A \cup B$ | the **union**: shingles in either set, each counted once | 7 shingles | 364| $\lvert \cdot \rvert$ | the number of items in a set | $\lvert A \cap B \rvert = 3$ | 365| $J(A, B)$ | the **Jaccard similarity**, from 0 (nothing shared) to 1 (identical) | 0.43 | 366 367**In words:** "the number of shingles the two pages share, divided by the 368number of different shingles they have between them." 369 370**With the numbers:** J = 3 / 7 = **0.43**. Real pipelines use 5-word 371shingles on whole pages; the scraper's copy of our 68-word river paragraph 372(one phrase changed, one sentence added) shares 78% of its shingles with 373the original. 374 375**In Python:** 376 377```python 378def shingles(text, k=2): 379 ws = text.split() 380 return {" ".join(ws[i:i + k]) for i in range(len(ws) - k + 1)} 381A = shingles("the cat sat on the mat") 382B = shingles("the cat sat on a mat") 383# |A ∩ B| and |A ∪ B| 384len(A & B), len(A | B) # → (3, 7) 385# J(A, B) 386round(len(A & B) / len(A | B), 2) # → 0.43 387``` 388 389**In code:** `shingles` cuts a text into k-word windows (5 by default) and `jaccard` divides the intersection by the union. 390 391### 1f. MinHash: estimating Jaccard from a few numbers 392 393Jaccard needs both full shingle sets side by side. With billions of pages, 394comparing every pair is impossible (a billion pages make about 5 × 10¹⁷ 395pairs). **MinHash** compresses each page into a short list of numbers, its 396**signature**, such that comparing two signatures estimates their Jaccard. 397 398**Everyday picture.** Shuffle a deck containing every shingle in the world 399and deal from the top. Stop at the first card that belongs to page A, and 400separately at the first card that belongs to page B. Those two "first cards" 401are the same card exactly when the first card from A's-or-B's pile happens 402to be one they share. The more they share, the likelier that is. 403 404**Tiny worked example.** A = {a, b, c} and B = {b, c, d}, so J = 2/4 = 0.5. 405Shuffle the four letters in all 24 possible orders. In 12 of them the first 406letter from A ∪ B is b or c (a shared letter), and then A's first and B's 407first are the same letter. 12 / 24 = 0.5, exactly J. 408 409$$ 410P\big[\min h(A) = \min h(B)\big] = J(A, B), 411\qquad 412\hat{J} = \frac{1}{k} \sum_{i=1}^{k} \mathbf{1}\big[\min h_i(A) = \min h_i(B)\big] 413$$ 414 415**Symbols** 416 417| Symbol | Meaning here | In the example | 418|---|---|---| 419| $h$ | 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 | 420| $\min h(A)$ | the smallest number $h$ gives any shingle of $A$: A's "first card" | | 421| $P[\ldots]$ | the probability that the statement in brackets is true | 12/24 | 422| $k$ | how many independent hash functions (the signature length) | 128 | 423| $h_i$ | the $i$-th hash function | | 424| $\mathbf{1}[\ldots]$ | the **indicator**: 1 if the statement is true, 0 if not | | 425| $\hat{J}$ | the estimate of $J$ (the hat means "estimated") | | 426 427**In words:** "under a random shuffle, two sets have the same first element 428with probability equal to their Jaccard similarity; so repeat with k 429shuffles and report the fraction of times the first elements agreed." 430 431**With the numbers:** all 24 orders of {a, b, c, d}: 12 agree, P = 0.5 = J. 432With k = 128 hashes, each pair of pages is compared with 128 numbers 433instead of hundreds of shingles, and the estimate's typical error is about 434√(J(1 − J)/k) = √(0.25/128) ≈ 0.044. 435 436**In Python:** 437 438```python 439import itertools, math 440A, B = {"a", "b", "c"}, {"b", "c", "d"} 441orders = list(itertools.permutations("abcd")) 442def first(order, s): 443 return next(x for x in order if x in s) 444agree = sum(first(o, A) == first(o, B) for o in orders) 445agree, len(orders) # → (12, 24) 446# P[min h(A) = min h(B)] equals J(A, B) 447agree / len(orders), len(A & B) / len(A | B) # → (0.5, 0.5) 448# typical error of the estimate with k = 128 hashes 449round(math.sqrt(0.5 * 0.5 / 128), 3) # → 0.044 450``` 451 452 453 454**Reading it:** the horizontal axis is the number of hash functions k (log 455scale) and the vertical axis the MinHash estimate. Each colour is one pair 456of sets and its dashed line is that pair's true Jaccard. On the left, with 457only a few hashes, the estimate can only be a coarse fraction and lurches 458around. Moving right, each line settles onto its dashed line. That is the 459whole promise of MinHash: a fixed, small signature per page, with error you 460choose by choosing k. 461 462**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. 463 464### 1g. Locality-sensitive hashing: finding candidate pairs without comparing all of them 465 466Signatures make each comparison cheap, but a billion pages still make too 467many pairs. **Locality-sensitive hashing (LSH)** avoids most comparisons: 468cut each signature into $b$ **bands** of $r$ numbers, and file every page 469into one bucket per band, keyed by that band's numbers. Only pages that 470share a bucket in at least one band are ever compared. 471 472```mermaid 473flowchart LR 474 D[Page] --> S[5-word shingles] 475 S --> H[k hash functions<br/>keep each minimum] 476 H --> G[Signature<br/>k numbers] 477 G --> B1[band 1: r numbers] --> K1[bucket] 478 G --> B2[band 2] --> K2[bucket] 479 G --> BB[band b] --> KB[bucket] 480 K1 & K2 & KB --> C[Candidate pairs:<br/>pages sharing any bucket] 481 C --> V[Check estimated Jaccard<br/>against the threshold] 482``` 483 484**Reading it:** a page flows left to right. Its shingles become a signature 485of k = b × r numbers, and the signature is sliced into b bands. Each band is 486a key into its own table of buckets. Two pages meet in the "candidate pairs" 487box only if all r numbers of some band match exactly; for everything else 488no comparison ever happens. The final box confirms each candidate with the 489full signature. 490 491$$ 492P(\text{candidate}) = 1 - \left(1 - s^{r}\right)^{b} 493$$ 494 495**Symbols** 496 497| Symbol | Meaning here | In the example | 498|---|---|---| 499| $s$ | the pair's Jaccard similarity | 0.8 or 0.3 | 500| $r$ | rows: numbers per band | 5 | 501| $b$ | number of bands | 20 | 502| $s^r$ | chance that all $r$ numbers of one band agree (each agrees with chance $s$) | $0.8^5 = 0.328$ | 503| $1 - s^r$ | chance that one band does *not* fully agree | 0.672 | 504| $(1 - s^r)^b$ | chance that *no* band agrees | $0.672^{20} = 0.00036$ | 505 506**In words:** "a pair becomes a candidate unless every one of its b bands 507has at least one disagreeing number." 508 509**With the numbers:** at s = 0.8: 1 − (1 − 0.328)²⁰ = **0.9996**, so near 510duplicates are almost never missed. At s = 0.3: 0.3⁵ = 0.00243, and 5111 − 0.99757²⁰ = **0.047**, so dissimilar pages are almost never compared. 512 513**In Python:** 514 515```python 516def p_candidate(s, b=20, r=5): 517 # 1 − (1 − s^r)^b 518 return 1 - (1 - s ** r) ** b 519round(p_candidate(0.8), 4) # → 0.9996 520round(p_candidate(0.3), 4) # → 0.0475 521# where the curve is steepest: about (1/b)^(1/r) 522round((1 / 20) ** (1 / 5), 2) # → 0.55 523``` 524 525 526 527**Reading it:** the horizontal axis is the true Jaccard of a pair and the 528vertical axis its chance of being compared. Every layout draws an S: pairs 529on the left are almost never compared (that is the saving) and pairs on the 530right almost always are (that is the recall). The dashed verticals mark 531(1/b)^(1/r), where each curve is steepest. Choosing b and r is choosing where 532the cliff sits: more rows per band push it right and make it sharper. 533 534**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. 535 536### 1h. Why duplicates hurt 537 538**Everyday picture.** A student who reads the same paragraph a hundred 539times can recite it, but has not learned a hundred paragraphs' worth. A 540model that sees a page thousands of times does the same: it spends capacity 541memorizing that page, and learns to recite it when prompted. 542 543**Tiny worked example.** A **bigram model** predicts each word from the 544word before it, by counting pairs. Train one on five short lines plus the 545boilerplate "click here to subscribe for free updates", and ask for the 546probability that it continues "click" into the whole boilerplate line: 547 548| Next-word step | Kept once | Kept 100 times | 549|---|---|---| 550| click → here | 1/2 | 100/101 | 551| here → to | 1/2 | 100/101 | 552| to → subscribe | 1/3 | 100/102 | 553| subscribe → for | 2/2 | 101/101 | 554| for → free | 1/3 | 100/102 | 555| free → updates | 1/2 | 100/101 | 556| **whole line** | **1/72 = 0.014** | **0.93** | 557 558Kept once, the model has many ways to continue "click". Repeated 100 times, 559the line becomes the only road, and the model regurgitates it 93% of the 560time. Lee et al. (2021) found a single 61-word sentence repeated more than 56160,000 times in the C4 dataset, and that models trained on deduplicated 562data emit memorized training text about ten times less often. Duplicates 563also waste compute and leak into test sets: a benchmark question copied 564across the web ends up in the training data, and the model's score on it 565measures recall, not skill. 566 567**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. 568 569### 1i. The whole belt on a small crawl 570 571`CRAWL_SAMPLE` holds seven pages. Running `curate` over it: 572 573| Page | What it is | Verdict | 574|---|---|---| 575| 0 | a clean English paragraph about river deltas | kept | 576| 1 | the same idea in French | language: fr | 577| 2 | a navigation bar | quality: too_short, too_few_alphabetic_words, missing_stop_words | 578| 3 | page 0 re-crawled with different capitals and spacing | exact duplicate of 0 | 579| 4 | a hashtag spam page | quality: too_many_symbols, missing_stop_words | 580| 5 | page 0 with a phrase changed and a link added | near duplicate of 0 | 581| 6 | a clean English paragraph about stars | kept | 582 583Two of seven survive. That is not unusual: large public pipelines keep only 584a small fraction of the raw crawl they start from. 585 586**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. 587 588### 1j. Data mixtures: how much of each source 589 590**Everyday picture.** A diet is not "all the food in the shop". You choose 591proportions: mostly staples, some vegetables, a little of the rich stuff. 592Pretraining data is mixed the same way from sources that differ in size and 593value: web text, code, books, encyclopedias, scientific papers, maths. 594 595Each source gets a **mixture weight**: its share of the tokens the model 596will see. Weights do not follow size. A small, valuable source (an 597encyclopedia) is often up-weighted, which means the model sees it more than 598once; a huge, noisy one (the web) is sampled less than one full pass. 599 600$$ 601\text{epochs}_i = \frac{w_i \, D}{N_i} 602$$ 603 604**Symbols** 605 606| Symbol | Meaning here | In the example | 607|---|---|---| 608| $i$ | one data source | the encyclopedia | 609| $w_i$ | its mixture weight: share of all training tokens drawn from it | 0.05 | 610| $D$ | the total training budget, in tokens | 1,000 billion | 611| $N_i$ | the tokens that source actually has | 20 billion | 612| $\text{epochs}_i$ | how many full passes over the source the mixture implies | 2.5 | 613 614**In words:** "the tokens you plan to draw from a source, divided by the 615tokens it has, is how many times the model reads it." 616 617**With the numbers:** web 0.80 × 1,000B / 900B = **0.89** passes; wiki 6180.05 × 1,000B / 20B = **2.5** passes; code 0.15 × 1,000B / 150B = **1.0**. 619The LLaMA paper's mixture has the same shape: Common Crawl is two thirds of 620the tokens at about 1.1 passes, while Wikipedia and books are 4.5% each but 621are read more than twice (2.45 and 2.23 passes). 622 623**In Python:** 624 625```python 626weights = {"web": 0.80, "wiki": 0.05, "code": 0.15} 627available = {"web": 900e9, "wiki": 20e9, "code": 150e9} 628D = 1000e9 629# epochs_i = w_i D / N_i 630{k: round(weights[k] * D / available[k], 2) for k in weights} # → {'web': 0.89, 'wiki': 2.5, 'code': 1.0} 631``` 632 633Repeating a small source a few times is fine; many more passes and the model 634starts memorizing it (the duplicate problem again, on purpose). Mixture 635weights are usually chosen by training small models on candidate mixtures 636and comparing them, and many recipes change the mixture near the end of 637training, up-weighting the cleanest sources. 638 639**In code:** `epochs_per_source` applies the formula to every source. 640 641### 1k. Token budgets: how much data for how big a model 642 643**Everyday picture.** A bigger brain can learn more, but only if you give 644it more to read. Given a fixed amount of study time (compute), there is a 645best split between "bigger brain" and "more reading". 646 647All budgets are counted in **tokens**, the pieces a tokenizer cuts text 648into (see `primer.ml.tokenization`); in English a token is roughly three 649quarters of a word. Two rules of thumb set the scale. 650 651$$ 652D_{\text{opt}} \approx 20\,N, 653\qquad 654C \approx 6\,N\,D 655$$ 656 657**Symbols** 658 659| Symbol | Meaning here | In the example | 660|---|---|---| 661| $N$ | the number of model parameters | 7 billion | 662| $D$ | the number of training tokens | 140 billion | 663| $D_{\text{opt}}$ | the compute-optimal token count for size $N$ (Hoffmann et al., 2022) | $20 \times 7 \times 10^9$ | 664| $C$ | training compute in FLOPs (floating-point operations) | $5.88 \times 10^{21}$ | 665| $6$ | 2 FLOPs per parameter per token in the forward pass (one multiply, one add), 4 in the backward pass | | 666| $\approx$ | "approximately equal": a rule of thumb, not an exact law | | 667 668**In words:** "for the best model per unit of compute, train on about 20 669tokens per parameter; and training costs about six operations per 670parameter per token." 671 672**With the numbers:** a 7B model's compute-optimal budget is 20 × 7 × 10⁹ = 673**140 billion tokens**, and training it costs 6 × 7 × 10⁹ × 1.4 × 10¹¹ = 674**5.88 × 10²¹ FLOPs**. At a sustained 400 teraFLOP/s per GPU (about 40% of an 675H100's bf16 peak) that is 5.88 × 10²¹ / 4 × 10¹⁴ ≈ 1.5 × 10⁷ GPU-seconds, 676about 4,100 GPU-hours. 677 678**In Python:** 679 680```python 681N = 7e9 682# D_opt ≈ 20 N 683D = 20 * N 684D # → 140000000000.0 685# C ≈ 6 N D 686C = 6 * N * D 687C # → 5.88e+21 688# GPU-hours at 400 teraFLOP/s each 689round(C / 400e12 / 3600) # → 4083 690``` 691 692In practice, models that will serve billions of requests are trained far 693past this point, because a smaller model trained longer is cheaper to run 694forever after. Llama 2 7B saw 2 trillion tokens (about 286 per parameter); 695Llama 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 696often the binding limit, and why the next topic exists. 697 698**In code:** `chinchilla_tokens` and `training_flops` are the two rules of thumb. 699 700### 1l. Synthetic data, and model collapse 701 702When good human text runs short, models generate more: rewritten web pages, 703worked maths solutions checked by a program, question-and-answer pairs 704drawn out of textbooks. Used carefully this works well: it can turn a 705messy page into a clear one, and a checked answer is clean signal. Used 706carelessly it has a known failure. 707 708**Everyday picture.** Photocopy a photo, then photocopy the copy, and keep 709going. Each copy is nearly right, but fine detail disappears first, and 710after enough rounds you have a grey smudge. A model trained on its own 711outputs loses the rare, surprising parts of the data (the tails) first. 712 713**Tiny worked example.** Draw 20 numbers from a bell curve with spread 1. 714Fit a bell curve to those 20 numbers, draw 20 new numbers from the fit, fit 715again, and repeat. Each fit is a slightly narrower guess on average, and the 716narrowing compounds. 717 718```mermaid 719flowchart LR 720 R[Real data<br/>spread 1.0] --> F1[Fit model 1] 721 F1 --> S1[Sample from model 1] 722 S1 --> F2[Fit model 2] 723 F2 --> S2[Sample from model 2] 724 S2 --> FN[... model n] 725 FN --> X[Spread shrinks:<br/>tails are forgotten first] 726``` 727 728**Reading it:** follow the chain left to right. Only the first model ever 729sees real data; every later model learns from the previous model's samples. 730Nothing in the loop can put back a rare value once a sample happens to miss 731it, so information only leaks out. The fix in practice is to keep real data 732in every generation's mix, and to filter synthetic data with checks that do 733not come from the same model. 734 735$$ 736\mathbb{E}\big[\hat{\sigma}^2_{t+1}\big] = \frac{n - 1}{n}\,\hat{\sigma}^2_t 737$$ 738 739**Symbols** 740 741| Symbol | Meaning here | In the example | 742|---|---|---| 743| $t$ | the generation number | 0, 1, 2, … | 744| $\hat{\sigma}^2_t$ | the fitted **variance** (spread squared) of generation $t$'s model | 1.0 at the start | 745| $n$ | samples each generation is fitted on | 20 | 746| $\mathbb{E}[\ldots]$ | the **expected value**: the average over many repeats of the experiment | | 747| $\frac{n-1}{n}$ | the shrink factor per generation of a maximum-likelihood fit | 0.95 | 748 749**In words:** "on average, each refit on n samples keeps only (n − 1)/n of 750the previous generation's variance." 751 752**With the numbers:** 0.95 per generation; after 200 generations the 753expected variance is 0.95²⁰⁰ ≈ 0.000035 of the original, a spread of about 7540.006. The single run below does not follow the average exactly, but it 755collapses all the same. 756 757**In Python:** 758 759```python 760import math 761n, generations = 20, 200 762# (n − 1)/n per generation, compounded 763shrink = ((n - 1) / n) ** generations 764f"{shrink:.1e}" # → '3.5e-05' 765# the spread is the square root of the variance 766round(math.sqrt(shrink), 4) # → 0.0059 767``` 768 769 770 771**Reading it:** the horizontal axis is the generation and the vertical axis 772the fitted spread, on a log scale so that halving always looks the same 773size. Each coloured line is one run with a different seed and the dashed 774line is the formula's average. The runs wander, up some generations and 775down others, but the trend is relentlessly downward, and a typical run 776falls even faster than the dashed line: the average is held up by rare runs 777that happen to stay wide. After 200 generations the "model" can only 778produce values in a sliver of the original range. 779Shumailov et al. (2023) showed the same effect in language models trained 780recursively on their own text. 781 782**In code:** `recursive_gaussian_fit` runs the fit-sample-refit loop and returns every generation's fitted spread. 783 784## 2. Why one GPU is not enough: the memory bill 785 786**Everyday picture.** To bake one cake you need the recipe, but to *learn* 787to bake you also need your notes on every attempt: what went wrong, how 788much to change, how confident you are in each change. Training a model is 789the same. The weights are the recipe; training also keeps a gradient for 790every weight (what to change) and the optimizer's running notes (how it has 791been changing), and those notes are bigger than the recipe. 792 793**Tiny worked example.** Take one parameter trained with Adam (see 794`primer.ml.optimizers`) in mixed precision, the standard recipe (section 4 795explains the number formats): 796 797| What is stored per parameter | Format | Bytes | 798|---|---|---| 799| the weight used in the forward and backward pass | bf16 | 2 | 800| its gradient | bf16 | 2 | 801| a full-precision **master copy** of the weight | fp32 | 4 | 802| Adam's momentum (running average of gradients) | fp32 | 4 | 803| Adam's variance (running average of squared gradients) | fp32 | 4 | 804| **total** | | **16** | 805 806$$ 807M_{\text{state}} = (2 + 2 + 4 + 4 + 4)\,\Psi = 16\,\Psi \ \text{bytes} 808$$ 809 810**Symbols** 811 812| Symbol | Meaning here | In the example | 813|---|---|---| 814| $\Psi$ | the number of parameters (Psi, the letter the ZeRO paper uses) | $7 \times 10^9$ | 815| $2 + 2$ | bytes for the bf16 weight and its bf16 gradient | | 816| $4 + 4 + 4$ | bytes for the fp32 master weight and Adam's two fp32 averages | | 817| $M_{\text{state}}$ | memory for the training state, before any activations | 112 GB | 818 819**In words:** "training with Adam in mixed precision costs sixteen bytes 820for every parameter, before you store a single activation." 821 822**With the numbers:** 16 × 7 × 10⁹ = **112 GB** for a 7B model. A widely used 823data-centre GPU, the H100, holds 80 GB (see `primer.ml.hardware`), so the 824state alone does not fit. Serving the same model needs only the 2-byte weights, 82514 GB: training costs eight times the memory of inference. 826 827**In Python:** 828 829```python 830params = 7e9 831bytes_per_param = 2 + 2 + 4 + 4 + 4 832bytes_per_param # → 16 833# M_state in GB 834bytes_per_param * params / 1e9 # → 112.0 835# inference needs only the bf16 weights 8362 * params / 1e9 # → 14.0 837``` 838 839### Activations: the memory that grows with the batch 840 841The backward pass needs the intermediate results of the forward pass (the 842**activations**) to compute gradients, so they are kept until it runs. Their 843size grows with the number of tokens in flight, not with the parameter count. 844Korthikanti et al. (2022) counted them for one transformer layer, with 16-bit 845activations: 846 847$$ 848A = L \cdot s\,b\,h \left(34 + \frac{5\,a\,s}{h}\right) \ \text{bytes} 849$$ 850 851**Symbols** 852 853| Symbol | Meaning here | For a 7B shape | 854|---|---|---| 855| $L$ | number of layers | 32 | 856| $s$ | sequence length in tokens | 4,096 | 857| $b$ | sequences per GPU in the batch | 1 | 858| $h$ | hidden width | 4,096 | 859| $a$ | attention heads | 32 | 860| $34\,s\,b\,h$ | the layer's ordinary tensors (inputs to each matrix multiply, norms, dropout masks) | | 861| $5\,a\,s^2 b$ | attention's $s \times s$ score and weight matrices, one per head (written as $s b h \cdot 5as/h$) | | 862 863**In words:** "each layer keeps about 34 bytes per token per hidden unit, 864plus attention's square score matrices, which grow with the square of the 865sequence length." 866 867**With the numbers:** per layer, 4096 × 4096 × (34 + 5 × 32 × 4096 / 4096) 868= 16.8 million × 194 bytes = 3.26 GB; over 32 layers, **104 GB**. Kernels 869like FlashAttention never store the score matrices (they recompute them in 870the backward pass), which removes the 5as/h term: 34 × 16.8 million × 32 = 871**18.3 GB**. 872 873**In Python:** 874 875```python 876L, s, b, h, a = 32, 4096, 1, 4096, 32 877# A = L · s b h (34 + 5 a s / h) 878round(L * s * b * h * (34 + 5 * a * s / h) / 1e9, 1) # → 104.2 879# without stored attention scores 880round(L * s * b * h * 34 / 1e9, 1) # → 18.3 881# activation checkpointing keeps only each layer's 2-byte input 882round(L * 2 * s * b * h / 1e9, 2) # → 1.07 883``` 884 885The last line is **activation checkpointing** (also called gradient 886checkpointing, Chen et al., 2016): keep only each layer's input, and during 887the backward pass re-run that layer's forward pass to rebuild what it needs. 888It trades about one extra forward pass (roughly a third more compute) for 889activation memory that no longer grows with depth. 890 891 892 893**Reading it:** each bar is the memory one GPU would need to train the 7B 894model alone, split by what it holds; the dashed line is an 80 GB GPU. The 895three bars differ only in how activations are handled: stored naively, 896without attention scores, or checkpointed. Even the leanest bar is above 897the line, because the 112 GB of weights, gradients and optimizer state is 898there in every bar. Shrinking activations is not enough: the state itself 899has to be split across GPUs. That is the next section. 900 901**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. 902 903## 3. Parallelism: many GPUs, one model 904 905**Everyday picture.** A restaurant can serve more diners in three ways. 906Open identical kitchens that each cook whole meals (**data parallelism**). 907Split one dish across cooks working side by side on the same step, one 908chopping the left half of the onions and one the right (**tensor 909parallelism**). Or build an assembly line, with each cook doing one stage 910and passing the plate on (**pipeline parallelism**). Real training runs use 911all three at once. 912 913### 3a. Data parallelism: identical copies, averaged gradients 914 915Every GPU holds a full copy of the model and processes a different slice of 916the batch. After the backward pass, the GPUs average their gradients so all 917copies take the same step and stay identical. This works because the 918gradient of an average loss is the average of the gradients. 919 920**Tiny worked example.** A one-weight model predicts y = w·x, with loss the 921mean of (w·x − y)². Batch: x = 1, 2, 3, 4 with y = 2, 4, 6, 8, and w = 1. On 922one GPU the gradient is (2/4)·Σ x(wx − y) = (2/4)·(−1 − 4 − 9 − 16) = **−15**. 923Split across two GPUs: GPU 1 gets x = 1, 2 and computes −5; GPU 2 gets 924x = 3, 4 and computes −25. Their average is (−5 − 25)/2 = **−15**, the same. 925 926$$ 927\nabla \mathcal{L} = \frac{1}{G} \sum_{k=1}^{G} \nabla \mathcal{L}_k 928$$ 929 930**Symbols** 931 932| Symbol | Meaning here | In the example | 933|---|---|---| 934| $\mathcal{L}$ | the loss averaged over the whole batch | | 935| $\nabla$ | "the gradient of": the slope of the loss for every weight | | 936| $G$ | number of GPUs, each with an equal share of the batch | 2 | 937| $\mathcal{L}_k$ | the loss averaged over GPU $k$'s share | | 938| $\sum_{k=1}^{G}$ | add up over every GPU | | 939 940**In words:** "the gradient for the whole batch is the average of the 941gradients each GPU computed on its own equal share." 942 943**With the numbers:** (−5 + −25) / 2 = −15, matching the single-GPU −15. 944 945**In Python:** 946 947```python 948def grad(xs, ys, w=1.0): 949 # d/dw of mean((w x − y)²) 950 return 2 / len(xs) * sum(x * (w * x - y) for x, y in zip(xs, ys)) 951grad([1, 2, 3, 4], [2, 4, 6, 8]) # → -15.0 952g1, g2 = grad([1, 2], [2, 4]), grad([3, 4], [6, 8]) 953g1, g2 # → (-5.0, -25.0) 954(g1 + g2) / 2 # → -15.0 955``` 956 957The averaging step is an **all-reduce**: every GPU contributes a vector and 958every GPU receives the sum. Sending every gradient to one GPU would jam its 959network link, so the standard algorithm is the **ring all-reduce**. 960 961```mermaid 962flowchart LR 963 G0[GPU 0<br/>chunks A B C D] -->|one chunk per step| G1[GPU 1<br/>chunks A B C D] 964 G1 -->|one chunk per step| G2[GPU 2<br/>chunks A B C D] 965 G2 -->|one chunk per step| G3[GPU 3<br/>chunks A B C D] 966 G3 -->|one chunk per step| G0 967``` 968 969**Reading it:** the four GPUs sit in a ring and each only ever talks to its 970right-hand neighbour. Each cuts its gradient into four chunks. Phase one 971(**reduce-scatter**): for three steps, each GPU passes one chunk to the right, 972where it is added to the neighbour's copy; afterwards each GPU owns the 973complete sum for one chunk. Phase two (**all-gather**): for three more steps 974the finished chunks travel round the ring, so everyone ends with all four 975sums. Every link is busy at every step, and no GPU is a bottleneck. 976 977$$ 978\text{sent per GPU} = \frac{2\,(G - 1)}{G}\,S 979$$ 980 981**Symbols** 982 983| Symbol | Meaning here | In the example | 984|---|---|---| 985| $G$ | GPUs in the ring | 4, or 1,000 | 986| $S$ | size of the gradient being summed | 14 GB (7B in bf16) | 987| $G - 1$ | steps in each of the two phases | 3 | 988| $\frac{1}{G}$ | each step moves one chunk, $1/G$ of the gradient | | 989| $2$ | two phases: reduce-scatter, then all-gather | | 990 991**In words:** "each GPU sends a little under twice its gradient, however 992many GPUs are in the ring." 993 994**With the numbers:** with 4 GPUs, 2 × 3/4 = 1.5 gradients' worth; for a 99514 GB gradient on 8 GPUs, 2 × 7/8 × 14 = **24.5 GB** per GPU per step; on 9961,000 GPUs, **27.97 GB**. The traffic per GPU barely grows, which is why 997data parallelism scales to thousands of GPUs. 998 999**In Python:** 1000 1001```python 1002def sent(G, S): 1003 # 2 (G − 1) / G · S 1004 return 2 * (G - 1) / G * S 1005sent(4, 1.0) # → 1.5 1006sent(8, 14e9) / 1e9 # → 24.5 1007round(sent(1000, 14e9) / 1e9, 2) # → 27.97 1008``` 1009 1010Data parallelism alone does not solve the memory problem: every GPU still 1011holds all 16 bytes per parameter. 1012 1013**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. 1014 1015### 3b. Sharding the state: ZeRO and FSDP 1016 1017**Everyday picture.** A reading group with one expensive textbook does not 1018buy a copy each. They tear it into chapters, each keeps one, and whoever 1019needs a chapter borrows it for the evening and hands it back. 1020 1021In plain data parallelism, every GPU stores an identical copy of the 1022optimizer state, the gradients and the weights: pure waste. **ZeRO** (the 1023Zero Redundancy Optimizer) keeps data parallelism's split of the batch but 1024gives each GPU only a 1/G shard of that state, in three stages. PyTorch's 1025**FSDP** (fully sharded data parallel) implements the third. 1026 1027```mermaid 1028flowchart TB 1029 subgraph S0["Stage 0: plain data parallel"] 1030 A0["every GPU: weights + gradients + optimizer state"] 1031 end 1032 subgraph S1["Stage 1: shard the optimizer state"] 1033 A1["every GPU: weights + gradients<br/>its 1/G of the optimizer state"] 1034 end 1035 subgraph S2["Stage 2: also shard the gradients"] 1036 A2["every GPU: weights<br/>its 1/G of gradients and optimizer state"] 1037 end 1038 subgraph S3["Stage 3 = FSDP: shard everything"] 1039 A3["every GPU: its 1/G of everything<br/>borrow each layer's weights just in time"] 1040 end 1041 S0 --> S1 --> S2 --> S3 1042``` 1043 1044**Reading it:** read top to bottom; each stage moves one more kind of state 1045from "every GPU keeps all of it" to "every GPU keeps a 1/G slice". Stage 1 1046works because each GPU only needs to update its own slice of the weights. 1047Stage 2 works because a reduce-scatter (the first half of the ring) can 1048deliver each GPU just the gradient slice it updates. Stage 3 goes furthest: 1049before computing a layer, the GPUs all-gather that layer's weights, use 1050them, and throw them away again. 1051 1052$$ 1053M_{\text{stage 0}} = 16\Psi, 1054\quad 1055M_{1} = 4\Psi + \frac{12\Psi}{G}, 1056\quad 1057M_{2} = 2\Psi + \frac{14\Psi}{G}, 1058\quad 1059M_{3} = \frac{16\Psi}{G} 1060$$ 1061 1062**Symbols** 1063 1064| Symbol | Meaning here | In the ZeRO paper's example | 1065|---|---|---| 1066| $\Psi$ | parameters | $7.5 \times 10^9$ | 1067| $G$ | GPUs sharing the state | 64 | 1068| $12\Psi$ | the fp32 master weights and Adam's two averages | 90 GB | 1069| $2\Psi$, $4\Psi$ | the bf16 weights; weights plus gradients | 15 GB; 30 GB | 1070| $M_{\text{stage}}$ | training state per GPU at that stage | | 1071 1072**In words:** "whatever is sharded is divided by the number of GPUs; what 1073is not sharded stays whole on every GPU." 1074 1075**With the numbers:** 7.5B parameters on 64 GPUs (the example in Figure 1 1076of the ZeRO paper): stage 0, **120 GB**; stage 1, 30 + 1.4 = **31.4 GB**; 1077stage 2, 15 + 1.6 = **16.6 GB**; stage 3, **1.9 GB**. The 7B model that could 1078not fit on one GPU now needs under 2 GB of state per GPU. 1079 1080**In Python:** 1081 1082```python 1083psi, G = 7.5e9, 64 1084stages = [16 * psi, 4 * psi + 12 * psi / G, 2 * psi + 14 * psi / G, 16 * psi / G] 1085[round(m / 1e9, 1) for m in stages] # → [120.0, 31.4, 16.6, 1.9] 1086``` 1087 1088The price is communication. Stages 1 and 2 cost the same traffic as a plain 1089all-reduce; stage 3 adds an all-gather of the weights in the forward pass 1090and again in the backward pass, about 1.5 times the traffic. Frameworks hide 1091it by fetching the next layer's weights while the current layer computes. 1092 1093 1094 1095**Reading it:** the horizontal axis is the number of GPUs (doubling at each 1096tick) and the vertical axis the state each GPU holds, both on log scales. The 1097dashed line is an 80 GB GPU. Stage 0 is flat: adding GPUs never helps. Stages 10981 and 2 fall at first and then level off at the part they do not shard (28 1099GB and 14 GB of weights and gradients). Only stage 3 keeps falling in a 1100straight line, because nothing is left unsharded. 1101 1102**In code:** `zero_memory_per_gpu` gives the state per GPU for any stage and GPU count. 1103 1104### 3c. Tensor parallelism: splitting one matrix multiply 1105 1106**Everyday picture.** Two people fill in one large multiplication table: 1107one does the left half of the columns, the other the right half. Neither 1108needs the other's work until the end, when the halves are placed side by 1109side. 1110 1111The biggest operations in a transformer are matrix multiplies, X·W. They 1112can be split across GPUs in two ways, and both give *exactly* the unsplit 1113answer. 1114 1115**Tiny worked example.** X = [1, 2] and 1116 1117W = [[1, 2, 3, 4], [5, 6, 7, 8]], so X·W = [11, 14, 17, 20]. 1118 1119- **By columns:** GPU 1 holds W's first two columns and computes [11, 14]; 1120 GPU 2 holds the last two and computes [17, 20]. Place side by side: 1121 [11, 14, 17, 20]. 1122- **By rows:** GPU 1 holds W's first row and X's first entry: 1 × [1, 2, 3, 4] 1123 = [1, 2, 3, 4]. GPU 2 holds the second: 2 × [5, 6, 7, 8] = [10, 12, 14, 16]. 1124 Add: [11, 14, 17, 20]. 1125 1126$$ 1127X W = \big[\,X W_1 \;\; X W_2\,\big] 1128\qquad\text{and}\qquad 1129X W = X_1 W_1 + X_2 W_2 1130$$ 1131 1132**Symbols** 1133 1134| Symbol | Meaning here | In the example | 1135|---|---|---| 1136| $X$ | the input activations, one row per token | [1, 2] | 1137| $W$ | a weight matrix | 2 × 4 | 1138| $[\,A \;\; B\,]$ | place two matrices side by side (**concatenate** columns) | | 1139| left: $W_1, W_2$ | $W$'s columns, split into two blocks | 2 × 2 each | 1140| right: $W_1, W_2$ | $W$'s rows, split into two blocks | 1 × 4 each | 1141| right: $X_1, X_2$ | $X$'s matching columns | [1] and [2] | 1142 1143**In words:** "split the weight by columns and each GPU produces some of 1144the output columns; split it by rows and each GPU produces a partial sum of 1145every output, and the partial sums add up to the answer." 1146 1147**With the numbers:** [11, 14] next to [17, 20], or [1, 2, 3, 4] + [10, 12, 114814, 16]: both give **[11, 14, 17, 20]**. 1149 1150**In Python:** 1151 1152```python 1153X = [1, 2] 1154W = [[1, 2, 3, 4], [5, 6, 7, 8]] 1155def matmul(x, w): 1156 return [sum(x[i] * w[i][j] for i in range(len(x))) for j in range(len(w[0]))] 1157matmul(X, W) # → [11, 14, 17, 20] 1158# by columns: each GPU computes half the output columns 1159matmul(X, [row[:2] for row in W]) + matmul(X, [row[2:] for row in W]) # → [11, 14, 17, 20] 1160# by rows: each GPU computes a partial sum of every output 1161p1, p2 = matmul(X[:1], W[:1]), matmul(X[1:], W[1:]) 1162[a + b for a, b in zip(p1, p2)] # → [11, 14, 17, 20] 1163``` 1164 1165Megatron-LM combines the two for a transformer's feed-forward layer, which 1166is Y = activation(X·W₁)·W₂: split W₁ by columns and W₂ by rows. 1167 1168```mermaid 1169flowchart LR 1170 X[X, full copy<br/>on both GPUs] --> A1["GPU 1: X · W1 left columns"] 1171 X --> A2["GPU 2: X · W1 right columns"] 1172 A1 --> R1[ReLU, locally] --> B1["· W2 top rows<br/>partial sum"] 1173 A2 --> R2[ReLU, locally] --> B2["· W2 bottom rows<br/>partial sum"] 1174 B1 & B2 --> AR[All-reduce:<br/>add the partial sums] --> Y[Y, full copy<br/>on both GPUs] 1175``` 1176 1177**Reading it:** the input is copied to both GPUs. The column split of W₁ 1178gives each GPU whole hidden units, so the activation function (applied 1179number by number) runs locally with no communication. The row split of W₂ 1180then consumes exactly those hidden units and produces a partial sum of the 1181output. One all-reduce at the end adds the two partial sums. Needing only 1182one all-reduce per block (the attention block is split the same way) is 1183what makes this practical, but it still happens inside every layer, so 1184tensor parallelism needs the fastest links available and usually stays 1185within one server of 8 GPUs. 1186 1187**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. 1188 1189### 3d. Pipeline parallelism, and the bubble 1190 1191**Everyday picture.** A car assembly line with four stations. The first 1192car takes four steps to roll off the end, and while it travels, stations 1193further down stand idle. Only when many cars are on the line at once is 1194every station busy. 1195 1196Pipeline parallelism gives each GPU a consecutive block of layers (a 1197**stage**). To keep the stages busy, the batch is cut into 1198**micro-batches** that follow each other down the line. The idle time at 1199the start and end is the **pipeline bubble**. 1200 1201```mermaid 1202flowchart LR 1203 MB[Micro-batches<br/>1, 2, 3, ...] --> S1[GPU 1<br/>layers 1-8] 1204 S1 -->|activations| S2[GPU 2<br/>layers 9-16] 1205 S2 -->|activations| S3[GPU 3<br/>layers 17-24] 1206 S3 -->|activations| S4[GPU 4<br/>layers 25-32] 1207 S4 -.->|gradients flow back| S1 1208``` 1209 1210**Reading it:** each GPU owns a quarter of the layers. Micro-batches enter 1211on the left one after another, and each GPU passes its output activations 1212to the next. When the forward passes are done, gradients flow back along 1213the dashed arrow in the reverse order. The only traffic is activations at 1214stage boundaries, far less than tensor parallelism's per-layer exchange, 1215which is why pipeline stages can sit on different servers. 1216 1217$$ 1218\text{bubble} = \frac{p - 1}{m + p - 1} 1219$$ 1220 1221**Symbols** 1222 1223| Symbol | Meaning here | In the example | 1224|---|---|---| 1225| $p$ | pipeline stages (GPUs in the line) | 4 | 1226| $m$ | micro-batches per batch | 1, 8 or 32 | 1227| $m + p - 1$ | time steps to push $m$ micro-batches through $p$ stages: $p$ to fill the line, then one per extra micro-batch | 11 when $m = 8$ | 1228| $p - 1$ | steps each stage spends idle while the line fills (and again while it drains) | 3 | 1229 1230**In words:** "the fraction of time each GPU sits idle is the fill time 1231divided by the total time; more micro-batches spread the same fill time 1232over more work." 1233 1234**With the numbers:** with 4 stages and 1 micro-batch, 3/4 = **75%** of the 1235time is bubble; with 8 micro-batches, 3/11 = **27%**; with 32, 3/35 = **8.6%**. 1236 1237**In Python:** 1238 1239```python 1240def bubble(p, m): 1241 # (p − 1) / (m + p − 1) 1242 return (p - 1) / (m + p - 1) 1243[round(bubble(4, m), 3) for m in (1, 8, 32)] # → [0.75, 0.273, 0.086] 1244``` 1245 1246 1247 1248**Reading it:** each row is a GPU (stage) and each column one time step. 1249Blue cells are forward passes and green cells backward passes, numbered by 1250micro-batch. The forward staircase runs down and to the right as each 1251micro-batch moves to the next stage; the backward staircase runs back up. 1252The white triangles in the corners are the bubble: stage 4 waits three steps 1253for the first micro-batch to arrive, and stage 1 waits three steps at the 1254end. Count them: 24 of the 88 cells are empty, 3/11 of the grid. 1255 1256 1257 1258**Reading it:** the horizontal axis is the number of micro-batches (log 1259scale) and the vertical axis the idle fraction. Each line is a pipeline 1260depth. To keep the bubble under about 10%, you need roughly ten times as 1261many micro-batches as stages, which pushes up the batch size. Smarter 1262schedules exist for exactly this reason: "one forward, one backward" (1F1B) 1263starts backward passes early to free activation memory, and interleaved 1264stages (Narayanan et al., 2021) give each GPU several smaller stages to 1265shrink the bubble further. 1266 1267**In code:** `pipeline_schedule` builds the grid (positive numbers forward, negative backward, 0 idle) and `bubble_fraction` is the formula. 1268 1269### 3e. All three at once 1270 1271```mermaid 1272flowchart TB 1273 subgraph DP["Data parallel: replicas see different data, all-reduce gradients"] 1274 subgraph R1["Replica 1"] 1275 direction LR 1276 P1["Pipeline stage 1<br/>one server: 8 GPUs, tensor parallel"] --> P2["Pipeline stage 2<br/>one server: 8 GPUs, tensor parallel"] 1277 end 1278 subgraph R2["Replica 2"] 1279 direction LR 1280 Q1["Pipeline stage 1<br/>8 GPUs, tensor parallel"] --> Q2["Pipeline stage 2<br/>8 GPUs, tensor parallel"] 1281 end 1282 end 1283``` 1284 1285**Reading it:** the three kinds nest, and each is placed where its traffic 1286fits. Tensor parallelism, which talks inside every layer, stays inside one 1287server on its fastest links. Pipeline parallelism, which only passes 1288activations between stages, spans servers. Data parallelism (often sharded 1289with ZeRO or FSDP), which talks once per step, wraps the whole thing and 1290multiplies it across the cluster. Llama 3 405B was trained this way on up 1291to 16,384 GPUs, adding a fourth kind (context parallelism) that splits very 1292long sequences. 1293 1294## 4. Mixed precision: doing the maths in fewer bits 1295 1296**Everyday picture.** A carpenter measures a room with a tape measure and a 1297table leg with calipers. Using calipers for everything would be slow; using 1298the tape measure for everything would give wobbly tables. Mixed precision 1299does each job with the coarsest number format that is good enough: the 1300heavy matrix multiplies in 16 (or 8) bits, and the few places where tiny 1301differences matter in 32. 1302 1303The payoff is large. Halving the bits halves the memory and the data moved, 1304and GPU tensor cores run 16-bit maths many times faster than 32-bit (see 1305`primer.ml.hardware`). 1306 1307### 4a. What a floating-point number is 1308 1309**Everyday picture.** Scientific notation, in binary. "6.02 × 10²³" has a 1310sign, a few significant digits and an exponent. A float is the same three 1311parts in bits: the exponent sets the **range** (how large or small a number 1312can be), and the fraction bits set the **precision** (how many significant 1313digits it keeps). 1314 1315**Tiny worked example.** In bf16, 1/3 is stored as sign 0, exponent field 1316125 and fraction field 0101011 in binary (43). The value is 13172^(125 − 127) × (1 + 43/128) = 0.25 × 1.3359375 = **0.333984375**. Seven 1318fraction bits keep only about three significant decimal digits, so bf16 1319cannot tell 1/3 from 0.33398. 1320 1321$$ 1322x = (-1)^{\text{sign}} \times 2^{\,E - \text{bias}} \times \left(1 + \frac{F}{2^{m}}\right) 1323$$ 1324 1325**Symbols** 1326 1327| Symbol | Meaning here | In the example | 1328|---|---|---| 1329| sign | 1 bit: 0 for positive, 1 for negative | 0 | 1330| $(-1)^{\text{sign}}$ | +1 or −1 | +1 | 1331| $E$ | the exponent field, read as a whole number | 125 | 1332| bias | a fixed offset, $2^{e-1} - 1$ for $e$ exponent bits, so negative exponents can be stored | 127 (bf16 and fp32) | 1333| $F$ | the fraction field, read as a whole number | 43 | 1334| $m$ | the number of fraction bits | 7 | 1335| $1 + F/2^m$ | the significant digits, between 1 and 2 (the leading 1 is implied, not stored) | 1.3359375 | 1336 1337**In words:** "a float is plus or minus a number between 1 and 2, scaled by 1338a power of two; exponent bits choose the power, fraction bits choose the 1339number between 1 and 2." 1340 1341**With the numbers:** 2^(−2) × (1 + 43/128) = **0.333984375**. The gap to the 1342next bf16 number above 1 is 2⁻⁷ = 0.0078; in fp32, with 23 fraction bits, 1343it is 2⁻²³ ≈ 0.00000012. 1344 1345**In Python:** 1346 1347```python 1348sign, E, bias, F, m = 0, 125, 127, 43, 7 1349# (−1)^sign × 2^(E − bias) × (1 + F / 2^m) 1350(-1) ** sign * 2.0 ** (E - bias) * (1 + F / 2 ** m) # → 0.333984375 1351# the gap after 1 (the "epsilon") in bf16 and fp32 13522.0 ** -7, 2.0 ** -23 # → (0.0078125, 1.1920928955078125e-07) 1353``` 1354 1355The formats that matter for training: 1356 1357| Format | Exponent bits | Fraction bits | Largest | Smallest normal | Gap after 1 | 1358|---|---|---|---|---|---| 1359| fp32 | 8 | 23 | 3.4 × 10³⁸ | 1.2 × 10⁻³⁸ | 1.2 × 10⁻⁷ | 1360| bf16 | 8 | 7 | 3.4 × 10³⁸ | 1.2 × 10⁻³⁸ | 0.0078 | 1361| fp16 | 5 | 10 | 65,504 | 6.1 × 10⁻⁵ | 0.00098 | 1362| fp8 E5M2 | 5 | 2 | 57,344 | 6.1 × 10⁻⁵ | 0.25 | 1363| fp8 E4M3 | 4 | 3 | 448 | 0.016 | 0.125 | 1364 1365Below the smallest normal number a format has a few **subnormal** values 1366that trade precision for extra range (fp16 reaches down to 6.0 × 10⁻⁸), and 1367below half of the smallest subnormal a number becomes 0: **underflow**. 1368Above the largest value is **overflow**, which becomes infinity. 1369 1370 1371 1372**Reading it:** each bar covers the magnitudes one format can represent, 1373on a log scale where each tick is a factor of 10. The darker part is the 1374normal range and the lighter tail on the left is the subnormals. bf16's bar 1375is as long as fp32's: same 8 exponent bits, same range. fp16's is a small 1376fraction of it, and fp8's smaller still. Range is what decides whether a 1377gradient survives; precision (not shown) decides how finely it is recorded. 1378 1379**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. 1380 1381### 4b. Underflow, and loss scaling 1382 1383**Everyday picture.** A kitchen scale that reads to the nearest gram 1384shows 0 for a pinch of saffron. Weigh the saffron together with a known 13851 kg jar, subtract the jar afterwards, and the pinch shows up. 1386 1387Gradients are often tiny. Many are smaller than fp16's smallest subnormal 1388(about 6 × 10⁻⁸), so a backward pass in fp16 silently turns them to 0 and 1389those weights stop learning. **Loss scaling** is the jar: multiply the loss 1390by a large number S before the backward pass, so every gradient is S times 1391larger (the chain rule passes the factor through unchanged); then divide 1392by S in fp32, before the update. 1393 1394$$ 1395g = \frac{\operatorname{fp16}\!\left(S \cdot \nabla \mathcal{L}\right)}{S} 1396$$ 1397 1398**Symbols** 1399 1400| Symbol | Meaning here | In the example | 1401|---|---|---| 1402| $\nabla \mathcal{L}$ | the true gradient of one weight | 10⁻⁸ | 1403| $S$ | the loss scale, usually a power of two so multiplying is exact | 65,536 = 2¹⁶ | 1404| $\operatorname{fp16}(\ldots)$ | "rounded to the nearest fp16 value", which is where the backward pass stores it | | 1405| $g$ | the gradient the optimizer receives, after dividing in fp32 | ≈ 10⁻⁸ | 1406 1407**In words:** "scale the loss up so the gradients are big enough for fp16, 1408compute them in fp16, and scale them back down in fp32." 1409 1410**With the numbers:** unscaled, 10⁻⁸ is under half of fp16's smallest value 1411(5.96 × 10⁻⁸) and rounds to **0**. Scaled, 10⁻⁸ × 65,536 = 6.55 × 10⁻⁴, a 1412normal fp16 number; it is stored as 6.5517 × 10⁻⁴ and divided back to 1413**9.997 × 10⁻⁹**, within 0.03% of the truth. 1414 1415**In Python:** 1416 1417```python 1418import numpy as np 1419grad, S = 1e-8, 65536 1420# without scaling: fp16 flushes it to zero 1421float(np.float16(grad)) # → 0.0 1422# with scaling: fp16 holds S · grad, then divide in higher precision 1423f"{float(np.float16(S * grad)) / S:.4e}" # → '9.9972e-09' 1424``` 1425 1426 1427 1428**Reading it:** the horizontal axis is a gradient's size in powers of two 1429and the height is how many gradients have that size. The shaded region on 1430the left is below fp16's smallest subnormal: whatever lands there becomes 14310. The grey histogram is the unscaled gradients, with the lost share 1432printed; the blue one is the same gradients times 2¹⁶, shifted 16 steps to 1433the right, clear of the shaded region and still far from the overflow wall 1434on the right. Loss scaling does not change the shape, only where it sits. 1435 1436A fixed scale is fragile: too small and gradients underflow, too large and 1437they overflow to infinity. **Dynamic loss scaling** adapts it: if any 1438gradient is infinite or not-a-number, skip the step and halve S; after a 1439long run of clean steps (2,000 is common), double S. 1440 1441```mermaid 1442flowchart TB 1443 M[fp32 master weights] -->|cast| W16[16-bit weights] 1444 W16 --> F[Forward pass<br/>16-bit matmuls] 1445 F --> L[Loss, in fp32] 1446 L -->|times S| B[Backward pass<br/>16-bit gradients] 1447 B --> CK{Any inf or NaN?} 1448 CK -- yes --> SK[Skip the step<br/>halve S] 1449 CK -- no --> U[Divide by S in fp32<br/>clip, then Adam update] 1450 U --> M 1451 SK --> M 1452``` 1453 1454**Reading it:** one training step goes round the loop once. The weights live 1455in fp32 (top) but are cast to 16 bits for the expensive forward and backward 1456passes. The loss is multiplied by S before the backward pass. The diamond is 1457dynamic loss scaling's check: an overflow means S was too big, so the step 1458is thrown away and S halved; otherwise the gradients are unscaled, clipped 1459and applied to the fp32 master copy. With **bf16** the scale can usually be 1460dropped altogether, because bf16 has fp32's range: the same 10⁻⁸ gradient is 1461stored as 1.0012 × 10⁻⁸ without any help. That is why bf16 became the 1462default for training on hardware that supports it. 1463 1464**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. 1465 1466### 4c. Why the master copy stays in fp32 1467 1468**Everyday picture.** Pour a teaspoon of water into a full bathtub and 1469measure with a bucket: the level has not changed, as far as the bucket can 1470tell. Do it a thousand times and you have added four litres, but every 1471single measurement still reads "no change". 1472 1473**Tiny worked example.** A weight of 1.0 receives an update of +0.0001 per 1474step. Next to 1.0, bf16 can only step in increments of 0.0078, so 1.0001 1475rounds straight back to 1.0. After 1,000 updates the bf16 weight is still 1476exactly 1.0; an fp32 weight has moved to 1.1, as it should. 1477 1478That is why the optimizer keeps an fp32 **master copy** of the weights and 1479applies updates to it, even though the matrix multiplies use 16-bit copies. 1480The 16-bit copy is re-made from the master after every step. 1481 1482**In Python:** 1483 1484```python 1485import numpy as np 1486w16, w32 = np.float16(1.0), np.float32(1.0) 1487for _ in range(1000): 1488 w16 = np.float16(w16 + np.float16(1e-4)) 1489 w32 = np.float32(w32 + np.float32(1e-4)) 1490# every update lost in 16 bits, all kept in 32 1491float(w16), round(float(w32), 4) # → (1.0, 1.1) 1492``` 1493 1494 1495 1496**Reading it:** the horizontal axis is the step and the vertical axis the 1497weight's value. The fp32 line rises steadily by 0.0001 per step. The bf16 1498and fp16 lines are flat at 1.0 for the whole run: each update is smaller 1499than half the gap to the next representable number, so it is rounded away 1500every time. The error is not noise that averages out; it is a systematic 1501loss of every small update, which is exactly what late training consists of. 1502 1503**In code:** `accumulate_updates` adds the same update many times, rounding the weight to a chosen format after each step. 1504 1505### 4d. fp8: the next halving 1506 1507**Everyday picture.** A ruler with only eight marks is useless for 1508measuring a hair, unless you first slide it under a magnifying glass set to 1509the right zoom. fp8 is that short ruler, and a per-tensor scale is the 1510magnifying glass. 1511 1512Recent GPUs multiply 8-bit floats at twice the 16-bit rate. With so few 1513bits, one format cannot cover everything, so training uses two: **E4M3** 1514(more precision, range to 448) for weights and activations, and **E5M2** 1515(more range, to 57,344) for gradients. Because the range is so short, every 1516tensor (or every small block of a tensor) gets its own scale factor, 1517chosen from its recent largest value, the same idea as loss scaling applied 1518everywhere. DeepSeek-V3 was trained largely in fp8 this way, keeping 1519sensitive parts (normalizations, the optimizer, the master weights) in 1520higher precision. 1521 1522## 5. Stability at scale: loss spikes, clipping, warmup and checkpoints 1523 1524**Everyday picture.** A ship on a months-long voyage does not assume the 1525sea stays calm. It trims the sails when gusts come (clipping), leaves port 1526slowly (warmup), keeps a log of its position (checkpoints), and has a 1527drill for when something goes wrong (rolling back). 1528 1529A large run does see storms. The loss curve occasionally jumps upward, a 1530**loss spike**, sometimes recovering by itself and sometimes diverging for 1531good. Causes include a batch of bad data, a learning rate a little too high 1532for the model's current state, attention logits growing without bound, and 1533numerical overflow. And separately from the maths, hardware fails: in the 1534Llama 3 405B run, 419 unexpected interruptions occurred over 54 days, 1535about one every three hours. 1536 1537```mermaid 1538flowchart LR 1539 T[Train step] --> W{Loss well above<br/>recent median?} 1540 W -- no --> CP{Checkpoint<br/>interval reached?} 1541 CP -- yes --> SV[Save weights,<br/>optimizer, data position] 1542 CP -- no --> T 1543 SV --> T 1544 W -- "yes, and it persists" --> RB[Roll back to a<br/>checkpoint before the spike] 1545 RB --> SKIP[Skip the batches<br/>around the spike] 1546 SKIP --> T 1547``` 1548 1549**Reading it:** the main loop is train, check, maybe save, repeat. The 1550upper diamond watches the loss. A brief blip is ignored, since clipping 1551usually absorbs it; a spike that persists sends the run back to an earlier 1552checkpoint, and the data batches that were in flight around the spike are 1553skipped. PaLM's authors did exactly this: restart about 100 steps before the 1554spike and skip roughly 200 to 500 batches, which removed the spikes. A 1555checkpoint must include the optimizer state and the position in the data, 1556or the resumed run is not the same run. 1557 1558### 5a. Spotting a spike 1559 1560**Tiny worked example.** Losses 3.0, 2.9, 2.8, 2.8, 2.7, then **8.5**. The 1561median of the previous four is 2.8, and 8.5 is more than twice that, so step 15625 is flagged. The next step, 4.0, is compared with the median of 2.8, 2.8, 15632.7, 8.5, which is still 2.8: one spike does not raise the bar, which is 1564why the rule uses the median and not the mean. 1565 1566$$ 1567\text{spike at } t \iff \mathcal{L}_t > f \cdot \operatorname{median}\big(\mathcal{L}_{t-w}, \ldots, \mathcal{L}_{t-1}\big) 1568$$ 1569 1570**Symbols** 1571 1572| Symbol | Meaning here | In the example | 1573|---|---|---| 1574| $\mathcal{L}_t$ | the training loss at step $t$ | 8.5 at $t = 5$ | 1575| $w$ | how many recent steps to look back over (the window) | 4 | 1576| $\operatorname{median}$ | the middle value after sorting; one outlier cannot move it far | 2.8 | 1577| $f$ | how many times the median counts as a spike | 2 | 1578| $\iff$ | "exactly when" | | 1579 1580**In words:** "flag a step when its loss is more than f times the median of 1581the last w losses." 1582 1583**With the numbers:** median(2.9, 2.8, 2.8, 2.7) = 2.8, and 8.5 > 2 × 2.8 = 15845.6, so step 5 is flagged; 4.0 < 5.6, so step 6 is not. 1585 1586**In Python:** 1587 1588```python 1589import statistics 1590losses = [3.0, 2.9, 2.8, 2.8, 2.7, 8.5, 4.0, 2.7] 1591w, f = 4, 2.0 1592# L_t > f · median(L_{t−w}, …, L_{t−1}) 1593[t for t in range(w, len(losses)) if losses[t] > f * statistics.median(losses[t - w:t])] # → [5] 1594``` 1595 1596**In code:** `detect_spikes` applies the median rule to a whole loss history. 1597 1598### 5b. Gradient clipping, when the gradient is spread across GPUs 1599 1600Clipping by global norm (built in `primer.ml.optimizers`) caps the length 1601of the whole gradient vector: if it is longer than c, shrink every entry by 1602the same factor. At scale there is a twist. With ZeRO or tensor parallelism 1603no GPU holds the whole gradient, so none can measure its length alone. Each 1604GPU sums the squares of its own shard, one all-reduce adds those sums (a 1605single number per GPU, so it is nearly free), and every GPU takes the 1606square root and applies the same factor. 1607 1608$$ 1609\lVert g \rVert = \sqrt{\sum_{k=1}^{G} \sum_{j} g_{kj}^2}, 1610\qquad 1611g \leftarrow g \cdot \min\!\left(1, \frac{c}{\lVert g \rVert}\right) 1612$$ 1613 1614**Symbols** 1615 1616| Symbol | Meaning here | In the example | 1617|---|---|---| 1618| $g$ | the whole gradient, all parameters | (1, 2, 2, 4) | 1619| $g_{kj}$ | entry $j$ of the shard on GPU $k$ | GPU 1: (1, 2); GPU 2: (2, 4) | 1620| $\sum_j g_{kj}^2$ | GPU $k$'s local sum of squares | 5 and 20 | 1621| $\lVert g \rVert$ | the length (**norm**) of the whole gradient | 5 | 1622| $c$ | the clipping threshold | 1 | 1623| $\min(1, \ldots)$ | never scale up, only down | 0.2 | 1624| $\leftarrow$ | "is replaced by" | | 1625 1626**In words:** "add up every GPU's sum of squares, take the square root to 1627get the total length, and if it is over the limit, shrink every entry on 1628every GPU by the same factor." 1629 1630**With the numbers:** 5 + 20 = 25, √25 = 5; with c = 1 the factor is 1/5, so 1631GPU 1's shard becomes (0.2, 0.4) and GPU 2's (0.4, 0.8). 1632 1633**In Python:** 1634 1635```python 1636import math 1637shards = [[1.0, 2.0], [2.0, 4.0]] 1638# each GPU's local sum of squares 1639local = [sum(x * x for x in s) for s in shards] 1640local # → [5.0, 20.0] 1641# all-reduce the sums, then the square root 1642norm = math.sqrt(sum(local)) 1643norm # → 5.0 1644c = 1.0 1645factor = min(1.0, c / norm) 1646[[x * factor for x in s] for s in shards] # → [[0.2, 0.4], [0.4, 0.8]] 1647``` 1648 1649 1650 1651**Reading it:** the horizontal axis is the training step and the vertical 1652axis the loss on clean data, on a log scale. Both runs have converged by 1653step 30, when one batch arrives with corrupted labels. Without clipping 1654(red), that single gradient is hundreds of times too large, the weight is 1655thrown far off, and the loss jumps to over 6,000 before slowly recovering. 1656With clipping at 1 (blue), the same batch can only move the weight by one 1657learning-rate step, and the loss rises to about 0.01. Clipping cannot tell a 1658bad batch from a good one; it just limits how much damage any one batch can 1659do. 1660 1661**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`. 1662 1663### 5c. Warmup 1664 1665**Everyday picture.** Nobody floors the accelerator in a car they have 1666never driven; they ease on until they know how it responds. 1667 1668At the start of training the weights are random and Adam's running averages 1669have seen only a handful of gradients, so its step sizes are unreliable. A 1670full learning rate at step 1 can throw the model into a region it never 1671recovers from. **Warmup** ramps the learning rate linearly from 0 to its 1672peak over the first few hundred to few thousand steps; `primer.ml.optimizers` 1673derives the warmup-then-cosine schedule with worked numbers. At scale it 1674matters more, not less: bigger models and bigger batches tolerate smaller 1675peak learning rates, and a spike early in a months-long run wastes the most. 1676 1677### 5d. Checkpoints: how often to save 1678 1679**Everyday picture.** Saving a long document every few seconds wastes time 1680on saving; saving once an hour risks losing an hour of work to a crash. 1681Somewhere in between is the least total waste. 1682 1683**Tiny worked example.** Suppose writing a checkpoint pauses training for 1 1684minute, and the cluster suffers a failure every 180 minutes on average. 1685Saving every 19 minutes spends 1/19 of the time saving, and each failure 1686loses on average half an interval, 9.5 minutes, every 180 minutes. Both 1687costs come to about 5.3%, and their total, 10.5%, is the smallest possible. 1688 1689$$ 1690\text{waste}(T) = \frac{C}{T} + \frac{T}{2M}, 1691\qquad 1692T^{*} = \sqrt{2\,C\,M} 1693$$ 1694 1695**Symbols** 1696 1697| Symbol | Meaning here | In the example | 1698|---|---|---| 1699| $T$ | time between checkpoints | 19 minutes | 1700| $C$ | time to write one checkpoint | 1 minute | 1701| $M$ | mean time between failures for the whole cluster | 180 minutes | 1702| $C/T$ | share of time spent saving | 1/19 | 1703| $T/(2M)$ | share of time redoing lost work: on average half an interval per failure | 9.5/180 | 1704| $T^{*}$ | the interval with the least total waste (Young, 1974) | 18.97 minutes | 1705 1706**In words:** "saving often wastes time saving and saving rarely wastes 1707time redoing; the best interval is the square root of twice the save time 1708times the time between failures." 1709 1710**With the numbers:** √(2 × 1 × 180) = √360 = **18.97 minutes**, and the 1711waste is 1/18.97 + 18.97/360 = 0.053 + 0.053 = **10.5%**. Cut the save time to 171210 seconds (by writing asynchronously, in the background, from every GPU's 1713shard at once) and the best interval falls to 7.7 minutes with only 4.3% 1714waste. 1715 1716**In Python:** 1717 1718```python 1719import math 1720C, M = 1.0, 180.0 1721# T* = √(2 C M) 1722T = math.sqrt(2 * C * M) 1723round(T, 2) # → 18.97 1724# waste(T) = C/T + T/(2M) 1725round(C / T + T / (2 * M), 3) # → 0.105 1726# a 10-second save 1727C = 1 / 6 1728round(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) 1729``` 1730 1731 1732 1733**Reading it:** the horizontal axis is the checkpoint interval in minutes 1734(log scale) and the vertical axis the share of time lost. Each curve is a 1735valley: on the left, saving too often; on the right, losing too much work 1736per failure. The dots mark √(2CM), the bottom of each valley. A faster save 1737moves the whole valley down and to the left, which is why large training 1738systems invest heavily in fast, asynchronous checkpointing: at thousands of 1739GPUs, failures are not an exception but the weather. 1740 1741**In code:** `wasted_fraction` is the waste formula and `optimal_checkpoint_interval` is Young's square-root rule. 1742 1743## In 20 seconds 1744 1745- **Data:** pretraining data is mostly web crawl, pushed through language 1746 ID, quality rules, a quality classifier, and exact and near-duplicate 1747 removal (MinHash with LSH); most of the crawl is discarded. Duplicates 1748 cause memorization and waste compute. 1749- **Mixture and budget:** sources are mixed by weight, not size; a 1750 compute-optimal budget is about 20 tokens per parameter, compute is about 1751 6·N·D, and models meant for heavy use are trained far longer. 1752- **Memory:** Adam in mixed precision needs 16 bytes per parameter before 1753 activations (112 GB for 7B), so training needs many GPUs. 1754- **Parallelism:** data parallel (all-reduce gradients), ZeRO/FSDP (shard 1755 the state), tensor parallel (split each matmul, inside a server) and 1756 pipeline parallel (split the layers, mind the bubble (p − 1)/(m + p − 1)). 1757- **Precision:** compute in bf16 (fp32's range, less precision) or fp8 with 1758 scaling; fp16 needs loss scaling; the master weights stay in fp32 so small 1759 updates are not rounded away. 1760- **Stability:** warmup, global-norm clipping, spike detection with 1761 rollback, and checkpoints every √(2·C·M). 1762 1763## Self-test questions 1764 1765**Why does a pretraining pipeline run deduplication after the quality 1766filters, not before?** 1767Language ID and quality rules look at one page at a time, so they are cheap 1768per page. Near-duplicate detection compares pages with one another, which is 1769the expensive step. Running the cheap filters first means the expensive one 1770sees far fewer pages. 1771 1772**How can MinHash estimate the overlap of two documents without comparing 1773their contents?** 1774Under a random ordering of all shingles, two sets have the same first 1775member with probability equal to their Jaccard similarity. A signature 1776records each document's first member under k random hash functions, so the 1777fraction of matching slots estimates the Jaccard. LSH then groups 1778signatures by bands so only likely pairs are ever compared. 1779 1780**Why do duplicated documents hurt a model, when more data usually helps?** 1781A repeated document gets many times the training signal of any other, so 1782the model memorizes it and tends to regurgitate it; the repeats also spend 1783compute that would have taught something new, and copies of benchmark 1784questions contaminate evaluations. 1785 1786**Where do the 16 bytes per parameter come from, and what do they mean for 1787a 7B model?** 17882 bytes for the bf16 weight, 2 for its gradient, and 12 for fp32 state: the 1789master weight and Adam's two running averages. 16 × 7 × 10⁹ = 112 GB, more 1790than one 80 GB GPU holds, before any activations. 1791 1792**What does each ZeRO stage shard, and what does it cost?** 1793Stage 1 shards the optimizer state, stage 2 also the gradients, stage 3 1794(FSDP) also the weights, dividing each by the number of GPUs. Stages 1 and 2 1795cost no more communication than plain data parallelism; stage 3 adds 1796all-gathers of each layer's weights in both passes, about 1.5 times the 1797traffic. 1798 1799**Why is tensor parallelism kept inside one server while pipeline 1800parallelism spans servers?** 1801Tensor parallelism exchanges partial results inside every layer, so it 1802needs the fastest links, which exist only between GPUs in the same server. 1803Pipeline parallelism only passes activations at stage boundaries, a small 1804and infrequent exchange that slower links between servers can carry. 1805 1806**What is the pipeline bubble, and how do you shrink it?** 1807The time stages sit idle while the pipeline fills and drains: (p − 1)/(m + 1808p − 1) of the schedule for p stages and m micro-batches. More micro-batches 1809shrink it (4 stages: 75% with 1, 8.6% with 32), as do schedules that 1810interleave forward and backward passes. 1811 1812**Why does fp16 training need loss scaling while bf16 usually does not?** 1813fp16 has 5 exponent bits, so its smallest value is about 6 × 10⁻⁸ and many 1814gradients underflow to zero; multiplying the loss by a large scale lifts 1815them into range. bf16 keeps fp32's 8 exponent bits and therefore its range, 1816giving up precision instead. 1817 1818**Why keep an fp32 copy of the weights if the maths runs in 16 bits?** 1819Late in training, updates are tiny compared with the weights. Next to 1.0 1820the bf16 grid spacing is about 0.008, so an update of 0.0001 rounds away 1821completely, every step. Applying updates to an fp32 master copy keeps them. 1822 1823**How often should a large run write checkpoints?** 1824Roughly every √(2·C·M), where C is the time to save and M the mean time 1825between failures: saving more often wastes time saving, less often wastes 1826work redone after failures. Faster, asynchronous saves allow more frequent 1827checkpoints and less waste. 1828 1829## The papers behind this lesson 1830 1831- **Rae et al., *Scaling Language Models: Methods, Analysis & Insights from 1832 Training Gopher* (2021)**: https://arxiv.org/abs/2112.11446. Among much 1833 else, published the simple quality rules (length, symbols, stop words) 1834 that many open pipelines still apply. 1835- **Lee et al., *Deduplicating Training Data Makes Language Models Better* 1836 (2021)**: https://arxiv.org/abs/2107.06499. Showed that exact and MinHash 1837 near-duplicate removal cuts verbatim memorization about tenfold without 1838 hurting quality. 1839- **Penedo et al., *The FineWeb Datasets: Decanting the Web for the Finest 1840 Text Data at Scale* (2024)**: https://arxiv.org/abs/2406.17557. Documented 1841 and ablated a full open curation pipeline over 96 Common Crawl snapshots, 1842 including the classifier-filtered FineWeb-Edu. 1843- **Hoffmann et al., *Training Compute-Optimal Large Language Models* 1844 (2022)**: https://arxiv.org/abs/2203.15556. Found that parameters and 1845 training tokens should grow together, about 20 tokens per parameter. 1846 [Annotated companion](../../papers/scaling-laws.html) 1847- **Shumailov et al., *The Curse of Recursion: Training on Generated Data 1848 Makes Models Forget* (2023)**: https://arxiv.org/abs/2305.17493. Showed 1849 that models trained recursively on their own outputs lose the tails of the 1850 original distribution: model collapse. 1851- **Rajbhandari et al., *ZeRO: Memory Optimizations Toward Training 1852 Trillion Parameter Models* (2019)**: https://arxiv.org/abs/1910.02054. 1853 Introduced the 16-bytes-per-parameter accounting and the three stages of 1854 sharding the training state across data-parallel GPUs. 1855 [Annotated companion](../../papers/zero.html) 1856- **Shoeybi et al., *Megatron-LM: Training Multi-Billion Parameter Language 1857 Models Using Model Parallelism* (2019)**: https://arxiv.org/abs/1909.08053. 1858 Split transformer layers across GPUs by columns and rows, with one 1859 all-reduce per block. 1860 [Annotated companion](../../papers/megatron-lm.html) 1861- **Huang et al., *GPipe: Efficient Training of Giant Neural Networks using 1862 Pipeline Parallelism* (2018)**: https://arxiv.org/abs/1811.06965. Split a 1863 model into stages fed by micro-batches, and analysed the resulting bubble. 1864- **Micikevicius et al., *Mixed Precision Training* (2017)**: 1865 https://arxiv.org/abs/1710.03740. Introduced the fp16 recipe: an fp32 1866 master copy of the weights, loss scaling, and fp32 accumulation. 1867- **Korthikanti et al., *Reducing Activation Recomputation in Large 1868 Transformer Models* (2022)**: https://arxiv.org/abs/2205.05198. Counted 1869 the activation memory of a transformer layer and showed how to recompute 1870 only the parts that are cheap to recompute. 1871 1872## Further reading 1873 1874- Micikevicius et al., *FP8 Formats for Deep Learning* (2022): https://arxiv.org/abs/2209.05433 1875- Kalamkar et al., *A Study of BFLOAT16 for Deep Learning Training* (2019): https://arxiv.org/abs/1905.12322 1876- Zhao et al., *PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel* (2023): https://arxiv.org/abs/2304.11277 1877- Narayanan et al., *Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM* (2021): https://arxiv.org/abs/2104.04473 1878- Chen et al., *Training Deep Nets with Sublinear Memory Cost* (activation checkpointing, 2016): https://arxiv.org/abs/1604.06174 1879- Chowdhery et al., *PaLM: Scaling Language Modeling with Pathways* (loss spikes and rollback, 2022): https://arxiv.org/abs/2204.02311 1880- Touvron et al., *LLaMA: Open and Efficient Foundation Language Models* (data mixture, 2023): https://arxiv.org/abs/2302.13971 1881- Llama Team, *The Llama 3 Herd of Models* (4D parallelism and failures at 16K GPUs, 2024): https://arxiv.org/abs/2407.21783 1882- DeepSeek-AI, *DeepSeek-V3 Technical Report* (fp8 training, 2024): https://arxiv.org/abs/2412.19437 1883- PyTorch automatic mixed precision: https://pytorch.org/docs/stable/amp.html 1884- PyTorch FullyShardedDataParallel: https://pytorch.org/docs/stable/fsdp.html 1885""" 1886 1887from __future__ import annotations 1888 1889import hashlib 1890import math 1891import re 1892import zlib 1893from collections import Counter 1894from dataclasses import dataclass, field 1895 1896import numpy as np 1897 1898from primer._show import banner, say, table, takeaway 1899from primer.ml.optimizers import clip_by_global_norm 1900 1901# --------------------------------------------------------------------------- 1902# 1. Data: language ID, quality filters, deduplication, mixtures, budgets 1903# --------------------------------------------------------------------------- 1904 1905_WORD = re.compile(r"[a-zà-ÿ0-9']+") 1906 1907# The ten commonest little words of each language. Real language ID (fastText's 1908# lid.176, CLD3) learns from character n-grams across 100+ languages; counting 1909# function words is the same idea at toy scale. 1910STOP_WORDS = { 1911 "en": {"the", "and", "of", "to", "is", "in", "that", "it", "was", "for"}, 1912 "fr": {"le", "la", "les", "et", "est", "un", "une", "des", "du", "dans"}, 1913 "de": {"der", "die", "das", "und", "ist", "nicht", "ein", "eine", "zu", "mit"}, 1914 "es": {"el", "los", "las", "y", "es", "una", "que", "por", "con", "del"}, 1915} 1916 1917 1918def words(text: str) -> list[str]: 1919 """Lowercased words: the unit every filter below counts.""" 1920 return _WORD.findall(text.lower()) 1921 1922 1923def detect_language(text: str, min_share: float = 0.1) -> str: 1924 """The language whose stop words make up the largest share of the text, or "unknown". 1925 1926 A page of product codes or numbers matches no language's little words at 1927 all, so it falls under `min_share` and is reported as unknown. 1928 """ 1929 ws = words(text) 1930 if not ws: 1931 return "unknown" 1932 shares = {lang: sum(w in sw for w in ws) / len(ws) for lang, sw in STOP_WORDS.items()} 1933 best = max(shares, key=shares.get) 1934 return best if shares[best] >= min_share else "unknown" 1935 1936 1937# Rae et al. (2021), the Gopher paper's "MassiveText" quality rules, Appendix A. 1938GOPHER_STOP_WORDS = {"the", "be", "to", "of", "and", "that", "have", "with"} 1939 1940 1941def quality_failures(text: str) -> list[str]: 1942 """Names of the heuristic quality rules a document breaks (empty list = keep it). 1943 1944 Each rule is cheap, obvious once stated, and removes a whole genre of junk: 1945 navigation snippets, hashtag spam, keyword stuffing, lists of links. 1946 """ 1947 raw = text.split() 1948 n = len(raw) 1949 lines = [l for l in text.splitlines() if l.strip()] or [text] 1950 failures = [] 1951 if n < 50: 1952 failures.append("too_short") 1953 if n > 100_000: 1954 failures.append("too_long") 1955 if n and not 3 <= sum(len(w) for w in raw) / n <= 10: 1956 failures.append("odd_word_length") 1957 symbols = text.count("#") + text.count("...") + text.count("…") 1958 if n and symbols / n > 0.1: 1959 failures.append("too_many_symbols") 1960 if sum(l.lstrip().startswith(("-", "*", "•")) for l in lines) / len(lines) > 0.9: 1961 failures.append("mostly_bullets") 1962 if sum(l.rstrip().endswith(("...", "…")) for l in lines) / len(lines) > 0.3: 1963 failures.append("mostly_ellipsis") 1964 if n and sum(any(c.isalpha() for c in w) for w in raw) / n < 0.8: 1965 failures.append("too_few_alphabetic_words") 1966 if len(GOPHER_STOP_WORDS & set(words(text))) < 2: 1967 failures.append("missing_stop_words") 1968 return failures 1969 1970 1971# Tiny labelled sets for the classifier: "reference" text in the style of an 1972# encyclopedia or textbook, and "raw crawl" text in the style of spam pages. 1973GOOD_EXAMPLES = [ 1974 "the river deposits sand at its mouth and over many years builds a delta", 1975 "a delta is formed where a river meets the sea and slows down", 1976 "the cell is the basic unit of life and every organism is made of cells", 1977 "light travels faster than sound which is why we see lightning before thunder", 1978 "the history of the city begins with a small settlement on the river", 1979 "photosynthesis converts light into chemical energy stored in sugar", 1980] 1981BAD_EXAMPLES = [ 1982 "click here buy now limited offer best price", 1983 "free free free win a prize click now", 1984 "best deals cheap price buy buy buy", 1985 "subscribe now click the link for free coupons", 1986 "hot singles best price click here now", 1987 "lose weight fast buy now free shipping", 1988] 1989 1990 1991@dataclass 1992class QualityClassifier: 1993 """A bag-of-words naive Bayes scorer: log-odds that a text looks like the reference set. 1994 1995 GPT-3 filtered Common Crawl with a linear classifier of this kind (reference 1996 text vs. raw crawl); FineWeb-Edu used an LLM's educational-value ratings to 1997 train one. The score is a sum of per-word votes, so it is fast enough to 1998 run over billions of pages. 1999 """ 2000 2001 good: Counter = field(default_factory=Counter) 2002 bad: Counter = field(default_factory=Counter) 2003 2004 @classmethod 2005 def trained_on_examples(cls, good=GOOD_EXAMPLES, bad=BAD_EXAMPLES) -> "QualityClassifier": 2006 clf = cls() 2007 for t in good: 2008 clf.good.update(words(t)) 2009 for t in bad: 2010 clf.bad.update(words(t)) 2011 return clf 2012 2013 def word_vote(self, w: str) -> float: 2014 """log P(w | good) − log P(w | bad), with add-one smoothing so unseen words never give log 0.""" 2015 vocab = len(set(self.good) | set(self.bad)) 2016 p_good = (self.good[w] + 1) / (sum(self.good.values()) + vocab) 2017 p_bad = (self.bad[w] + 1) / (sum(self.bad.values()) + vocab) 2018 return math.log(p_good / p_bad) 2019 2020 def score(self, text: str) -> float: 2021 """Sum of word votes over words the classifier has seen. Above 0 leans reference, below 0 leans spam.""" 2022 known = set(self.good) | set(self.bad) 2023 return sum(self.word_vote(w) for w in words(text) if w in known) 2024 2025 2026def normalize_for_exact(text: str) -> str: 2027 """Lowercase and collapse whitespace, so trivially reformatted copies hash the same.""" 2028 return " ".join(text.lower().split()) 2029 2030 2031def exact_dedup(docs: list[str]) -> list[str]: 2032 """Keep the first copy of each document, comparing a hash of its normalized text.""" 2033 seen, kept = set(), [] 2034 for d in docs: 2035 key = hashlib.sha1(normalize_for_exact(d).encode()).hexdigest() 2036 if key not in seen: 2037 seen.add(key) 2038 kept.append(d) 2039 return kept 2040 2041 2042def shingles(text: str, k: int = 5) -> set[str]: 2043 """The set of overlapping k-word windows ("shingles") in a text.""" 2044 ws = words(text) 2045 return {" ".join(ws[i : i + k]) for i in range(max(1, len(ws) - k + 1))} 2046 2047 2048def jaccard(a: set, b: set) -> float: 2049 """|A ∩ B| / |A ∪ B|: shared shingles over all distinct shingles.""" 2050 return len(a & b) / len(a | b) if a | b else 1.0 2051 2052 2053_MERSENNE = (1 << 31) - 1 # a prime, and a·x + b stays below 2⁶³ so int64 never overflows 2054 2055 2056class MinHasher: 2057 """k random hash functions; a document's signature is its smallest hash under each. 2058 2059 Two documents agree in one slot with probability exactly their Jaccard 2060 similarity, so the fraction of agreeing slots estimates it, from k numbers 2061 per document instead of the whole shingle set. 2062 """ 2063 2064 def __init__(self, num_perm: int = 128, seed: int = 0): 2065 rng = np.random.default_rng(seed) 2066 # h_i(x) = (a_i·x + b_i) mod p: a cheap family that behaves like random orderings. 2067 self.a = rng.integers(1, _MERSENNE, num_perm, dtype=np.int64) 2068 self.b = rng.integers(0, _MERSENNE, num_perm, dtype=np.int64) 2069 2070 def signature(self, shingle_set: set[str]) -> np.ndarray: 2071 # crc32 rather than hash(): Python salts hash() per process, which would break determinism. 2072 x = np.array([zlib.crc32(s.encode()) % _MERSENNE for s in shingle_set], dtype=np.int64) 2073 if x.size == 0: 2074 return np.full(self.a.shape, _MERSENNE, dtype=np.int64) 2075 return ((self.a[:, None] * x[None, :] + self.b[:, None]) % _MERSENNE).min(axis=1) # (num_perm,) 2076 2077 2078def estimate_jaccard(sig_a: np.ndarray, sig_b: np.ndarray) -> float: 2079 """Fraction of signature slots where the two minimums agree.""" 2080 return float(np.mean(sig_a == sig_b)) 2081 2082 2083def lsh_candidate_probability(s: float, bands: int, rows: int) -> float: 2084 """Chance that a pair with Jaccard s shares at least one whole band: 1 − (1 − sʳ)ᵇ.""" 2085 return 1 - (1 - s**rows) ** bands 2086 2087 2088def near_duplicate_pairs(docs: list[str], bands: int = 20, rows: int = 5, threshold: float = 0.7, seed: int = 0) -> list[tuple[int, int]]: 2089 """Pairs (i, j), i < j, whose MinHash similarity is at least `threshold`. 2090 2091 Instead of comparing all n² pairs, each signature is cut into `bands` bands 2092 of `rows` numbers and each band is hashed into a bucket. Only documents that 2093 land in the same bucket for some band are compared at all. 2094 """ 2095 hasher = MinHasher(bands * rows, seed) 2096 sigs = [hasher.signature(shingles(d)) for d in docs] 2097 buckets: dict[tuple, list[int]] = {} 2098 for i, sig in enumerate(sigs): 2099 for band in range(bands): 2100 buckets.setdefault((band, *sig[band * rows : (band + 1) * rows].tolist()), []).append(i) 2101 candidates = {(i, j) for ids in buckets.values() for i in ids for j in ids if i < j} 2102 return sorted((i, j) for i, j in candidates if estimate_jaccard(sigs[i], sigs[j]) >= threshold) 2103 2104 2105CRAWL_SAMPLE = [ 2106 # 0: a clean English paragraph 2107 "The river carries sand and small stones from the mountains to the sea. Over thousands of years " 2108 "that sand settles at the mouth of the river and builds a wide, flat delta. Farmers have grown rice " 2109 "on these deltas for centuries, because the soil is rich and the water is close. When the river " 2110 "floods, it brings new soil with it, which is why the land stays fertile.", 2111 # 1: a clean French paragraph (kept by a French pipeline, not by this English one) 2112 "Le fleuve transporte le sable et les pierres des montagnes vers la mer. Avec le temps, le sable se " 2113 "dépose à l'embouchure et forme un delta. Les paysans cultivent le riz dans les deltas depuis des " 2114 "siècles, car la terre est riche et l'eau est proche.", 2115 # 2: a navigation bar 2116 "Home | About us | Contact | Privacy policy | Log in", 2117 # 3: document 0 again, re-crawled with different capitalization and spacing 2118 "the river carries sand and small stones from the mountains to the sea. Over thousands of years " 2119 "that sand settles at the mouth of the river and builds a wide, flat delta. Farmers have grown rice " 2120 "on these deltas for centuries, because the soil is rich and the water is close. When the river " 2121 "floods, it brings new soil with it, which is why the land stays fertile.", 2122 # 4: hashtag spam 2123 " ".join(["#deal #sale #win the best price today"] * 12), 2124 # 5: document 0, lightly edited by a scraper site 2125 "The river carries sand and small stones from the mountains to the sea. Over thousands of years " 2126 "that sand settles at the mouth of the river and builds a wide, flat delta. Farmers have grown rice " 2127 "on these deltas for hundreds of years, because the soil is rich and the water is close. When the " 2128 "river floods, it brings new soil with it, which is why the land stays fertile. Read more on our site.", 2129 # 6: a different clean English paragraph 2130 "Stars are born inside cold clouds of gas and dust. When part of a cloud becomes dense enough, gravity " 2131 "pulls it together faster than the gas can push back. The centre heats up as it shrinks, and once it " 2132 "reaches about ten million degrees, hydrogen begins to fuse into helium. That is the moment a star " 2133 "switches on, and it will shine with that fuel for millions or billions of years.", 2134] 2135 2136 2137def curate(docs: list[str], language: str = "en") -> list[str]: 2138 """Run the pipeline in order (language, quality, exact dedup, near dedup); one verdict per doc. 2139 2140 Cheap checks run first so the expensive ones see fewer documents: language 2141 ID and the heuristics are a pass over each page, while near-duplicate 2142 detection compares pages with each other. 2143 """ 2144 verdicts: list[str] = [] 2145 kept_hashes: dict[str, int] = {} 2146 kept: list[int] = [] 2147 hasher = MinHasher(100, seed=0) 2148 kept_sigs: dict[int, np.ndarray] = {} 2149 for i, d in enumerate(docs): 2150 lang = detect_language(d) 2151 if lang != language: 2152 verdicts.append(f"language: {lang}") 2153 continue 2154 failures = quality_failures(d) 2155 if failures: 2156 verdicts.append("quality: " + ", ".join(failures)) 2157 continue 2158 key = hashlib.sha1(normalize_for_exact(d).encode()).hexdigest() 2159 if key in kept_hashes: 2160 verdicts.append(f"exact duplicate of {kept_hashes[key]}") 2161 continue 2162 sig = hasher.signature(shingles(d)) 2163 twin = next((j for j in kept if estimate_jaccard(sig, kept_sigs[j]) >= 0.7), None) 2164 if twin is not None: 2165 verdicts.append(f"near duplicate of {twin}") 2166 continue 2167 kept_hashes[key], kept_sigs[i] = i, sig 2168 kept.append(i) 2169 verdicts.append("kept") 2170 return verdicts 2171 2172 2173BOILERPLATE = "click here to subscribe for free updates" 2174OTHER_LINES = [ 2175 "click the map to see the river", 2176 "we walk here every day", 2177 "we went to the river", 2178 "subscribe for the weekly news", 2179 "free maps for every school", 2180] 2181 2182 2183def verbatim_probability(copies: int, line: str = BOILERPLATE) -> float: 2184 """Probability a bigram model, trained on the small corpus plus `copies` of the line, continues the line's 2185 first word into the whole line. 2186 2187 A bigram model predicts each word from the one before it, by counting. 2188 Every duplicate adds the same counts again, so the line's path through the 2189 counts becomes the only road: that is memorization in its simplest form. 2190 """ 2191 corpus = OTHER_LINES + [line] * copies 2192 pairs = Counter((a, b) for s in corpus for a, b in zip(s.split(), s.split()[1:])) 2193 firsts = Counter(a for (a, _), c in pairs.items() for _ in range(c)) 2194 ws = line.split() 2195 return float(np.prod([pairs[(a, b)] / firsts[a] for a, b in zip(ws, ws[1:])])) 2196 2197 2198def epochs_per_source(weights: dict[str, float], sizes: dict[str, float], budget: float) -> dict[str, float]: 2199 """How many passes over each source a mixture implies: weight × budget / tokens available.""" 2200 return {k: weights[k] * budget / sizes[k] for k in weights} 2201 2202 2203def chinchilla_tokens(params: float) -> float: 2204 """Compute-optimal training tokens for a model size: about 20 per parameter (Hoffmann et al., 2022).""" 2205 return 20 * params 2206 2207 2208def training_flops(params: float, tokens: float) -> float: 2209 """C ≈ 6·N·D: 2 FLOPs per parameter per token forward, 4 backward.""" 2210 return 6 * params * tokens 2211 2212 2213def recursive_gaussian_fit(generations: int = 200, n: int = 20, seed: int = 0) -> list[float]: 2214 """Model collapse in miniature: fit a bell curve, sample from the fit, refit on the samples, repeat. 2215 2216 Returns the fitted standard deviation of every generation. Each refit on a 2217 finite sample loses a little of the tails, and the losses compound. 2218 """ 2219 rng = np.random.default_rng(seed) 2220 mu, sigma = 0.0, 1.0 2221 stds = [] 2222 for _ in range(generations): 2223 sample = rng.normal(mu, sigma, n) 2224 mu, sigma = float(sample.mean()), float(sample.std()) # maximum-likelihood fit (divides by n) 2225 stds.append(sigma) 2226 return stds 2227 2228 2229# --------------------------------------------------------------------------- 2230# 2. Memory: why one GPU is not enough 2231# --------------------------------------------------------------------------- 2232 2233 2234def training_memory(params: float) -> dict[str, float]: 2235 """Bytes of training state for Adam in mixed precision, the "16 bytes per parameter" rule.""" 2236 parts = { 2237 "weights (bf16)": 2 * params, 2238 "gradients (bf16)": 2 * params, 2239 "master weights (fp32)": 4 * params, 2240 "Adam momentum (fp32)": 4 * params, 2241 "Adam variance (fp32)": 4 * params, 2242 } 2243 return {**parts, "total": sum(parts.values())} 2244 2245 2246def activation_bytes(seq: int, batch: int, hidden: int, heads: int, layers: int, 2247 store_scores: bool = True, checkpointed: bool = False) -> float: 2248 """Activations saved for the backward pass, in bytes (Korthikanti et al., 2022, no parallelism). 2249 2250 Per layer: s·b·h·(34 + 5·a·s/h). The 34·s·b·h part is the layer's ordinary 2251 intermediate tensors; the 5·a·s² part is attention's score matrices, which 2252 FlashAttention-style kernels recompute instead of storing 2253 (`store_scores=False`). With activation checkpointing only each layer's 2254 16-bit input (2·s·b·h bytes) is kept, and the rest is recomputed. 2255 """ 2256 sbh = seq * batch * hidden 2257 if checkpointed: 2258 return 2 * sbh * layers 2259 per_layer = sbh * (34 + (5 * heads * seq / hidden if store_scores else 0)) 2260 return per_layer * layers 2261 2262 2263def zero_memory_per_gpu(params: float, n_gpus: int, stage: int, optimizer_bytes: int = 12) -> float: 2264 """Bytes of training state per GPU under ZeRO stage 0 (plain data parallel) to 3 (FSDP). 2265 2266 Stage 1 shards the optimizer state, stage 2 also the gradients, stage 3 also 2267 the weights. Whatever is sharded is divided by the number of GPUs. 2268 """ 2269 p, k, n = params, optimizer_bytes, n_gpus 2270 return {0: (2 + 2 + k) * p, 1: 2 * p + 2 * p + k * p / n, 2: 2 * p + (2 + k) * p / n, 3: (2 + 2 + k) * p / n}[stage] 2271 2272 2273# --------------------------------------------------------------------------- 2274# 3. Parallelism: data, tensor and pipeline 2275# --------------------------------------------------------------------------- 2276 2277 2278def linear_regression_gradient(X: np.ndarray, y: np.ndarray, w: np.ndarray) -> np.ndarray: 2279 """Gradient of the mean squared error mean((Xw − y)²) with respect to w: (2/n)·Xᵀ(Xw − y).""" 2280 return 2 / len(y) * X.T @ (X @ w - y) 2281 2282 2283def ring_all_reduce(arrays: list[np.ndarray]) -> tuple[list[np.ndarray], list[int]]: 2284 """Sum equal-length arrays across N simulated workers arranged in a ring. 2285 2286 Each worker cuts its array into N chunks. Reduce-scatter: for N − 1 steps, 2287 every worker passes one chunk to its right-hand neighbour, which adds it to 2288 its own copy; afterwards worker i owns the complete sum of one chunk. 2289 All-gather: for N − 1 more steps the finished chunks travel round the ring 2290 so everyone ends with every sum. Returns each worker's result and how many 2291 numbers each worker sent. 2292 """ 2293 n = len(arrays) 2294 chunks = [np.array_split(a.astype(float).copy(), n) for a in arrays] # chunks[worker][chunk] 2295 sent = [0] * n 2296 for step in range(n - 1): # reduce-scatter 2297 # Worker i sends chunk (i − step) to worker i + 1. Read everything first: all sends are simultaneous. 2298 outgoing = [(i, (i - step) % n, chunks[i][(i - step) % n].copy()) for i in range(n)] 2299 for i, c, data in outgoing: 2300 chunks[(i + 1) % n][c] += data 2301 sent[i] += data.size 2302 for step in range(n - 1): # all-gather: worker i now owns the full sum of chunk (i + 1) 2303 outgoing = [(i, (i + 1 - step) % n, chunks[i][(i + 1 - step) % n].copy()) for i in range(n)] 2304 for i, c, data in outgoing: 2305 chunks[(i + 1) % n][c] = data 2306 sent[i] += data.size 2307 return [np.concatenate(c) for c in chunks], sent 2308 2309 2310def all_reduce_traffic(size: float, n: int) -> float: 2311 """Data each worker sends in a ring all-reduce of `size`: 2·(N − 1)/N · size, almost flat in N.""" 2312 return 2 * (n - 1) / n * size 2313 2314 2315def column_parallel_matmul(X: np.ndarray, W: np.ndarray, n: int) -> np.ndarray: 2316 """X @ W with W's columns split across n devices; the pieces are placed side by side (an all-gather).""" 2317 return np.concatenate([X @ W_i for W_i in np.array_split(W, n, axis=1)], axis=1) 2318 2319 2320def row_parallel_matmul(X: np.ndarray, W: np.ndarray, n: int) -> np.ndarray: 2321 """X @ W with W's rows (and X's matching columns) split across n devices; partial results are summed 2322 (an all-reduce).""" 2323 parts = zip(np.array_split(X, n, axis=1), np.array_split(W, n, axis=0)) 2324 return sum(X_i @ W_i for X_i, W_i in parts) 2325 2326 2327def tensor_parallel_mlp(X: np.ndarray, W1: np.ndarray, W2: np.ndarray, n: int) -> np.ndarray: 2328 """ReLU(X W1) W2 as Megatron-LM splits it: W1 by columns, W2 by rows, one all-reduce at the end. 2329 2330 Splitting W1 by columns gives each device whole hidden units, so the 2331 elementwise ReLU can run locally with no communication at all. 2332 """ 2333 W1s, W2s = np.array_split(W1, n, axis=1), np.array_split(W2, n, axis=0) 2334 partials = [np.maximum(X @ a, 0) @ b for a, b in zip(W1s, W2s)] # each device: (tokens, d_model) 2335 return sum(partials) # the one all-reduce 2336 2337 2338def pipeline_schedule(stages: int, microbatches: int) -> np.ndarray: 2339 """GPipe's fill-and-drain schedule as a (stages, time) grid. 2340 2341 Cell value j > 0: forward pass of micro-batch j; −j: its backward pass; 2342 0: the stage is idle (the bubble). Forward passes flow down the stages, 2343 then backward passes flow back up. 2344 """ 2345 p, m = stages, microbatches 2346 span = m + p - 1 2347 grid = np.zeros((p, 2 * span), dtype=int) 2348 for s in range(p): 2349 for j in range(m): 2350 grid[s, s + j] = j + 1 # stage s can start micro-batch j once stage s−1 has finished it 2351 grid[s, span + (p - 1 - s) + j] = -(j + 1) # backward starts at the last stage 2352 return grid 2353 2354 2355def bubble_fraction(stages: int, microbatches: int) -> float: 2356 """Share of the schedule each stage spends idle: (p − 1) / (m + p − 1).""" 2357 return (stages - 1) / (microbatches + stages - 1) 2358 2359 2360# --------------------------------------------------------------------------- 2361# 4. Mixed precision: number formats, loss scaling, master weights 2362# --------------------------------------------------------------------------- 2363 2364 2365@dataclass(frozen=True) 2366class FloatFormat: 2367 """A binary floating-point format: 1 sign bit, `exp_bits` exponent bits, `man_bits` fraction bits.""" 2368 2369 name: str 2370 exp_bits: int 2371 man_bits: int 2372 max_override: float | None = None # E4M3 spends its top code on NaN only, so it reaches 448 not 240 2373 2374 @property 2375 def bias(self) -> int: 2376 return 2 ** (self.exp_bits - 1) - 1 2377 2378 @property 2379 def max_value(self) -> float: 2380 if self.max_override is not None: 2381 return self.max_override 2382 return (2 - 2.0**-self.man_bits) * 2.0**self.bias 2383 2384 @property 2385 def min_normal(self) -> float: 2386 return 2.0 ** (1 - self.bias) 2387 2388 @property 2389 def min_subnormal(self) -> float: 2390 return 2.0 ** (1 - self.bias - self.man_bits) 2391 2392 @property 2393 def epsilon(self) -> float: 2394 """Gap between 1 and the next representable number: the format's relative precision.""" 2395 return 2.0**-self.man_bits 2396 2397 2398FP32 = FloatFormat("fp32", 8, 23) 2399BF16 = FloatFormat("bf16", 8, 7) 2400FP16 = FloatFormat("fp16", 5, 10) 2401FP8_E4M3 = FloatFormat("fp8 E4M3", 4, 3, max_override=448.0) 2402FP8_E5M2 = FloatFormat("fp8 E5M2", 5, 2) 2403FORMATS = [FP32, BF16, FP16, FP8_E5M2, FP8_E4M3] 2404 2405 2406def quantize(x, fmt: FloatFormat): 2407 """Round x to the nearest value `fmt` can hold (ties to even), as the hardware does. 2408 2409 Between 2ᵉ and 2ᵉ⁺¹ a format has 2^man_bits evenly spaced values, so the 2410 spacing is 2^(e − man_bits). Below the smallest normal number the spacing 2411 stops shrinking (subnormals), and anything under half the smallest 2412 subnormal rounds to 0: that is underflow. Above the largest value is 2413 overflow, shown here as infinity. 2414 """ 2415 x = np.asarray(x, dtype=np.float64) 2416 mag = np.abs(x) 2417 e = np.floor(np.log2(np.where(mag > 0, mag, 1.0))) 2418 e = np.maximum(e, 1 - fmt.bias) # subnormals share the smallest normal exponent's spacing 2419 step = 2.0 ** (e - fmt.man_bits) 2420 q = np.round(mag / step) * step # np.round rounds halves to even, like IEEE hardware 2421 q = np.where(q > fmt.max_value, np.inf, q) 2422 out = np.sign(x) * q 2423 return float(out) if out.ndim == 0 else out 2424 2425 2426def scaled_gradient_roundtrip(grad: float, scale: float, fmt: FloatFormat) -> float: 2427 """Multiply by the loss scale, store in `fmt` (where the backward pass happens), divide back in fp32.""" 2428 return quantize(grad * scale, fmt) / scale 2429 2430 2431@dataclass 2432class DynamicLossScaler: 2433 """Keep the loss scale as large as possible without overflowing. 2434 2435 Overflow (an inf or NaN gradient) means the scale is too big: skip this 2436 step and halve it. A long run of clean steps means there may be headroom: 2437 double it. 2438 """ 2439 2440 scale: float = 2.0**16 2441 growth_interval: int = 2000 2442 clean_steps: int = 0 2443 2444 def update(self, grads_finite: bool) -> bool: 2445 """Record one step's outcome; return True if the optimizer should apply this step.""" 2446 if not grads_finite: 2447 self.scale /= 2 2448 self.clean_steps = 0 2449 return False 2450 self.clean_steps += 1 2451 if self.clean_steps == self.growth_interval: 2452 self.scale *= 2 2453 self.clean_steps = 0 2454 return True 2455 2456 2457def accumulate_updates(w0: float, update: float, steps: int, fmt: FloatFormat) -> float: 2458 """Add `update` to a weight `steps` times, rounding the weight to `fmt` after every step.""" 2459 w = quantize(w0, fmt) 2460 for _ in range(steps): 2461 w = quantize(w + update, fmt) 2462 return w 2463 2464 2465# --------------------------------------------------------------------------- 2466# 5. Stability: clipping, spikes, checkpoints 2467# --------------------------------------------------------------------------- 2468 2469 2470def sharded_global_norm(shards: list[np.ndarray]) -> float: 2471 """‖g‖ when g is split across GPUs: each sums its own squares, one all-reduce adds the sums, then √.""" 2472 local = [float(np.sum(s**2)) for s in shards] # one number per GPU 2473 return math.sqrt(sum(local)) # the all-reduce of those numbers, then the square root 2474 2475 2476def train_through_a_bad_batch(clip: float | None, steps: int = 60, bad_step: int = 30, lr: float = 0.1, 2477 seed: int = 0) -> list[float]: 2478 """SGD on y = 3x, with one corrupted batch (labels flipped and blown up 100×) at `bad_step`. 2479 2480 Returns the loss on clean held-out data after every step. With `clip`, each 2481 gradient is scaled down to at most that length before the update. 2482 """ 2483 rng = np.random.default_rng(seed) 2484 x_eval = rng.standard_normal(256) 2485 w = np.array([2.5]) 2486 losses = [] 2487 for t in range(steps): 2488 x = rng.standard_normal(32) 2489 y = -300 * x if t == bad_step else 3 * x 2490 g = linear_regression_gradient(x[:, None], y, w) 2491 if clip is not None: 2492 (g,), _ = clip_by_global_norm([g], clip) 2493 w = w - lr * g 2494 losses.append(float(np.mean((w[0] * x_eval - 3 * x_eval) ** 2))) 2495 return losses 2496 2497 2498def detect_spikes(losses: list[float], window: int = 50, factor: float = 2.0) -> list[int]: 2499 """Steps whose loss exceeds `factor` × the median of the previous `window` losses. 2500 2501 The median, not the mean, so one spike does not raise the bar for the next. 2502 """ 2503 return [i for i in range(window, len(losses)) if losses[i] > factor * float(np.median(losses[i - window : i]))] 2504 2505 2506def wasted_fraction(interval: float, cost: float, mtbf: float) -> float: 2507 """Share of time lost to checkpointing: saving (cost/interval) plus redoing lost work (interval/(2·mtbf)).""" 2508 return cost / interval + interval / (2 * mtbf) 2509 2510 2511def optimal_checkpoint_interval(cost: float, mtbf: float) -> float: 2512 """Young's (1974) interval √(2·cost·mtbf), which minimises `wasted_fraction`.""" 2513 return math.sqrt(2 * cost * mtbf) 2514 2515 2516# --------------------------------------------------------------------------- 2517# 6. Figures (rendered into the HTML docs by `make figures`) 2518# --------------------------------------------------------------------------- 2519 2520 2521def figures() -> dict: 2522 """Plot this lesson's data. matplotlib is imported here, and only here, 2523 so the lesson itself needs nothing beyond NumPy.""" 2524 import matplotlib 2525 2526 matplotlib.use("Agg") 2527 import matplotlib.pyplot as plt 2528 2529 BLUE, RED, GREEN, AMBER, MUTED = "#2563eb", "#dc2626", "#059669", "#d97706", "#9ca3af" 2530 figs = {} 2531 2532 # --- 1. MinHash estimates converge on the true Jaccard ------------------- 2533 fig, ax = plt.subplots(figsize=(6, 3.4)) 2534 ks = np.unique(np.logspace(0, np.log10(256), 40).astype(int)) 2535 hasher = MinHasher(256, seed=1) 2536 for true_j, color in ((0.2, BLUE), (0.5, GREEN), (0.8, RED)): 2537 # |A| = |B| = 100 with overlap o gives J = o / (200 − o), so o = 200·J / (1 + J). 2538 o = round(200 * true_j / (1 + true_j)) 2539 a = {f"s{i}" for i in range(100)} 2540 b = {f"s{i}" for i in range(100 - o, 200 - o)} 2541 sa, sb = hasher.signature(a), hasher.signature(b) 2542 ax.plot(ks, [estimate_jaccard(sa[:k], sb[:k]) for k in ks], color=color, label=f"true J = {jaccard(a, b):.2f}") 2543 ax.axhline(jaccard(a, b), color=color, ls="--", lw=1) 2544 ax.set_xscale("log", base=2) 2545 ax.set_xlabel("number of hash functions k") 2546 ax.set_ylabel("MinHash estimate of J") 2547 ax.set_ylim(-0.05, 1.05) 2548 ax.set_title("MinHash: more hashes, tighter estimate") 2549 ax.legend(frameon=False, loc="center right", bbox_to_anchor=(1.0, 0.64)) 2550 figs["minhash_convergence"] = fig 2551 2552 # --- 2. LSH S-curves ----------------------------------------------------- 2553 fig, ax = plt.subplots(figsize=(6, 3.4)) 2554 s = np.linspace(0, 1, 201) 2555 for (b, r), color in (((50, 2), BLUE), ((20, 5), GREEN), ((10, 10), RED)): 2556 ax.plot(s, lsh_candidate_probability(s, b, r), color=color, label=f"{b} bands × {r} rows") 2557 ax.axvline((1 / b) ** (1 / r), color=color, ls="--", lw=1) 2558 ax.set_xlabel("true Jaccard similarity of a pair") 2559 ax.set_ylabel("chance the pair is compared") 2560 ax.set_title("LSH banding: an S-curve you can place") 2561 ax.legend(frameon=False, loc="lower right") 2562 figs["lsh_s_curve"] = fig 2563 2564 # --- 3. Model collapse --------------------------------------------------- 2565 fig, ax = plt.subplots(figsize=(6, 3.4)) 2566 gens = np.arange(1, 201) 2567 for seed, color in ((0, BLUE), (1, GREEN), (2, AMBER)): 2568 ax.semilogy(gens, recursive_gaussian_fit(200, 20, seed), color=color, lw=1.2, label=f"seed {seed}") 2569 ax.semilogy(gens, np.sqrt((19 / 20) ** gens), color="#4b5563", ls="--", label="average shrink √(0.95ᵗ)") 2570 ax.set_xlabel("generation (each fitted to 20 samples of the last)") 2571 ax.set_ylabel("fitted spread σ") 2572 ax.set_title("Training on your own samples: the spread collapses") 2573 ax.legend(frameon=False, loc="lower left") 2574 figs["model_collapse"] = fig 2575 2576 # --- 4. The 7B memory bill ------------------------------------------------ 2577 fig, ax = plt.subplots(figsize=(6.4, 3.6)) 2578 state = training_memory(7e9) 2579 acts = { 2580 "naive": activation_bytes(4096, 1, 4096, 32, 32), 2581 "no stored scores": activation_bytes(4096, 1, 4096, 32, 32, store_scores=False), 2582 "checkpointed": activation_bytes(4096, 1, 4096, 32, 32, checkpointed=True), 2583 } 2584 colors = [BLUE, "#60a5fa", GREEN, "#34d399", "#a7f3d0"] 2585 x = np.arange(len(acts)) 2586 bottom = np.zeros(len(acts)) 2587 for (name, val), color in zip([(k, v) for k, v in state.items() if k != "total"], colors): 2588 ax.bar(x, val / 1e9, bottom=bottom, color=color, label=name, width=0.55) 2589 bottom += val / 1e9 2590 ax.bar(x, [a / 1e9 for a in acts.values()], bottom=bottom, color=AMBER, label="activations", width=0.55) 2591 for xi, a in zip(x, acts.values()): 2592 ax.text(xi, bottom[xi] + a / 1e9 + 3, f"{(state['total'] + a) / 1e9:.0f} GB", ha="center") 2593 ax.axhline(80, color=RED, ls="--", label="one 80 GB GPU") 2594 ax.set_xticks(x, list(acts)) 2595 ax.set_ylabel("GB on one GPU") 2596 ax.set_ylim(0, 245) 2597 ax.set_title("Training a 7B model on one GPU: it does not fit") 2598 ax.legend(frameon=False, fontsize=8, loc="upper right") 2599 figs["memory_7b"] = fig 2600 2601 # --- 5. ZeRO stages ------------------------------------------------------- 2602 fig, ax = plt.subplots(figsize=(6, 3.6)) 2603 ns = 2 ** np.arange(0, 11) 2604 for stage, color in zip(range(4), (MUTED, BLUE, GREEN, RED)): 2605 ax.loglog(ns, [zero_memory_per_gpu(7e9, int(n), stage) / 1e9 for n in ns], "o-", color=color, ms=3, 2606 label=["stage 0: plain data parallel", "stage 1: + shard optimizer", "stage 2: + shard gradients", 2607 "stage 3 (FSDP): + shard weights"][stage]) 2608 ax.axhline(80, color="#4b5563", ls="--") 2609 ax.text(2**7, 95, "80 GB GPU", color="#4b5563", zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1)) 2610 ax.set_xscale("log", base=2) 2611 ax.set_xlabel("GPUs sharing the state") 2612 ax.set_ylabel("training state per GPU (GB)") 2613 ax.set_title("ZeRO: sharding divides the 16 bytes per parameter") 2614 ax.legend(frameon=False, fontsize=8, loc="lower left") 2615 figs["zero_stages"] = fig 2616 2617 # --- 6. The pipeline schedule -------------------------------------------- 2618 from matplotlib.colors import ListedColormap 2619 2620 p, m = 4, 8 2621 grid = pipeline_schedule(p, m) 2622 fig, ax = plt.subplots(figsize=(8, 2.6)) 2623 kind = np.sign(grid) + 1 # 0 backward, 1 idle, 2 forward 2624 ax.imshow(kind, cmap=ListedColormap(["#86efac", "#ffffff", "#93c5fd"]), aspect="auto", vmin=0, vmax=2) 2625 for (row, col), v in np.ndenumerate(grid): 2626 if v: 2627 ax.text(col, row, str(abs(v)), ha="center", va="center", fontsize=8) 2628 ax.set_xticks(np.arange(-0.5, grid.shape[1], 1), minor=True) 2629 ax.set_yticks(np.arange(-0.5, p, 1), minor=True) 2630 ax.grid(which="minor", color="#d1d5db", lw=0.6) 2631 ax.grid(which="major", visible=False) 2632 ax.tick_params(which="minor", length=0) 2633 ax.set_yticks(range(p), [f"GPU {i + 1}" for i in range(p)]) 2634 ax.set_xticks(range(0, grid.shape[1], 2)) 2635 ax.set_xlabel("time step (blue: forward of micro-batch n, green: backward, white: idle)") 2636 ax.set_title(f"GPipe schedule, {p} stages, {m} micro-batches: bubble = {bubble_fraction(p, m):.0%}") 2637 figs["pipeline_schedule"] = fig 2638 2639 # --- 7. Bubble fraction --------------------------------------------------- 2640 fig, ax = plt.subplots(figsize=(6, 3.4)) 2641 ms = np.arange(1, 257) 2642 for p, color in zip((2, 4, 8, 16), (BLUE, GREEN, AMBER, RED)): 2643 ax.semilogx(ms, [bubble_fraction(p, int(m)) for m in ms], color=color, label=f"{p} stages") 2644 ax.axhline(0.1, color=MUTED, ls="--") 2645 ax.set_xlabel("micro-batches per batch") 2646 ax.set_ylabel("share of time idle") 2647 ax.set_title("The pipeline bubble: (p − 1) / (m + p − 1)") 2648 ax.legend(frameon=False) 2649 figs["bubble_fraction"] = fig 2650 2651 # --- 8. Float ranges ------------------------------------------------------ 2652 fig, ax = plt.subplots(figsize=(6.4, 3.0)) 2653 for i, fmt in enumerate(FORMATS): 2654 lo, mid, hi = np.log10(fmt.min_subnormal), np.log10(fmt.min_normal), np.log10(fmt.max_value) 2655 ax.barh(i, mid - lo, left=lo, color="#bfdbfe", height=0.6) 2656 ax.barh(i, hi - mid, left=mid, color=BLUE, height=0.6) 2657 ax.set_yticks(range(len(FORMATS)), [f.name for f in FORMATS]) 2658 ax.invert_yaxis() 2659 ax.set_xlabel("log₁₀ of magnitude (light: subnormals, dark: normal range)") 2660 ax.set_title("Range of each format: bf16 keeps all of fp32's") 2661 figs["float_ranges"] = fig 2662 2663 # --- 9. Loss scaling ------------------------------------------------------ 2664 rng = np.random.default_rng(0) 2665 grads = np.exp(rng.normal(np.log(2.0**-22), 3.0, 1_000_000)) # sizes spread over many powers of two 2666 fig, ax = plt.subplots(figsize=(6.4, 3.4)) 2667 bins = np.arange(-50, 20, 0.5) 2668 lost = float(np.mean(quantize(grads, FP16) == 0)) 2669 lost_scaled = float(np.mean(quantize(grads * 2.0**16, FP16) == 0)) 2670 ax.hist(np.log2(grads), bins=bins, color=MUTED, alpha=0.8, label=f"unscaled: {lost:.0%} become 0 in fp16") 2671 ax.hist(np.log2(grads * 2.0**16), bins=bins, color=BLUE, alpha=0.6, 2672 label=f"× 2¹⁶: {lost_scaled:.2%} become 0") 2673 top = ax.get_ylim()[1] * 1.5 # headroom above the peaks for the labels and legend 2674 ax.set_ylim(0, top) 2675 ax.axvspan(-50, np.log2(FP16.min_subnormal / 2), color=RED, alpha=0.12) 2676 ax.axvline(np.log2(FP16.max_value), color=RED, ls="--") 2677 ax.text(np.log2(FP16.max_value) - 0.5, top * 0.93, "fp16 overflow", color=RED, ha="right") 2678 ax.text(-49, top * 0.93, "flushed to 0", color=RED) 2679 ax.set_xlabel("gradient size, log₂") 2680 ax.set_ylabel("number of gradients") 2681 ax.set_title("Loss scaling slides the gradients into fp16's range") 2682 ax.legend(frameon=False, loc="upper center", bbox_to_anchor=(0.5, 0.88), fontsize=8) 2683 figs["loss_scaling"] = fig 2684 2685 # --- 10. Master weights --------------------------------------------------- 2686 fig, ax = plt.subplots(figsize=(6, 3.2)) 2687 for fmt, color, ls in ((FP32, BLUE, "-"), (BF16, RED, "-"), (FP16, AMBER, ":")): 2688 w, path = quantize(1.0, fmt), [] 2689 for _ in range(1000): 2690 w = quantize(w + 1e-4, fmt) 2691 path.append(w) 2692 ax.plot(range(1, 1001), path, color=color, ls=ls, lw=2, label=fmt.name) 2693 ax.set_xlabel("update step (each adds 0.0001)") 2694 ax.set_ylabel("weight value") 2695 ax.set_title("Tiny updates vanish in 16 bits; an fp32 master keeps them") 2696 ax.legend(frameon=False) 2697 figs["master_weights"] = fig 2698 2699 # --- 11. One bad batch, with and without clipping ------------------------ 2700 fig, ax = plt.subplots(figsize=(6, 3.4)) 2701 for clip, color, label in ((None, RED, "no clipping"), (1.0, BLUE, "clip global norm at 1")): 2702 losses = np.maximum(train_through_a_bad_batch(clip), 1e-12) 2703 ax.semilogy(losses, color=color, label=label) 2704 ax.axvline(30, color=MUTED, ls="--") 2705 ax.text(31, 1e-8, "corrupted batch", color="#4b5563") 2706 ax.set_xlabel("training step") 2707 ax.set_ylabel("clean held-out loss") 2708 ax.set_title("One bad batch: clipping limits the damage") 2709 ax.legend(frameon=False, loc="upper right") 2710 figs["bad_batch"] = fig 2711 2712 # --- 12. Checkpoint interval ---------------------------------------------- 2713 fig, ax = plt.subplots(figsize=(6, 3.4)) 2714 T = np.logspace(0, np.log10(300), 200) 2715 for cost, color, label in ((1.0, BLUE, "1-minute save"), (1 / 6, GREEN, "10-second save")): 2716 wasted = np.array([wasted_fraction(t, cost, 180) for t in T]) 2717 ax.semilogx(T, np.where(wasted <= 0.5, wasted, np.nan), color=color, label=label) # above 50% leaves the chart 2718 best = optimal_checkpoint_interval(cost, 180) 2719 ax.plot(best, wasted_fraction(best, cost, 180), "o", color=color) 2720 ax.annotate(f"{best:.1f} min, {wasted_fraction(best, cost, 180):.1%}", (best, wasted_fraction(best, cost, 180)), 2721 textcoords="offset points", xytext=(6, -14), color=color, zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1)) 2722 ax.set_xlabel("minutes between checkpoints (failure every 180 min)") 2723 ax.set_ylabel("share of time wasted") 2724 ax.set_ylim(0, 0.5) 2725 ax.set_title("Checkpoint interval: a valley at √(2·C·M)") 2726 ax.legend(frameon=False) 2727 figs["checkpoint_interval"] = fig 2728 2729 return figs 2730 2731 2732# --------------------------------------------------------------------------- 2733# 7. Narrated walkthrough 2734# --------------------------------------------------------------------------- 2735 2736 2737def demo() -> None: 2738 banner("1. Cleaning a crawl: language, quality, exact and near duplicates") 2739 say( 2740 """ 2741 Seven pages arrive from a crawl. Cheap per-page checks run first 2742 (language ID, quality rules); the comparisons between pages 2743 (exact hash, then MinHash) run last, on what survived. 2744 """ 2745 ) 2746 labels = ["river paragraph", "French version", "navigation bar", "re-crawled copy", "hashtag spam", 2747 "scraper's edited copy", "stars paragraph"] 2748 table(["page", "what it is", "verdict"], [(i, l, v) for i, (l, v) in enumerate(zip(labels, curate(CRAWL_SAMPLE)))]) 2749 clf = QualityClassifier.trained_on_examples() 2750 for text in ("the delta is formed when the river deposits sand over many years", 2751 "click here buy now best price free free free"): 2752 print(f" classifier score {clf.score(text):+7.2f} {text!r}") 2753 print() 2754 takeaway("Most of a crawl never reaches the model: two of seven pages survive here.") 2755 2756 banner("2. MinHash and LSH: near duplicates without comparing everything") 2757 a, b = CRAWL_SAMPLE[0], CRAWL_SAMPLE[5] 2758 true_j = jaccard(shingles(a), shingles(b)) 2759 rows = [] 2760 for k in (8, 32, 128, 512): 2761 h = MinHasher(k, seed=0) 2762 rows.append((k, true_j, estimate_jaccard(h.signature(shingles(a)), h.signature(shingles(b))))) 2763 table(["hashes k", "true Jaccard", "MinHash estimate"], rows, floatfmt=".3f") 2764 table(["similarity s", "P(candidate), 20 bands × 5 rows"], 2765 [(s, lsh_candidate_probability(s, 20, 5)) for s in (0.3, 0.5, 0.7, 0.8, 0.9)], floatfmt=".4f") 2766 pairs = near_duplicate_pairs([a, b, CRAWL_SAMPLE[6]]) 2767 say(f"LSH over pages 0, 5 and 6 (listed as 0, 1, 2) flags {pairs}: pages 0 and 5, and nothing else.") 2768 2769 banner("3. Why duplicates hurt: a bigram model memorizes what it sees often") 2770 table(["copies of the boilerplate line", "P(model continues 'click' into the whole line)"], 2771 [(c, verbatim_probability(c)) for c in (1, 10, 100, 1000)], floatfmt=".3f") 2772 takeaway("Repetition turns learning into recitation; deduplication is also privacy and eval hygiene.") 2773 2774 banner("4. Mixtures, budgets and synthetic data") 2775 epochs = epochs_per_source({"web": 0.80, "wiki": 0.05, "code": 0.15}, 2776 {"web": 900e9, "wiki": 20e9, "code": 150e9}, budget=1000e9) 2777 table(["source", "weight", "passes over it"], [(k, w, epochs[k]) for k, w in (("web", 0.80), ("wiki", 0.05), ("code", 0.15))], 2778 floatfmt=".2f") 2779 say(f"A 7B model's compute-optimal budget: {chinchilla_tokens(7e9) / 1e9:.0f}B tokens, " 2780 f"costing {training_flops(7e9, chinchilla_tokens(7e9)):.2e} FLOPs.") 2781 stds = recursive_gaussian_fit() 2782 say(f"Refit on your own samples 200 times: spread 1.0 -> {stds[49]:.3f} after 50 generations, " 2783 f"{stds[-1]:.1e} after 200. Keep real data in the mix.") 2784 2785 banner("5. The memory bill: why one GPU is not enough") 2786 mem = training_memory(7e9) 2787 table(["what", "GB for 7B"], [(k, v / 1e9) for k, v in mem.items()], floatfmt=".1f") 2788 table(["activations at 4,096 tokens", "GB"], [ 2789 ("stored naively", activation_bytes(4096, 1, 4096, 32, 32) / 1e9), 2790 ("without attention scores", activation_bytes(4096, 1, 4096, 32, 32, store_scores=False) / 1e9), 2791 ("activation checkpointing", activation_bytes(4096, 1, 4096, 32, 32, checkpointed=True) / 1e9), 2792 ], floatfmt=".1f") 2793 table(["ZeRO stage", "GB per GPU, 7.5B on 64 GPUs"], [(s, zero_memory_per_gpu(7.5e9, 64, s) / 1e9) for s in range(4)], 2794 floatfmt=".2f") 2795 2796 banner("6. Parallelism: all-reduce, tensor splits, the pipeline bubble") 2797 grads = [np.arange(8.0) * (i + 1) for i in range(4)] 2798 results, sent = ring_all_reduce(grads) 2799 say(f"Ring all-reduce over 4 workers: everyone ends with {results[0].tolist()}; " 2800 f"each sent {sent[0]} numbers, 1.5 times its own 8.") 2801 rng = np.random.default_rng(0) 2802 X, W1, W2 = rng.standard_normal((3, 8)), rng.standard_normal((8, 16)), rng.standard_normal((16, 8)) 2803 err = np.abs(tensor_parallel_mlp(X, W1, W2, n=4) - np.maximum(X @ W1, 0) @ W2).max() 2804 say(f"Feed-forward layer split over 4 devices (columns, then rows): largest difference from one device = {err:.1e}.") 2805 print(pipeline_schedule(4, 4)) 2806 print() 2807 table(["micro-batches (4 stages)", "bubble"], [(m, bubble_fraction(4, m)) for m in (1, 4, 8, 32)], floatfmt=".3f") 2808 2809 banner("7. Mixed precision: range, precision, loss scaling, master weights") 2810 table(["format", "largest", "smallest normal", "gap after 1"], 2811 [(f.name, f"{f.max_value:.3g}", f"{f.min_normal:.3g}", f"{f.epsilon:.3g}") for f in FORMATS]) 2812 say(f"A gradient of 1e-8 in fp16: {quantize(1e-8, FP16)}. Scaled by 65536, stored, unscaled: " 2813 f"{scaled_gradient_roundtrip(1e-8, 65536, FP16):.4e}. In bf16 with no scaling: {quantize(1e-8, BF16):.4e}.") 2814 scaler = DynamicLossScaler(scale=65536.0, growth_interval=3) 2815 events = [True, True, False, True, True, True] 2816 log = [] 2817 for ok in events: 2818 applied = scaler.update(ok) 2819 log.append(("clean" if ok else "overflow", "applied" if applied else "skipped", scaler.scale)) 2820 table(["gradients", "step", "scale after"], log, floatfmt=".0f") 2821 say(f"1,000 updates of 0.0001 to a weight of 1.0: fp32 ends at {accumulate_updates(1.0, 1e-4, 1000, FP32):.4f}, " 2822 f"bf16 at {accumulate_updates(1.0, 1e-4, 1000, BF16):.4f}.") 2823 takeaway("bf16 for the maths (fp32's range), fp32 for the master weights and optimizer state.") 2824 2825 banner("8. Stability: clipping, spikes, checkpoints") 2826 no_clip, clipped = train_through_a_bad_batch(None), train_through_a_bad_batch(1.0) 2827 say(f"One corrupted batch at step 30: peak loss {max(no_clip):,.0f} without clipping, " 2828 f"{max(clipped[30:]):.4f} with clipping.") 2829 say(f"Spikes flagged in [3.0, 2.9, 2.8, 2.8, 2.7, 8.5, 4.0, 2.7]: steps " 2830 f"{detect_spikes([3.0, 2.9, 2.8, 2.8, 2.7, 8.5, 4.0, 2.7], window=4, factor=2.0)}.") 2831 table(["save time (min)", "best interval (min)", "time wasted"], 2832 [(c, optimal_checkpoint_interval(c, 180), wasted_fraction(optimal_checkpoint_interval(c, 180), c, 180)) 2833 for c in (5.0, 1.0, 1 / 6)], floatfmt=".3f") 2834 takeaway("At thousands of GPUs, failures are the weather: clip, watch, roll back, and save often and fast.") 2835 2836 2837if __name__ == "__main__": 2838 demo()
1919def words(text: str) -> list[str]: 1920 """Lowercased words: the unit every filter below counts.""" 1921 return _WORD.findall(text.lower())
Lowercased words: the unit every filter below counts.
1924def detect_language(text: str, min_share: float = 0.1) -> str: 1925 """The language whose stop words make up the largest share of the text, or "unknown". 1926 1927 A page of product codes or numbers matches no language's little words at 1928 all, so it falls under `min_share` and is reported as unknown. 1929 """ 1930 ws = words(text) 1931 if not ws: 1932 return "unknown" 1933 shares = {lang: sum(w in sw for w in ws) / len(ws) for lang, sw in STOP_WORDS.items()} 1934 best = max(shares, key=shares.get) 1935 return best if shares[best] >= min_share else "unknown"
The language whose stop words make up the largest share of the text, or "unknown".
A page of product codes or numbers matches no language's little words at
all, so it falls under min_share and is reported as unknown.
1942def quality_failures(text: str) -> list[str]: 1943 """Names of the heuristic quality rules a document breaks (empty list = keep it). 1944 1945 Each rule is cheap, obvious once stated, and removes a whole genre of junk: 1946 navigation snippets, hashtag spam, keyword stuffing, lists of links. 1947 """ 1948 raw = text.split() 1949 n = len(raw) 1950 lines = [l for l in text.splitlines() if l.strip()] or [text] 1951 failures = [] 1952 if n < 50: 1953 failures.append("too_short") 1954 if n > 100_000: 1955 failures.append("too_long") 1956 if n and not 3 <= sum(len(w) for w in raw) / n <= 10: 1957 failures.append("odd_word_length") 1958 symbols = text.count("#") + text.count("...") + text.count("…") 1959 if n and symbols / n > 0.1: 1960 failures.append("too_many_symbols") 1961 if sum(l.lstrip().startswith(("-", "*", "•")) for l in lines) / len(lines) > 0.9: 1962 failures.append("mostly_bullets") 1963 if sum(l.rstrip().endswith(("...", "…")) for l in lines) / len(lines) > 0.3: 1964 failures.append("mostly_ellipsis") 1965 if n and sum(any(c.isalpha() for c in w) for w in raw) / n < 0.8: 1966 failures.append("too_few_alphabetic_words") 1967 if len(GOPHER_STOP_WORDS & set(words(text))) < 2: 1968 failures.append("missing_stop_words") 1969 return failures
Names of the heuristic quality rules a document breaks (empty list = keep it).
Each rule is cheap, obvious once stated, and removes a whole genre of junk: navigation snippets, hashtag spam, keyword stuffing, lists of links.
1992@dataclass 1993class QualityClassifier: 1994 """A bag-of-words naive Bayes scorer: log-odds that a text looks like the reference set. 1995 1996 GPT-3 filtered Common Crawl with a linear classifier of this kind (reference 1997 text vs. raw crawl); FineWeb-Edu used an LLM's educational-value ratings to 1998 train one. The score is a sum of per-word votes, so it is fast enough to 1999 run over billions of pages. 2000 """ 2001 2002 good: Counter = field(default_factory=Counter) 2003 bad: Counter = field(default_factory=Counter) 2004 2005 @classmethod 2006 def trained_on_examples(cls, good=GOOD_EXAMPLES, bad=BAD_EXAMPLES) -> "QualityClassifier": 2007 clf = cls() 2008 for t in good: 2009 clf.good.update(words(t)) 2010 for t in bad: 2011 clf.bad.update(words(t)) 2012 return clf 2013 2014 def word_vote(self, w: str) -> float: 2015 """log P(w | good) − log P(w | bad), with add-one smoothing so unseen words never give log 0.""" 2016 vocab = len(set(self.good) | set(self.bad)) 2017 p_good = (self.good[w] + 1) / (sum(self.good.values()) + vocab) 2018 p_bad = (self.bad[w] + 1) / (sum(self.bad.values()) + vocab) 2019 return math.log(p_good / p_bad) 2020 2021 def score(self, text: str) -> float: 2022 """Sum of word votes over words the classifier has seen. Above 0 leans reference, below 0 leans spam.""" 2023 known = set(self.good) | set(self.bad) 2024 return sum(self.word_vote(w) for w in words(text) if w in known)
A bag-of-words naive Bayes scorer: log-odds that a text looks like the reference set.
GPT-3 filtered Common Crawl with a linear classifier of this kind (reference text vs. raw crawl); FineWeb-Edu used an LLM's educational-value ratings to train one. The score is a sum of per-word votes, so it is fast enough to run over billions of pages.
2014 def word_vote(self, w: str) -> float: 2015 """log P(w | good) − log P(w | bad), with add-one smoothing so unseen words never give log 0.""" 2016 vocab = len(set(self.good) | set(self.bad)) 2017 p_good = (self.good[w] + 1) / (sum(self.good.values()) + vocab) 2018 p_bad = (self.bad[w] + 1) / (sum(self.bad.values()) + vocab) 2019 return math.log(p_good / p_bad)
log P(w | good) − log P(w | bad), with add-one smoothing so unseen words never give log 0.
2021 def score(self, text: str) -> float: 2022 """Sum of word votes over words the classifier has seen. Above 0 leans reference, below 0 leans spam.""" 2023 known = set(self.good) | set(self.bad) 2024 return sum(self.word_vote(w) for w in words(text) if w in known)
Sum of word votes over words the classifier has seen. Above 0 leans reference, below 0 leans spam.
2027def normalize_for_exact(text: str) -> str: 2028 """Lowercase and collapse whitespace, so trivially reformatted copies hash the same.""" 2029 return " ".join(text.lower().split())
Lowercase and collapse whitespace, so trivially reformatted copies hash the same.
2032def exact_dedup(docs: list[str]) -> list[str]: 2033 """Keep the first copy of each document, comparing a hash of its normalized text.""" 2034 seen, kept = set(), [] 2035 for d in docs: 2036 key = hashlib.sha1(normalize_for_exact(d).encode()).hexdigest() 2037 if key not in seen: 2038 seen.add(key) 2039 kept.append(d) 2040 return kept
Keep the first copy of each document, comparing a hash of its normalized text.
2043def shingles(text: str, k: int = 5) -> set[str]: 2044 """The set of overlapping k-word windows ("shingles") in a text.""" 2045 ws = words(text) 2046 return {" ".join(ws[i : i + k]) for i in range(max(1, len(ws) - k + 1))}
The set of overlapping k-word windows ("shingles") in a text.
2049def jaccard(a: set, b: set) -> float: 2050 """|A ∩ B| / |A ∪ B|: shared shingles over all distinct shingles.""" 2051 return len(a & b) / len(a | b) if a | b else 1.0
|A ∩ B| / |A ∪ B|: shared shingles over all distinct shingles.
2057class MinHasher: 2058 """k random hash functions; a document's signature is its smallest hash under each. 2059 2060 Two documents agree in one slot with probability exactly their Jaccard 2061 similarity, so the fraction of agreeing slots estimates it, from k numbers 2062 per document instead of the whole shingle set. 2063 """ 2064 2065 def __init__(self, num_perm: int = 128, seed: int = 0): 2066 rng = np.random.default_rng(seed) 2067 # h_i(x) = (a_i·x + b_i) mod p: a cheap family that behaves like random orderings. 2068 self.a = rng.integers(1, _MERSENNE, num_perm, dtype=np.int64) 2069 self.b = rng.integers(0, _MERSENNE, num_perm, dtype=np.int64) 2070 2071 def signature(self, shingle_set: set[str]) -> np.ndarray: 2072 # crc32 rather than hash(): Python salts hash() per process, which would break determinism. 2073 x = np.array([zlib.crc32(s.encode()) % _MERSENNE for s in shingle_set], dtype=np.int64) 2074 if x.size == 0: 2075 return np.full(self.a.shape, _MERSENNE, dtype=np.int64) 2076 return ((self.a[:, None] * x[None, :] + self.b[:, None]) % _MERSENNE).min(axis=1) # (num_perm,)
k random hash functions; a document's signature is its smallest hash under each.
Two documents agree in one slot with probability exactly their Jaccard similarity, so the fraction of agreeing slots estimates it, from k numbers per document instead of the whole shingle set.
2065 def __init__(self, num_perm: int = 128, seed: int = 0): 2066 rng = np.random.default_rng(seed) 2067 # h_i(x) = (a_i·x + b_i) mod p: a cheap family that behaves like random orderings. 2068 self.a = rng.integers(1, _MERSENNE, num_perm, dtype=np.int64) 2069 self.b = rng.integers(0, _MERSENNE, num_perm, dtype=np.int64)
2071 def signature(self, shingle_set: set[str]) -> np.ndarray: 2072 # crc32 rather than hash(): Python salts hash() per process, which would break determinism. 2073 x = np.array([zlib.crc32(s.encode()) % _MERSENNE for s in shingle_set], dtype=np.int64) 2074 if x.size == 0: 2075 return np.full(self.a.shape, _MERSENNE, dtype=np.int64) 2076 return ((self.a[:, None] * x[None, :] + self.b[:, None]) % _MERSENNE).min(axis=1) # (num_perm,)
2079def estimate_jaccard(sig_a: np.ndarray, sig_b: np.ndarray) -> float: 2080 """Fraction of signature slots where the two minimums agree.""" 2081 return float(np.mean(sig_a == sig_b))
Fraction of signature slots where the two minimums agree.
2084def lsh_candidate_probability(s: float, bands: int, rows: int) -> float: 2085 """Chance that a pair with Jaccard s shares at least one whole band: 1 − (1 − sʳ)ᵇ.""" 2086 return 1 - (1 - s**rows) ** bands
Chance that a pair with Jaccard s shares at least one whole band: 1 − (1 − sʳ)ᵇ.
2089def near_duplicate_pairs(docs: list[str], bands: int = 20, rows: int = 5, threshold: float = 0.7, seed: int = 0) -> list[tuple[int, int]]: 2090 """Pairs (i, j), i < j, whose MinHash similarity is at least `threshold`. 2091 2092 Instead of comparing all n² pairs, each signature is cut into `bands` bands 2093 of `rows` numbers and each band is hashed into a bucket. Only documents that 2094 land in the same bucket for some band are compared at all. 2095 """ 2096 hasher = MinHasher(bands * rows, seed) 2097 sigs = [hasher.signature(shingles(d)) for d in docs] 2098 buckets: dict[tuple, list[int]] = {} 2099 for i, sig in enumerate(sigs): 2100 for band in range(bands): 2101 buckets.setdefault((band, *sig[band * rows : (band + 1) * rows].tolist()), []).append(i) 2102 candidates = {(i, j) for ids in buckets.values() for i in ids for j in ids if i < j} 2103 return sorted((i, j) for i, j in candidates if estimate_jaccard(sigs[i], sigs[j]) >= threshold)
Pairs (i, j), i < j, whose MinHash similarity is at least threshold.
Instead of comparing all n² pairs, each signature is cut into bands bands
of rows numbers and each band is hashed into a bucket. Only documents that
land in the same bucket for some band are compared at all.
2138def curate(docs: list[str], language: str = "en") -> list[str]: 2139 """Run the pipeline in order (language, quality, exact dedup, near dedup); one verdict per doc. 2140 2141 Cheap checks run first so the expensive ones see fewer documents: language 2142 ID and the heuristics are a pass over each page, while near-duplicate 2143 detection compares pages with each other. 2144 """ 2145 verdicts: list[str] = [] 2146 kept_hashes: dict[str, int] = {} 2147 kept: list[int] = [] 2148 hasher = MinHasher(100, seed=0) 2149 kept_sigs: dict[int, np.ndarray] = {} 2150 for i, d in enumerate(docs): 2151 lang = detect_language(d) 2152 if lang != language: 2153 verdicts.append(f"language: {lang}") 2154 continue 2155 failures = quality_failures(d) 2156 if failures: 2157 verdicts.append("quality: " + ", ".join(failures)) 2158 continue 2159 key = hashlib.sha1(normalize_for_exact(d).encode()).hexdigest() 2160 if key in kept_hashes: 2161 verdicts.append(f"exact duplicate of {kept_hashes[key]}") 2162 continue 2163 sig = hasher.signature(shingles(d)) 2164 twin = next((j for j in kept if estimate_jaccard(sig, kept_sigs[j]) >= 0.7), None) 2165 if twin is not None: 2166 verdicts.append(f"near duplicate of {twin}") 2167 continue 2168 kept_hashes[key], kept_sigs[i] = i, sig 2169 kept.append(i) 2170 verdicts.append("kept") 2171 return verdicts
Run the pipeline in order (language, quality, exact dedup, near dedup); one verdict per doc.
Cheap checks run first so the expensive ones see fewer documents: language ID and the heuristics are a pass over each page, while near-duplicate detection compares pages with each other.
2184def verbatim_probability(copies: int, line: str = BOILERPLATE) -> float: 2185 """Probability a bigram model, trained on the small corpus plus `copies` of the line, continues the line's 2186 first word into the whole line. 2187 2188 A bigram model predicts each word from the one before it, by counting. 2189 Every duplicate adds the same counts again, so the line's path through the 2190 counts becomes the only road: that is memorization in its simplest form. 2191 """ 2192 corpus = OTHER_LINES + [line] * copies 2193 pairs = Counter((a, b) for s in corpus for a, b in zip(s.split(), s.split()[1:])) 2194 firsts = Counter(a for (a, _), c in pairs.items() for _ in range(c)) 2195 ws = line.split() 2196 return float(np.prod([pairs[(a, b)] / firsts[a] for a, b in zip(ws, ws[1:])]))
Probability a bigram model, trained on the small corpus plus copies of the line, continues the line's
first word into the whole line.
A bigram model predicts each word from the one before it, by counting. Every duplicate adds the same counts again, so the line's path through the counts becomes the only road: that is memorization in its simplest form.
2199def epochs_per_source(weights: dict[str, float], sizes: dict[str, float], budget: float) -> dict[str, float]: 2200 """How many passes over each source a mixture implies: weight × budget / tokens available.""" 2201 return {k: weights[k] * budget / sizes[k] for k in weights}
How many passes over each source a mixture implies: weight × budget / tokens available.
2204def chinchilla_tokens(params: float) -> float: 2205 """Compute-optimal training tokens for a model size: about 20 per parameter (Hoffmann et al., 2022).""" 2206 return 20 * params
Compute-optimal training tokens for a model size: about 20 per parameter (Hoffmann et al., 2022).
2209def training_flops(params: float, tokens: float) -> float: 2210 """C ≈ 6·N·D: 2 FLOPs per parameter per token forward, 4 backward.""" 2211 return 6 * params * tokens
C ≈ 6·N·D: 2 FLOPs per parameter per token forward, 4 backward.
2214def recursive_gaussian_fit(generations: int = 200, n: int = 20, seed: int = 0) -> list[float]: 2215 """Model collapse in miniature: fit a bell curve, sample from the fit, refit on the samples, repeat. 2216 2217 Returns the fitted standard deviation of every generation. Each refit on a 2218 finite sample loses a little of the tails, and the losses compound. 2219 """ 2220 rng = np.random.default_rng(seed) 2221 mu, sigma = 0.0, 1.0 2222 stds = [] 2223 for _ in range(generations): 2224 sample = rng.normal(mu, sigma, n) 2225 mu, sigma = float(sample.mean()), float(sample.std()) # maximum-likelihood fit (divides by n) 2226 stds.append(sigma) 2227 return stds
Model collapse in miniature: fit a bell curve, sample from the fit, refit on the samples, repeat.
Returns the fitted standard deviation of every generation. Each refit on a finite sample loses a little of the tails, and the losses compound.
2235def training_memory(params: float) -> dict[str, float]: 2236 """Bytes of training state for Adam in mixed precision, the "16 bytes per parameter" rule.""" 2237 parts = { 2238 "weights (bf16)": 2 * params, 2239 "gradients (bf16)": 2 * params, 2240 "master weights (fp32)": 4 * params, 2241 "Adam momentum (fp32)": 4 * params, 2242 "Adam variance (fp32)": 4 * params, 2243 } 2244 return {**parts, "total": sum(parts.values())}
Bytes of training state for Adam in mixed precision, the "16 bytes per parameter" rule.
2247def activation_bytes(seq: int, batch: int, hidden: int, heads: int, layers: int, 2248 store_scores: bool = True, checkpointed: bool = False) -> float: 2249 """Activations saved for the backward pass, in bytes (Korthikanti et al., 2022, no parallelism). 2250 2251 Per layer: s·b·h·(34 + 5·a·s/h). The 34·s·b·h part is the layer's ordinary 2252 intermediate tensors; the 5·a·s² part is attention's score matrices, which 2253 FlashAttention-style kernels recompute instead of storing 2254 (`store_scores=False`). With activation checkpointing only each layer's 2255 16-bit input (2·s·b·h bytes) is kept, and the rest is recomputed. 2256 """ 2257 sbh = seq * batch * hidden 2258 if checkpointed: 2259 return 2 * sbh * layers 2260 per_layer = sbh * (34 + (5 * heads * seq / hidden if store_scores else 0)) 2261 return per_layer * layers
Activations saved for the backward pass, in bytes (Korthikanti et al., 2022, no parallelism).
Per layer: s·b·h·(34 + 5·a·s/h). The 34·s·b·h part is the layer's ordinary
intermediate tensors; the 5·a·s² part is attention's score matrices, which
FlashAttention-style kernels recompute instead of storing
(store_scores=False). With activation checkpointing only each layer's
16-bit input (2·s·b·h bytes) is kept, and the rest is recomputed.
2264def zero_memory_per_gpu(params: float, n_gpus: int, stage: int, optimizer_bytes: int = 12) -> float: 2265 """Bytes of training state per GPU under ZeRO stage 0 (plain data parallel) to 3 (FSDP). 2266 2267 Stage 1 shards the optimizer state, stage 2 also the gradients, stage 3 also 2268 the weights. Whatever is sharded is divided by the number of GPUs. 2269 """ 2270 p, k, n = params, optimizer_bytes, n_gpus 2271 return {0: (2 + 2 + k) * p, 1: 2 * p + 2 * p + k * p / n, 2: 2 * p + (2 + k) * p / n, 3: (2 + 2 + k) * p / n}[stage]
Bytes of training state per GPU under ZeRO stage 0 (plain data parallel) to 3 (FSDP).
Stage 1 shards the optimizer state, stage 2 also the gradients, stage 3 also the weights. Whatever is sharded is divided by the number of GPUs.
2279def linear_regression_gradient(X: np.ndarray, y: np.ndarray, w: np.ndarray) -> np.ndarray: 2280 """Gradient of the mean squared error mean((Xw − y)²) with respect to w: (2/n)·Xᵀ(Xw − y).""" 2281 return 2 / len(y) * X.T @ (X @ w - y)
Gradient of the mean squared error mean((Xw − y)²) with respect to w: (2/n)·Xᵀ(Xw − y).
2284def ring_all_reduce(arrays: list[np.ndarray]) -> tuple[list[np.ndarray], list[int]]: 2285 """Sum equal-length arrays across N simulated workers arranged in a ring. 2286 2287 Each worker cuts its array into N chunks. Reduce-scatter: for N − 1 steps, 2288 every worker passes one chunk to its right-hand neighbour, which adds it to 2289 its own copy; afterwards worker i owns the complete sum of one chunk. 2290 All-gather: for N − 1 more steps the finished chunks travel round the ring 2291 so everyone ends with every sum. Returns each worker's result and how many 2292 numbers each worker sent. 2293 """ 2294 n = len(arrays) 2295 chunks = [np.array_split(a.astype(float).copy(), n) for a in arrays] # chunks[worker][chunk] 2296 sent = [0] * n 2297 for step in range(n - 1): # reduce-scatter 2298 # Worker i sends chunk (i − step) to worker i + 1. Read everything first: all sends are simultaneous. 2299 outgoing = [(i, (i - step) % n, chunks[i][(i - step) % n].copy()) for i in range(n)] 2300 for i, c, data in outgoing: 2301 chunks[(i + 1) % n][c] += data 2302 sent[i] += data.size 2303 for step in range(n - 1): # all-gather: worker i now owns the full sum of chunk (i + 1) 2304 outgoing = [(i, (i + 1 - step) % n, chunks[i][(i + 1 - step) % n].copy()) for i in range(n)] 2305 for i, c, data in outgoing: 2306 chunks[(i + 1) % n][c] = data 2307 sent[i] += data.size 2308 return [np.concatenate(c) for c in chunks], sent
Sum equal-length arrays across N simulated workers arranged in a ring.
Each worker cuts its array into N chunks. Reduce-scatter: for N − 1 steps, every worker passes one chunk to its right-hand neighbour, which adds it to its own copy; afterwards worker i owns the complete sum of one chunk. All-gather: for N − 1 more steps the finished chunks travel round the ring so everyone ends with every sum. Returns each worker's result and how many numbers each worker sent.
2311def all_reduce_traffic(size: float, n: int) -> float: 2312 """Data each worker sends in a ring all-reduce of `size`: 2·(N − 1)/N · size, almost flat in N.""" 2313 return 2 * (n - 1) / n * size
Data each worker sends in a ring all-reduce of size: 2·(N − 1)/N · size, almost flat in N.
2316def column_parallel_matmul(X: np.ndarray, W: np.ndarray, n: int) -> np.ndarray: 2317 """X @ W with W's columns split across n devices; the pieces are placed side by side (an all-gather).""" 2318 return np.concatenate([X @ W_i for W_i in np.array_split(W, n, axis=1)], axis=1)
X @ W with W's columns split across n devices; the pieces are placed side by side (an all-gather).
2321def row_parallel_matmul(X: np.ndarray, W: np.ndarray, n: int) -> np.ndarray: 2322 """X @ W with W's rows (and X's matching columns) split across n devices; partial results are summed 2323 (an all-reduce).""" 2324 parts = zip(np.array_split(X, n, axis=1), np.array_split(W, n, axis=0)) 2325 return sum(X_i @ W_i for X_i, W_i in parts)
X @ W with W's rows (and X's matching columns) split across n devices; partial results are summed (an all-reduce).
2328def tensor_parallel_mlp(X: np.ndarray, W1: np.ndarray, W2: np.ndarray, n: int) -> np.ndarray: 2329 """ReLU(X W1) W2 as Megatron-LM splits it: W1 by columns, W2 by rows, one all-reduce at the end. 2330 2331 Splitting W1 by columns gives each device whole hidden units, so the 2332 elementwise ReLU can run locally with no communication at all. 2333 """ 2334 W1s, W2s = np.array_split(W1, n, axis=1), np.array_split(W2, n, axis=0) 2335 partials = [np.maximum(X @ a, 0) @ b for a, b in zip(W1s, W2s)] # each device: (tokens, d_model) 2336 return sum(partials) # the one all-reduce
ReLU(X W1) W2 as Megatron-LM splits it: W1 by columns, W2 by rows, one all-reduce at the end.
Splitting W1 by columns gives each device whole hidden units, so the elementwise ReLU can run locally with no communication at all.
2339def pipeline_schedule(stages: int, microbatches: int) -> np.ndarray: 2340 """GPipe's fill-and-drain schedule as a (stages, time) grid. 2341 2342 Cell value j > 0: forward pass of micro-batch j; −j: its backward pass; 2343 0: the stage is idle (the bubble). Forward passes flow down the stages, 2344 then backward passes flow back up. 2345 """ 2346 p, m = stages, microbatches 2347 span = m + p - 1 2348 grid = np.zeros((p, 2 * span), dtype=int) 2349 for s in range(p): 2350 for j in range(m): 2351 grid[s, s + j] = j + 1 # stage s can start micro-batch j once stage s−1 has finished it 2352 grid[s, span + (p - 1 - s) + j] = -(j + 1) # backward starts at the last stage 2353 return grid
GPipe's fill-and-drain schedule as a (stages, time) grid.
Cell value j > 0: forward pass of micro-batch j; −j: its backward pass; 0: the stage is idle (the bubble). Forward passes flow down the stages, then backward passes flow back up.
2356def bubble_fraction(stages: int, microbatches: int) -> float: 2357 """Share of the schedule each stage spends idle: (p − 1) / (m + p − 1).""" 2358 return (stages - 1) / (microbatches + stages - 1)
Share of the schedule each stage spends idle: (p − 1) / (m + p − 1).
2366@dataclass(frozen=True) 2367class FloatFormat: 2368 """A binary floating-point format: 1 sign bit, `exp_bits` exponent bits, `man_bits` fraction bits.""" 2369 2370 name: str 2371 exp_bits: int 2372 man_bits: int 2373 max_override: float | None = None # E4M3 spends its top code on NaN only, so it reaches 448 not 240 2374 2375 @property 2376 def bias(self) -> int: 2377 return 2 ** (self.exp_bits - 1) - 1 2378 2379 @property 2380 def max_value(self) -> float: 2381 if self.max_override is not None: 2382 return self.max_override 2383 return (2 - 2.0**-self.man_bits) * 2.0**self.bias 2384 2385 @property 2386 def min_normal(self) -> float: 2387 return 2.0 ** (1 - self.bias) 2388 2389 @property 2390 def min_subnormal(self) -> float: 2391 return 2.0 ** (1 - self.bias - self.man_bits) 2392 2393 @property 2394 def epsilon(self) -> float: 2395 """Gap between 1 and the next representable number: the format's relative precision.""" 2396 return 2.0**-self.man_bits
2407def quantize(x, fmt: FloatFormat): 2408 """Round x to the nearest value `fmt` can hold (ties to even), as the hardware does. 2409 2410 Between 2ᵉ and 2ᵉ⁺¹ a format has 2^man_bits evenly spaced values, so the 2411 spacing is 2^(e − man_bits). Below the smallest normal number the spacing 2412 stops shrinking (subnormals), and anything under half the smallest 2413 subnormal rounds to 0: that is underflow. Above the largest value is 2414 overflow, shown here as infinity. 2415 """ 2416 x = np.asarray(x, dtype=np.float64) 2417 mag = np.abs(x) 2418 e = np.floor(np.log2(np.where(mag > 0, mag, 1.0))) 2419 e = np.maximum(e, 1 - fmt.bias) # subnormals share the smallest normal exponent's spacing 2420 step = 2.0 ** (e - fmt.man_bits) 2421 q = np.round(mag / step) * step # np.round rounds halves to even, like IEEE hardware 2422 q = np.where(q > fmt.max_value, np.inf, q) 2423 out = np.sign(x) * q 2424 return float(out) if out.ndim == 0 else out
Round x to the nearest value fmt can hold (ties to even), as the hardware does.
Between 2ᵉ and 2ᵉ⁺¹ a format has 2^man_bits evenly spaced values, so the spacing is 2^(e − man_bits). Below the smallest normal number the spacing stops shrinking (subnormals), and anything under half the smallest subnormal rounds to 0: that is underflow. Above the largest value is overflow, shown here as infinity.
2427def scaled_gradient_roundtrip(grad: float, scale: float, fmt: FloatFormat) -> float: 2428 """Multiply by the loss scale, store in `fmt` (where the backward pass happens), divide back in fp32.""" 2429 return quantize(grad * scale, fmt) / scale
Multiply by the loss scale, store in fmt (where the backward pass happens), divide back in fp32.
2432@dataclass 2433class DynamicLossScaler: 2434 """Keep the loss scale as large as possible without overflowing. 2435 2436 Overflow (an inf or NaN gradient) means the scale is too big: skip this 2437 step and halve it. A long run of clean steps means there may be headroom: 2438 double it. 2439 """ 2440 2441 scale: float = 2.0**16 2442 growth_interval: int = 2000 2443 clean_steps: int = 0 2444 2445 def update(self, grads_finite: bool) -> bool: 2446 """Record one step's outcome; return True if the optimizer should apply this step.""" 2447 if not grads_finite: 2448 self.scale /= 2 2449 self.clean_steps = 0 2450 return False 2451 self.clean_steps += 1 2452 if self.clean_steps == self.growth_interval: 2453 self.scale *= 2 2454 self.clean_steps = 0 2455 return True
Keep the loss scale as large as possible without overflowing.
Overflow (an inf or NaN gradient) means the scale is too big: skip this step and halve it. A long run of clean steps means there may be headroom: double it.
2445 def update(self, grads_finite: bool) -> bool: 2446 """Record one step's outcome; return True if the optimizer should apply this step.""" 2447 if not grads_finite: 2448 self.scale /= 2 2449 self.clean_steps = 0 2450 return False 2451 self.clean_steps += 1 2452 if self.clean_steps == self.growth_interval: 2453 self.scale *= 2 2454 self.clean_steps = 0 2455 return True
Record one step's outcome; return True if the optimizer should apply this step.
2458def accumulate_updates(w0: float, update: float, steps: int, fmt: FloatFormat) -> float: 2459 """Add `update` to a weight `steps` times, rounding the weight to `fmt` after every step.""" 2460 w = quantize(w0, fmt) 2461 for _ in range(steps): 2462 w = quantize(w + update, fmt) 2463 return w
Add update to a weight steps times, rounding the weight to fmt after every step.
2477def train_through_a_bad_batch(clip: float | None, steps: int = 60, bad_step: int = 30, lr: float = 0.1, 2478 seed: int = 0) -> list[float]: 2479 """SGD on y = 3x, with one corrupted batch (labels flipped and blown up 100×) at `bad_step`. 2480 2481 Returns the loss on clean held-out data after every step. With `clip`, each 2482 gradient is scaled down to at most that length before the update. 2483 """ 2484 rng = np.random.default_rng(seed) 2485 x_eval = rng.standard_normal(256) 2486 w = np.array([2.5]) 2487 losses = [] 2488 for t in range(steps): 2489 x = rng.standard_normal(32) 2490 y = -300 * x if t == bad_step else 3 * x 2491 g = linear_regression_gradient(x[:, None], y, w) 2492 if clip is not None: 2493 (g,), _ = clip_by_global_norm([g], clip) 2494 w = w - lr * g 2495 losses.append(float(np.mean((w[0] * x_eval - 3 * x_eval) ** 2))) 2496 return losses
SGD on y = 3x, with one corrupted batch (labels flipped and blown up 100×) at bad_step.
Returns the loss on clean held-out data after every step. With clip, each
gradient is scaled down to at most that length before the update.
2499def detect_spikes(losses: list[float], window: int = 50, factor: float = 2.0) -> list[int]: 2500 """Steps whose loss exceeds `factor` × the median of the previous `window` losses. 2501 2502 The median, not the mean, so one spike does not raise the bar for the next. 2503 """ 2504 return [i for i in range(window, len(losses)) if losses[i] > factor * float(np.median(losses[i - window : i]))]
Steps whose loss exceeds factor × the median of the previous window losses.
The median, not the mean, so one spike does not raise the bar for the next.
2507def wasted_fraction(interval: float, cost: float, mtbf: float) -> float: 2508 """Share of time lost to checkpointing: saving (cost/interval) plus redoing lost work (interval/(2·mtbf)).""" 2509 return cost / interval + interval / (2 * mtbf)
Share of time lost to checkpointing: saving (cost/interval) plus redoing lost work (interval/(2·mtbf)).
2512def optimal_checkpoint_interval(cost: float, mtbf: float) -> float: 2513 """Young's (1974) interval √(2·cost·mtbf), which minimises `wasted_fraction`.""" 2514 return math.sqrt(2 * cost * mtbf)
Young's (1974) interval √(2·cost·mtbf), which minimises wasted_fraction.
2522def figures() -> dict: 2523 """Plot this lesson's data. matplotlib is imported here, and only here, 2524 so the lesson itself needs nothing beyond NumPy.""" 2525 import matplotlib 2526 2527 matplotlib.use("Agg") 2528 import matplotlib.pyplot as plt 2529 2530 BLUE, RED, GREEN, AMBER, MUTED = "#2563eb", "#dc2626", "#059669", "#d97706", "#9ca3af" 2531 figs = {} 2532 2533 # --- 1. MinHash estimates converge on the true Jaccard ------------------- 2534 fig, ax = plt.subplots(figsize=(6, 3.4)) 2535 ks = np.unique(np.logspace(0, np.log10(256), 40).astype(int)) 2536 hasher = MinHasher(256, seed=1) 2537 for true_j, color in ((0.2, BLUE), (0.5, GREEN), (0.8, RED)): 2538 # |A| = |B| = 100 with overlap o gives J = o / (200 − o), so o = 200·J / (1 + J). 2539 o = round(200 * true_j / (1 + true_j)) 2540 a = {f"s{i}" for i in range(100)} 2541 b = {f"s{i}" for i in range(100 - o, 200 - o)} 2542 sa, sb = hasher.signature(a), hasher.signature(b) 2543 ax.plot(ks, [estimate_jaccard(sa[:k], sb[:k]) for k in ks], color=color, label=f"true J = {jaccard(a, b):.2f}") 2544 ax.axhline(jaccard(a, b), color=color, ls="--", lw=1) 2545 ax.set_xscale("log", base=2) 2546 ax.set_xlabel("number of hash functions k") 2547 ax.set_ylabel("MinHash estimate of J") 2548 ax.set_ylim(-0.05, 1.05) 2549 ax.set_title("MinHash: more hashes, tighter estimate") 2550 ax.legend(frameon=False, loc="center right", bbox_to_anchor=(1.0, 0.64)) 2551 figs["minhash_convergence"] = fig 2552 2553 # --- 2. LSH S-curves ----------------------------------------------------- 2554 fig, ax = plt.subplots(figsize=(6, 3.4)) 2555 s = np.linspace(0, 1, 201) 2556 for (b, r), color in (((50, 2), BLUE), ((20, 5), GREEN), ((10, 10), RED)): 2557 ax.plot(s, lsh_candidate_probability(s, b, r), color=color, label=f"{b} bands × {r} rows") 2558 ax.axvline((1 / b) ** (1 / r), color=color, ls="--", lw=1) 2559 ax.set_xlabel("true Jaccard similarity of a pair") 2560 ax.set_ylabel("chance the pair is compared") 2561 ax.set_title("LSH banding: an S-curve you can place") 2562 ax.legend(frameon=False, loc="lower right") 2563 figs["lsh_s_curve"] = fig 2564 2565 # --- 3. Model collapse --------------------------------------------------- 2566 fig, ax = plt.subplots(figsize=(6, 3.4)) 2567 gens = np.arange(1, 201) 2568 for seed, color in ((0, BLUE), (1, GREEN), (2, AMBER)): 2569 ax.semilogy(gens, recursive_gaussian_fit(200, 20, seed), color=color, lw=1.2, label=f"seed {seed}") 2570 ax.semilogy(gens, np.sqrt((19 / 20) ** gens), color="#4b5563", ls="--", label="average shrink √(0.95ᵗ)") 2571 ax.set_xlabel("generation (each fitted to 20 samples of the last)") 2572 ax.set_ylabel("fitted spread σ") 2573 ax.set_title("Training on your own samples: the spread collapses") 2574 ax.legend(frameon=False, loc="lower left") 2575 figs["model_collapse"] = fig 2576 2577 # --- 4. The 7B memory bill ------------------------------------------------ 2578 fig, ax = plt.subplots(figsize=(6.4, 3.6)) 2579 state = training_memory(7e9) 2580 acts = { 2581 "naive": activation_bytes(4096, 1, 4096, 32, 32), 2582 "no stored scores": activation_bytes(4096, 1, 4096, 32, 32, store_scores=False), 2583 "checkpointed": activation_bytes(4096, 1, 4096, 32, 32, checkpointed=True), 2584 } 2585 colors = [BLUE, "#60a5fa", GREEN, "#34d399", "#a7f3d0"] 2586 x = np.arange(len(acts)) 2587 bottom = np.zeros(len(acts)) 2588 for (name, val), color in zip([(k, v) for k, v in state.items() if k != "total"], colors): 2589 ax.bar(x, val / 1e9, bottom=bottom, color=color, label=name, width=0.55) 2590 bottom += val / 1e9 2591 ax.bar(x, [a / 1e9 for a in acts.values()], bottom=bottom, color=AMBER, label="activations", width=0.55) 2592 for xi, a in zip(x, acts.values()): 2593 ax.text(xi, bottom[xi] + a / 1e9 + 3, f"{(state['total'] + a) / 1e9:.0f} GB", ha="center") 2594 ax.axhline(80, color=RED, ls="--", label="one 80 GB GPU") 2595 ax.set_xticks(x, list(acts)) 2596 ax.set_ylabel("GB on one GPU") 2597 ax.set_ylim(0, 245) 2598 ax.set_title("Training a 7B model on one GPU: it does not fit") 2599 ax.legend(frameon=False, fontsize=8, loc="upper right") 2600 figs["memory_7b"] = fig 2601 2602 # --- 5. ZeRO stages ------------------------------------------------------- 2603 fig, ax = plt.subplots(figsize=(6, 3.6)) 2604 ns = 2 ** np.arange(0, 11) 2605 for stage, color in zip(range(4), (MUTED, BLUE, GREEN, RED)): 2606 ax.loglog(ns, [zero_memory_per_gpu(7e9, int(n), stage) / 1e9 for n in ns], "o-", color=color, ms=3, 2607 label=["stage 0: plain data parallel", "stage 1: + shard optimizer", "stage 2: + shard gradients", 2608 "stage 3 (FSDP): + shard weights"][stage]) 2609 ax.axhline(80, color="#4b5563", ls="--") 2610 ax.text(2**7, 95, "80 GB GPU", color="#4b5563", zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1)) 2611 ax.set_xscale("log", base=2) 2612 ax.set_xlabel("GPUs sharing the state") 2613 ax.set_ylabel("training state per GPU (GB)") 2614 ax.set_title("ZeRO: sharding divides the 16 bytes per parameter") 2615 ax.legend(frameon=False, fontsize=8, loc="lower left") 2616 figs["zero_stages"] = fig 2617 2618 # --- 6. The pipeline schedule -------------------------------------------- 2619 from matplotlib.colors import ListedColormap 2620 2621 p, m = 4, 8 2622 grid = pipeline_schedule(p, m) 2623 fig, ax = plt.subplots(figsize=(8, 2.6)) 2624 kind = np.sign(grid) + 1 # 0 backward, 1 idle, 2 forward 2625 ax.imshow(kind, cmap=ListedColormap(["#86efac", "#ffffff", "#93c5fd"]), aspect="auto", vmin=0, vmax=2) 2626 for (row, col), v in np.ndenumerate(grid): 2627 if v: 2628 ax.text(col, row, str(abs(v)), ha="center", va="center", fontsize=8) 2629 ax.set_xticks(np.arange(-0.5, grid.shape[1], 1), minor=True) 2630 ax.set_yticks(np.arange(-0.5, p, 1), minor=True) 2631 ax.grid(which="minor", color="#d1d5db", lw=0.6) 2632 ax.grid(which="major", visible=False) 2633 ax.tick_params(which="minor", length=0) 2634 ax.set_yticks(range(p), [f"GPU {i + 1}" for i in range(p)]) 2635 ax.set_xticks(range(0, grid.shape[1], 2)) 2636 ax.set_xlabel("time step (blue: forward of micro-batch n, green: backward, white: idle)") 2637 ax.set_title(f"GPipe schedule, {p} stages, {m} micro-batches: bubble = {bubble_fraction(p, m):.0%}") 2638 figs["pipeline_schedule"] = fig 2639 2640 # --- 7. Bubble fraction --------------------------------------------------- 2641 fig, ax = plt.subplots(figsize=(6, 3.4)) 2642 ms = np.arange(1, 257) 2643 for p, color in zip((2, 4, 8, 16), (BLUE, GREEN, AMBER, RED)): 2644 ax.semilogx(ms, [bubble_fraction(p, int(m)) for m in ms], color=color, label=f"{p} stages") 2645 ax.axhline(0.1, color=MUTED, ls="--") 2646 ax.set_xlabel("micro-batches per batch") 2647 ax.set_ylabel("share of time idle") 2648 ax.set_title("The pipeline bubble: (p − 1) / (m + p − 1)") 2649 ax.legend(frameon=False) 2650 figs["bubble_fraction"] = fig 2651 2652 # --- 8. Float ranges ------------------------------------------------------ 2653 fig, ax = plt.subplots(figsize=(6.4, 3.0)) 2654 for i, fmt in enumerate(FORMATS): 2655 lo, mid, hi = np.log10(fmt.min_subnormal), np.log10(fmt.min_normal), np.log10(fmt.max_value) 2656 ax.barh(i, mid - lo, left=lo, color="#bfdbfe", height=0.6) 2657 ax.barh(i, hi - mid, left=mid, color=BLUE, height=0.6) 2658 ax.set_yticks(range(len(FORMATS)), [f.name for f in FORMATS]) 2659 ax.invert_yaxis() 2660 ax.set_xlabel("log₁₀ of magnitude (light: subnormals, dark: normal range)") 2661 ax.set_title("Range of each format: bf16 keeps all of fp32's") 2662 figs["float_ranges"] = fig 2663 2664 # --- 9. Loss scaling ------------------------------------------------------ 2665 rng = np.random.default_rng(0) 2666 grads = np.exp(rng.normal(np.log(2.0**-22), 3.0, 1_000_000)) # sizes spread over many powers of two 2667 fig, ax = plt.subplots(figsize=(6.4, 3.4)) 2668 bins = np.arange(-50, 20, 0.5) 2669 lost = float(np.mean(quantize(grads, FP16) == 0)) 2670 lost_scaled = float(np.mean(quantize(grads * 2.0**16, FP16) == 0)) 2671 ax.hist(np.log2(grads), bins=bins, color=MUTED, alpha=0.8, label=f"unscaled: {lost:.0%} become 0 in fp16") 2672 ax.hist(np.log2(grads * 2.0**16), bins=bins, color=BLUE, alpha=0.6, 2673 label=f"× 2¹⁶: {lost_scaled:.2%} become 0") 2674 top = ax.get_ylim()[1] * 1.5 # headroom above the peaks for the labels and legend 2675 ax.set_ylim(0, top) 2676 ax.axvspan(-50, np.log2(FP16.min_subnormal / 2), color=RED, alpha=0.12) 2677 ax.axvline(np.log2(FP16.max_value), color=RED, ls="--") 2678 ax.text(np.log2(FP16.max_value) - 0.5, top * 0.93, "fp16 overflow", color=RED, ha="right") 2679 ax.text(-49, top * 0.93, "flushed to 0", color=RED) 2680 ax.set_xlabel("gradient size, log₂") 2681 ax.set_ylabel("number of gradients") 2682 ax.set_title("Loss scaling slides the gradients into fp16's range") 2683 ax.legend(frameon=False, loc="upper center", bbox_to_anchor=(0.5, 0.88), fontsize=8) 2684 figs["loss_scaling"] = fig 2685 2686 # --- 10. Master weights --------------------------------------------------- 2687 fig, ax = plt.subplots(figsize=(6, 3.2)) 2688 for fmt, color, ls in ((FP32, BLUE, "-"), (BF16, RED, "-"), (FP16, AMBER, ":")): 2689 w, path = quantize(1.0, fmt), [] 2690 for _ in range(1000): 2691 w = quantize(w + 1e-4, fmt) 2692 path.append(w) 2693 ax.plot(range(1, 1001), path, color=color, ls=ls, lw=2, label=fmt.name) 2694 ax.set_xlabel("update step (each adds 0.0001)") 2695 ax.set_ylabel("weight value") 2696 ax.set_title("Tiny updates vanish in 16 bits; an fp32 master keeps them") 2697 ax.legend(frameon=False) 2698 figs["master_weights"] = fig 2699 2700 # --- 11. One bad batch, with and without clipping ------------------------ 2701 fig, ax = plt.subplots(figsize=(6, 3.4)) 2702 for clip, color, label in ((None, RED, "no clipping"), (1.0, BLUE, "clip global norm at 1")): 2703 losses = np.maximum(train_through_a_bad_batch(clip), 1e-12) 2704 ax.semilogy(losses, color=color, label=label) 2705 ax.axvline(30, color=MUTED, ls="--") 2706 ax.text(31, 1e-8, "corrupted batch", color="#4b5563") 2707 ax.set_xlabel("training step") 2708 ax.set_ylabel("clean held-out loss") 2709 ax.set_title("One bad batch: clipping limits the damage") 2710 ax.legend(frameon=False, loc="upper right") 2711 figs["bad_batch"] = fig 2712 2713 # --- 12. Checkpoint interval ---------------------------------------------- 2714 fig, ax = plt.subplots(figsize=(6, 3.4)) 2715 T = np.logspace(0, np.log10(300), 200) 2716 for cost, color, label in ((1.0, BLUE, "1-minute save"), (1 / 6, GREEN, "10-second save")): 2717 wasted = np.array([wasted_fraction(t, cost, 180) for t in T]) 2718 ax.semilogx(T, np.where(wasted <= 0.5, wasted, np.nan), color=color, label=label) # above 50% leaves the chart 2719 best = optimal_checkpoint_interval(cost, 180) 2720 ax.plot(best, wasted_fraction(best, cost, 180), "o", color=color) 2721 ax.annotate(f"{best:.1f} min, {wasted_fraction(best, cost, 180):.1%}", (best, wasted_fraction(best, cost, 180)), 2722 textcoords="offset points", xytext=(6, -14), color=color, zorder=3, bbox=dict(facecolor="white", edgecolor="none", pad=1)) 2723 ax.set_xlabel("minutes between checkpoints (failure every 180 min)") 2724 ax.set_ylabel("share of time wasted") 2725 ax.set_ylim(0, 0.5) 2726 ax.set_title("Checkpoint interval: a valley at √(2·C·M)") 2727 ax.legend(frameon=False) 2728 figs["checkpoint_interval"] = fig 2729 2730 return figs
Plot this lesson's data. matplotlib is imported here, and only here, so the lesson itself needs nothing beyond NumPy.
2738def demo() -> None: 2739 banner("1. Cleaning a crawl: language, quality, exact and near duplicates") 2740 say( 2741 """ 2742 Seven pages arrive from a crawl. Cheap per-page checks run first 2743 (language ID, quality rules); the comparisons between pages 2744 (exact hash, then MinHash) run last, on what survived. 2745 """ 2746 ) 2747 labels = ["river paragraph", "French version", "navigation bar", "re-crawled copy", "hashtag spam", 2748 "scraper's edited copy", "stars paragraph"] 2749 table(["page", "what it is", "verdict"], [(i, l, v) for i, (l, v) in enumerate(zip(labels, curate(CRAWL_SAMPLE)))]) 2750 clf = QualityClassifier.trained_on_examples() 2751 for text in ("the delta is formed when the river deposits sand over many years", 2752 "click here buy now best price free free free"): 2753 print(f" classifier score {clf.score(text):+7.2f} {text!r}") 2754 print() 2755 takeaway("Most of a crawl never reaches the model: two of seven pages survive here.") 2756 2757 banner("2. MinHash and LSH: near duplicates without comparing everything") 2758 a, b = CRAWL_SAMPLE[0], CRAWL_SAMPLE[5] 2759 true_j = jaccard(shingles(a), shingles(b)) 2760 rows = [] 2761 for k in (8, 32, 128, 512): 2762 h = MinHasher(k, seed=0) 2763 rows.append((k, true_j, estimate_jaccard(h.signature(shingles(a)), h.signature(shingles(b))))) 2764 table(["hashes k", "true Jaccard", "MinHash estimate"], rows, floatfmt=".3f") 2765 table(["similarity s", "P(candidate), 20 bands × 5 rows"], 2766 [(s, lsh_candidate_probability(s, 20, 5)) for s in (0.3, 0.5, 0.7, 0.8, 0.9)], floatfmt=".4f") 2767 pairs = near_duplicate_pairs([a, b, CRAWL_SAMPLE[6]]) 2768 say(f"LSH over pages 0, 5 and 6 (listed as 0, 1, 2) flags {pairs}: pages 0 and 5, and nothing else.") 2769 2770 banner("3. Why duplicates hurt: a bigram model memorizes what it sees often") 2771 table(["copies of the boilerplate line", "P(model continues 'click' into the whole line)"], 2772 [(c, verbatim_probability(c)) for c in (1, 10, 100, 1000)], floatfmt=".3f") 2773 takeaway("Repetition turns learning into recitation; deduplication is also privacy and eval hygiene.") 2774 2775 banner("4. Mixtures, budgets and synthetic data") 2776 epochs = epochs_per_source({"web": 0.80, "wiki": 0.05, "code": 0.15}, 2777 {"web": 900e9, "wiki": 20e9, "code": 150e9}, budget=1000e9) 2778 table(["source", "weight", "passes over it"], [(k, w, epochs[k]) for k, w in (("web", 0.80), ("wiki", 0.05), ("code", 0.15))], 2779 floatfmt=".2f") 2780 say(f"A 7B model's compute-optimal budget: {chinchilla_tokens(7e9) / 1e9:.0f}B tokens, " 2781 f"costing {training_flops(7e9, chinchilla_tokens(7e9)):.2e} FLOPs.") 2782 stds = recursive_gaussian_fit() 2783 say(f"Refit on your own samples 200 times: spread 1.0 -> {stds[49]:.3f} after 50 generations, " 2784 f"{stds[-1]:.1e} after 200. Keep real data in the mix.") 2785 2786 banner("5. The memory bill: why one GPU is not enough") 2787 mem = training_memory(7e9) 2788 table(["what", "GB for 7B"], [(k, v / 1e9) for k, v in mem.items()], floatfmt=".1f") 2789 table(["activations at 4,096 tokens", "GB"], [ 2790 ("stored naively", activation_bytes(4096, 1, 4096, 32, 32) / 1e9), 2791 ("without attention scores", activation_bytes(4096, 1, 4096, 32, 32, store_scores=False) / 1e9), 2792 ("activation checkpointing", activation_bytes(4096, 1, 4096, 32, 32, checkpointed=True) / 1e9), 2793 ], floatfmt=".1f") 2794 table(["ZeRO stage", "GB per GPU, 7.5B on 64 GPUs"], [(s, zero_memory_per_gpu(7.5e9, 64, s) / 1e9) for s in range(4)], 2795 floatfmt=".2f") 2796 2797 banner("6. Parallelism: all-reduce, tensor splits, the pipeline bubble") 2798 grads = [np.arange(8.0) * (i + 1) for i in range(4)] 2799 results, sent = ring_all_reduce(grads) 2800 say(f"Ring all-reduce over 4 workers: everyone ends with {results[0].tolist()}; " 2801 f"each sent {sent[0]} numbers, 1.5 times its own 8.") 2802 rng = np.random.default_rng(0) 2803 X, W1, W2 = rng.standard_normal((3, 8)), rng.standard_normal((8, 16)), rng.standard_normal((16, 8)) 2804 err = np.abs(tensor_parallel_mlp(X, W1, W2, n=4) - np.maximum(X @ W1, 0) @ W2).max() 2805 say(f"Feed-forward layer split over 4 devices (columns, then rows): largest difference from one device = {err:.1e}.") 2806 print(pipeline_schedule(4, 4)) 2807 print() 2808 table(["micro-batches (4 stages)", "bubble"], [(m, bubble_fraction(4, m)) for m in (1, 4, 8, 32)], floatfmt=".3f") 2809 2810 banner("7. Mixed precision: range, precision, loss scaling, master weights") 2811 table(["format", "largest", "smallest normal", "gap after 1"], 2812 [(f.name, f"{f.max_value:.3g}", f"{f.min_normal:.3g}", f"{f.epsilon:.3g}") for f in FORMATS]) 2813 say(f"A gradient of 1e-8 in fp16: {quantize(1e-8, FP16)}. Scaled by 65536, stored, unscaled: " 2814 f"{scaled_gradient_roundtrip(1e-8, 65536, FP16):.4e}. In bf16 with no scaling: {quantize(1e-8, BF16):.4e}.") 2815 scaler = DynamicLossScaler(scale=65536.0, growth_interval=3) 2816 events = [True, True, False, True, True, True] 2817 log = [] 2818 for ok in events: 2819 applied = scaler.update(ok) 2820 log.append(("clean" if ok else "overflow", "applied" if applied else "skipped", scaler.scale)) 2821 table(["gradients", "step", "scale after"], log, floatfmt=".0f") 2822 say(f"1,000 updates of 0.0001 to a weight of 1.0: fp32 ends at {accumulate_updates(1.0, 1e-4, 1000, FP32):.4f}, " 2823 f"bf16 at {accumulate_updates(1.0, 1e-4, 1000, BF16):.4f}.") 2824 takeaway("bf16 for the maths (fp32's range), fp32 for the master weights and optimizer state.") 2825 2826 banner("8. Stability: clipping, spikes, checkpoints") 2827 no_clip, clipped = train_through_a_bad_batch(None), train_through_a_bad_batch(1.0) 2828 say(f"One corrupted batch at step 30: peak loss {max(no_clip):,.0f} without clipping, " 2829 f"{max(clipped[30:]):.4f} with clipping.") 2830 say(f"Spikes flagged in [3.0, 2.9, 2.8, 2.8, 2.7, 8.5, 4.0, 2.7]: steps " 2831 f"{detect_spikes([3.0, 2.9, 2.8, 2.8, 2.7, 8.5, 4.0, 2.7], window=4, factor=2.0)}.") 2832 table(["save time (min)", "best interval (min)", "time wasted"], 2833 [(c, optimal_checkpoint_interval(c, 180), wasted_fraction(optimal_checkpoint_interval(c, 180), c, 180)) 2834 for c in (5.0, 1.0, 1 / 6)], floatfmt=".3f") 2835 takeaway("At thousands of GPUs, failures are the weather: clip, watch, roll back, and save often and fast.")