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_stages makes 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.

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

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

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

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

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

Signatures make each comparison cheap, but a billion pages still make too many pairs. Locality-sensitive hashing (LSH) avoids most comparisons: cut each signature into $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

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

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

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

1h. Why duplicates hurt

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

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

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

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

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

1i. The whole belt on a small crawl

CRAWL_SAMPLE holds seven pages. Running curate over it:

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

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

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

1j. Data mixtures: how much of each source

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

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

Level 3: the formula and its symbols

$$ \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

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

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

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

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.

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

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

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

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.

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

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

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

3c. Tensor parallelism: splitting one matrix multiply

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

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

Tiny worked example. X = [1, 2] and

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

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

$$ 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]

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

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

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

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

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

3e. All three at once

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.

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

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

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

4b. Underflow, and loss scaling

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

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

Level 3: the formula and its symbols

$$ 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'

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

Reading it: the horizontal axis is a gradient's size in powers of two and the height is how many gradients have that size. The shaded region on the left is below fp16's smallest subnormal: whatever lands there becomes

  1. 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)

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

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

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

4d. fp8: the next halving

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

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

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]]

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

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

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

5c. Warmup

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

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

5d. Checkpoints: how often to save

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

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

Level 3: the formula and its symbols

$$ \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)

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

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

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

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

on GitHub
   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![Three pairs of sets with true Jaccard 0.2, 0.5 and 0.8: with a handful of hashes the estimates jump around; by 256 hashes each sits close to its dashed true value](figures/primer.ml.pretraining.minhash_convergence.svg)
 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![Probability of becoming a candidate against true similarity, for three band layouts of about 100 hashes: each is an S-curve, steeper and further right as rows per band grow](figures/primer.ml.pretraining.lsh_s_curve.svg)
 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![Fitted spread over 200 generations of refitting on 20 samples, for three random seeds on a log scale: every run falls by four to five orders of magnitude, even faster than the dashed line from the formula](figures/primer.ml.pretraining.model_collapse.svg)
 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![Stacked bars for a 7B model at 4,096 tokens: the 112 GB of training state alone passes the 80 GB line, and naive activations add another 104 GB](figures/primer.ml.pretraining.memory_7b.svg)
 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![Training state per GPU for a 7B model as the number of GPUs grows from 1 to 1,024: stage 0 stays at 112 GB, stages 1 and 2 flatten at their unsharded floors, and stage 3 keeps falling](figures/primer.ml.pretraining.zero_stages.svg)
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![The GPipe schedule for 4 stages and 8 micro-batches as a grid of GPUs by time step: forward passes form a staircase down, backward passes a staircase back up, and the empty triangles in the corners are the bubble](figures/primer.ml.pretraining.pipeline_schedule.svg)
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![Bubble fraction against number of micro-batches for 2, 4, 8 and 16 stages: every curve starts high and falls towards zero, and deeper pipelines need more micro-batches for the same efficiency](figures/primer.ml.pretraining.bubble_fraction.svg)
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![Horizontal bars on a log scale showing the range of each format from its smallest subnormal to its largest value: fp32 and bf16 span the same huge range, fp16 and fp8 are far narrower](figures/primer.ml.pretraining.float_ranges.svg)
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![Histogram of the size of a million simulated gradients on a log scale: unscaled, a large share falls left of fp16's smallest value and is lost; multiplied by 65,536 the whole histogram shifts right into fp16's range](figures/primer.ml.pretraining.loss_scaling.svg)
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![A weight receiving 1,000 updates of 0.0001: in fp32 it climbs in a straight line from 1.0 to 1.1; stored in bf16 or fp16 it never leaves 1.0](figures/primer.ml.pretraining.master_weights.svg)
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![Clean held-out loss over 60 steps of training on y = 3x with one corrupted batch at step 30: without clipping the loss leaps above 6,000 and takes dozens of steps to come back; with clipping it barely moves](figures/primer.ml.pretraining.bad_batch.svg)
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![Share of time wasted against checkpoint interval for a 1-minute and a 10-second save with a failure every 3 hours: each curve is a valley whose floor is marked at the square-root interval](figures/primer.ml.pretraining.checkpoint_interval.svg)
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()
Level 3: the code, function by function.
STOP_WORDS = {'en': {'and', 'is', 'was', 'for', 'to', 'it', 'the', 'that', 'in', 'of'}, 'fr': {'et', 'dans', 'du', 'des', 'est', 'un', 'une', 'le', 'la', 'les'}, 'de': {'nicht', 'ein', 'und', 'eine', 'ist', 'die', 'zu', 'das', 'der', 'mit'}, 'es': {'por', 'que', 'del', 'una', 'las', 'es', 'y', 'el', 'con', 'los'}}
def words(text: str) -> list[str]: on GitHub
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.

def detect_language(text: str, min_share: float = 0.1) -> str: on GitHub
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.

GOPHER_STOP_WORDS = {'and', 'to', 'the', 'that', 'have', 'with', 'be', 'of'}
def quality_failures(text: str) -> list[str]: on GitHub
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.

GOOD_EXAMPLES = ['the river deposits sand at its mouth and over many years builds a delta', 'a delta is formed where a river meets the sea and slows down', 'the cell is the basic unit of life and every organism is made of cells', 'light travels faster than sound which is why we see lightning before thunder', 'the history of the city begins with a small settlement on the river', 'photosynthesis converts light into chemical energy stored in sugar']
BAD_EXAMPLES = ['click here buy now limited offer best price', 'free free free win a prize click now', 'best deals cheap price buy buy buy', 'subscribe now click the link for free coupons', 'hot singles best price click here now', 'lose weight fast buy now free shipping']
@dataclass
class QualityClassifier: on GitHub
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.

QualityClassifier( good: collections.Counter = <factory>, bad: collections.Counter = <factory>)
good: collections.Counter
bad: collections.Counter
@classmethod
def trained_on_examples( cls, good=['the river deposits sand at its mouth and over many years builds a delta', 'a delta is formed where a river meets the sea and slows down', 'the cell is the basic unit of life and every organism is made of cells', 'light travels faster than sound which is why we see lightning before thunder', 'the history of the city begins with a small settlement on the river', 'photosynthesis converts light into chemical energy stored in sugar'], bad=['click here buy now limited offer best price', 'free free free win a prize click now', 'best deals cheap price buy buy buy', 'subscribe now click the link for free coupons', 'hot singles best price click here now', 'lose weight fast buy now free shipping']) -> QualityClassifier: on GitHub
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
def word_vote(self, w: str) -> float: on GitHub
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.

def score(self, text: str) -> float: on GitHub
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.

def normalize_for_exact(text: str) -> str: on GitHub
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.

def exact_dedup(docs: list[str]) -> list[str]: on GitHub
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.

def shingles(text: str, k: int = 5) -> set[str]: on GitHub
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.

def jaccard(a: set, b: set) -> float: on GitHub
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.

class MinHasher: on GitHub
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.

MinHasher(num_perm: int = 128, seed: int = 0) on GitHub
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)
a
b
def signature(self, shingle_set: set[str]) -> numpy.ndarray: on GitHub
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,)
def estimate_jaccard(sig_a: numpy.ndarray, sig_b: numpy.ndarray) -> float: on GitHub
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.

def lsh_candidate_probability(s: float, bands: int, rows: int) -> float: on GitHub
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ʳ)ᵇ.

def near_duplicate_pairs( docs: list[str], bands: int = 20, rows: int = 5, threshold: float = 0.7, seed: int = 0) -> list[tuple[int, int]]: on GitHub
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.

CRAWL_SAMPLE = ['The river carries sand and small stones from the mountains to the sea. Over thousands of years that sand settles at the mouth of the river and builds a wide, flat delta. Farmers have grown rice on these deltas for centuries, because the soil is rich and the water is close. When the river floods, it brings new soil with it, which is why the land stays fertile.', "Le fleuve transporte le sable et les pierres des montagnes vers la mer. Avec le temps, le sable se dépose à l'embouchure et forme un delta. Les paysans cultivent le riz dans les deltas depuis des siècles, car la terre est riche et l'eau est proche.", 'Home | About us | Contact | Privacy policy | Log in', 'the river carries sand and small stones from the mountains to the sea. Over thousands of years that sand settles at the mouth of the river and builds a wide, flat delta. Farmers have grown rice on these deltas for centuries, because the soil is rich and the water is close. When the river floods, it brings new soil with it, which is why the land stays fertile.', '#deal #sale #win the best price today #deal #sale #win the best price today #deal #sale #win the best price today #deal #sale #win the best price today #deal #sale #win the best price today #deal #sale #win the best price today #deal #sale #win the best price today #deal #sale #win the best price today #deal #sale #win the best price today #deal #sale #win the best price today #deal #sale #win the best price today #deal #sale #win the best price today', 'The river carries sand and small stones from the mountains to the sea. Over thousands of years that sand settles at the mouth of the river and builds a wide, flat delta. Farmers have grown rice on these deltas for hundreds of years, because the soil is rich and the water is close. When the river floods, it brings new soil with it, which is why the land stays fertile. Read more on our site.', 'Stars are born inside cold clouds of gas and dust. When part of a cloud becomes dense enough, gravity pulls it together faster than the gas can push back. The centre heats up as it shrinks, and once it reaches about ten million degrees, hydrogen begins to fuse into helium. That is the moment a star switches on, and it will shine with that fuel for millions or billions of years.']
def curate(docs: list[str], language: str = 'en') -> list[str]: on GitHub
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.

BOILERPLATE = 'click here to subscribe for free updates'
OTHER_LINES = ['click the map to see the river', 'we walk here every day', 'we went to the river', 'subscribe for the weekly news', 'free maps for every school']
def verbatim_probability( copies: int, line: str = 'click here to subscribe for free updates') -> float: on GitHub
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.

def epochs_per_source( weights: dict[str, float], sizes: dict[str, float], budget: float) -> dict[str, float]: on GitHub
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.

def chinchilla_tokens(params: float) -> float: on GitHub
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).

def training_flops(params: float, tokens: float) -> float: on GitHub
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.

def recursive_gaussian_fit(generations: int = 200, n: int = 20, seed: int = 0) -> list[float]: on GitHub
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.

def training_memory(params: float) -> dict[str, float]: on GitHub
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.

def activation_bytes( seq: int, batch: int, hidden: int, heads: int, layers: int, store_scores: bool = True, checkpointed: bool = False) -> float: on GitHub
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.

def zero_memory_per_gpu( params: float, n_gpus: int, stage: int, optimizer_bytes: int = 12) -> float: on GitHub
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.

def linear_regression_gradient(X: numpy.ndarray, y: numpy.ndarray, w: numpy.ndarray) -> numpy.ndarray: on GitHub
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).

def ring_all_reduce(arrays: list[numpy.ndarray]) -> tuple[list[numpy.ndarray], list[int]]: on GitHub
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.

def all_reduce_traffic(size: float, n: int) -> float: on GitHub
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.

def column_parallel_matmul(X: numpy.ndarray, W: numpy.ndarray, n: int) -> numpy.ndarray: on GitHub
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).

def row_parallel_matmul(X: numpy.ndarray, W: numpy.ndarray, n: int) -> numpy.ndarray: on GitHub
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).

def tensor_parallel_mlp( X: numpy.ndarray, W1: numpy.ndarray, W2: numpy.ndarray, n: int) -> numpy.ndarray: on GitHub
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.

def pipeline_schedule(stages: int, microbatches: int) -> numpy.ndarray: on GitHub
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.

def bubble_fraction(stages: int, microbatches: int) -> float: on GitHub
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).

@dataclass(frozen=True)
class FloatFormat: on GitHub
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

A binary floating-point format: 1 sign bit, exp_bits exponent bits, man_bits fraction bits.

FloatFormat( name: str, exp_bits: int, man_bits: int, max_override: float | None = None)
name: str
exp_bits: int
man_bits: int
max_override: float | None = None
bias: int on GitHub
2375    @property
2376    def bias(self) -> int:
2377        return 2 ** (self.exp_bits - 1) - 1
max_value: float on GitHub
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
min_normal: float on GitHub
2385    @property
2386    def min_normal(self) -> float:
2387        return 2.0 ** (1 - self.bias)
min_subnormal: float on GitHub
2389    @property
2390    def min_subnormal(self) -> float:
2391        return 2.0 ** (1 - self.bias - self.man_bits)
epsilon: float on GitHub
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

Gap between 1 and the next representable number: the format's relative precision.

FP32 = FloatFormat(name='fp32', exp_bits=8, man_bits=23, max_override=None)
BF16 = FloatFormat(name='bf16', exp_bits=8, man_bits=7, max_override=None)
FP16 = FloatFormat(name='fp16', exp_bits=5, man_bits=10, max_override=None)
FP8_E4M3 = FloatFormat(name='fp8 E4M3', exp_bits=4, man_bits=3, max_override=448.0)
FP8_E5M2 = FloatFormat(name='fp8 E5M2', exp_bits=5, man_bits=2, max_override=None)
FORMATS = [FloatFormat(name='fp32', exp_bits=8, man_bits=23, max_override=None), FloatFormat(name='bf16', exp_bits=8, man_bits=7, max_override=None), FloatFormat(name='fp16', exp_bits=5, man_bits=10, max_override=None), FloatFormat(name='fp8 E5M2', exp_bits=5, man_bits=2, max_override=None), FloatFormat(name='fp8 E4M3', exp_bits=4, man_bits=3, max_override=448.0)]
def quantize(x, fmt: FloatFormat): on GitHub
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.

def scaled_gradient_roundtrip( grad: float, scale: float, fmt: FloatFormat) -> float: on GitHub
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.

@dataclass
class DynamicLossScaler: on GitHub
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.

DynamicLossScaler( scale: float = 65536.0, growth_interval: int = 2000, clean_steps: int = 0)
scale: float = 65536.0
growth_interval: int = 2000
clean_steps: int = 0
def update(self, grads_finite: bool) -> bool: on GitHub
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.

def accumulate_updates( w0: float, update: float, steps: int, fmt: FloatFormat) -> float: on GitHub
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.

def sharded_global_norm(shards: list[numpy.ndarray]) -> float: on GitHub
2471def sharded_global_norm(shards: list[np.ndarray]) -> float:
2472    """‖g‖ when g is split across GPUs: each sums its own squares, one all-reduce adds the sums, then √."""
2473    local = [float(np.sum(s**2)) for s in shards]  # one number per GPU
2474    return math.sqrt(sum(local))  # the all-reduce of those numbers, then the square root

‖g‖ when g is split across GPUs: each sums its own squares, one all-reduce adds the sums, then √.

def train_through_a_bad_batch( clip: float | None, steps: int = 60, bad_step: int = 30, lr: float = 0.1, seed: int = 0) -> list[float]: on GitHub
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.

def detect_spikes(losses: list[float], window: int = 50, factor: float = 2.0) -> list[int]: on GitHub
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.

def wasted_fraction(interval: float, cost: float, mtbf: float) -> float: on GitHub
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)).

def optimal_checkpoint_interval(cost: float, mtbf: float) -> float: on GitHub
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.

def figures() -> dict: on GitHub
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.

def demo() -> None: on GitHub
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.")