Pre-training
Pre-training is the step that turns a randomly initialised transformer into a base model. You run next-token prediction over trillions of tokens of filtered text and code, and almost all of the model's knowledge and raw capability comes from it. Interviewers ask about it because it pulls together statistics (the objective and scaling laws), data engineering (the pipeline is now where most of the quality gains come from), numerical optimisation (optimisers, schedules, stability) and cost reasoning (FLOPs, GPU-hours, MFU). This page teaches each piece well enough for you to reason about it out loud and back it up with numbers.
TL;DR: the 8–12 things to be able to say out loud
- Objective: minimise the average negative log-likelihood of the next token (cross-entropy). Perplexity is \(e^{\text{loss}}\). Bits-per-byte normalises by raw bytes, so models with different tokenizers can be compared.
- Why it generalises: predicting text well enough forces the model to compress grammar, facts, reasoning patterns and code semantics. The loss falls smoothly with scale, and downstream skills follow it, though not always smoothly.
- Data is the main lever. Typical pipeline: extract text from HTML → language ID → heuristic filters → deduplication (exact, MinHash-LSH, suffix arrays) → model-based quality classifiers (FineWeb-Edu, DCLM) → toxicity and PII filtering → decontamination against eval sets.
- Mixture and ordering matter. Domain weights are tuned with proxy models (DoReMi, RegMix). The best data is saved for a final annealing / mid-training phase run while the learning rate decays.
- Repetition limit: up to about 4 epochs of the same data costs almost nothing. Beyond roughly 16 epochs, extra passes are nearly worthless (Muennighoff+ 2023).
- Compute: \(C \approx 6ND\) training FLOPs (\(N\) parameters, \(D\) tokens). Chinchilla-optimal is about 20 tokens per parameter. Modern models are deliberately over-trained (100–2000+ tokens per parameter) because inference cost dominates over a model's lifetime.
- Kaplan vs Chinchilla: Kaplan said to grow the model faster than the data (\(N \propto C^{0.73}\)). Chinchilla said to grow them equally (\(N, D \propto C^{0.5}\)). The gap is mostly explained by methodology: how parameters were counted, the LR schedule, and the small scale of Kaplan's runs.
- Optimiser: AdamW is the default. It stores two moments, so with mixed precision you need about 16 bytes per parameter. Schedule: warmup then cosine, or WSD (warmup–stable–decay), which lets you branch off cooldowns at any point. Muon (orthogonalised momentum) is the main 2025–26 challenger and has been used at the trillion-parameter scale (Kimi K2).
- Hyperparameter transfer: μP lets you tune on a small model and reuse the learning rate and init at large width. Scaling-law fits also predict the best LR and batch size as a function of compute.
- Stability toolkit: gradient clipping, z-loss, QK-norm, logit soft-capping, careful init, removing biases, and restarting from a checkpoint while skipping bad batches when a spike hits.
- Stages: pre-train (general web) → mid-train or anneal (code, math, reasoning, synthetic) → long-context extension → post-training (SFT, RL).
- Back-of-envelope: 70B parameters × 15T tokens ≈ 6.3e24 FLOPs. At 40% MFU on H100s that is about 4.4M GPU-hours, or about 3 weeks on 8k GPUs.
1. The objective: next-token prediction
Cross-entropy, the autoregressive factorisation, and teacher forcing
A decoder-only language model defines a distribution over a sequence by the chain rule:
$$p_\theta(x_1,\dots,x_T) = \prod_{t=1}^{T} p_\theta(x_t \mid x_{<t}).$$Pre-training minimises the average negative log-likelihood (NLL) of the training tokens:
$$\mathcal{L}(\theta) = -\frac{1}{T}\sum_{t=1}^{T} \log p_\theta(x_t \mid x_{<t}), \qquad p_\theta(\cdot \mid x_{<t}) = \operatorname{softmax}(W_U h_t).$$This is the cross-entropy between the empirical data distribution and the model. Minimising it is the same as minimising \(\mathrm{KL}(p_{\text{data}} \,\|\, p_\theta)\) plus the data's own entropy, which is a constant. So the loss can never go below the irreducible entropy of text. That floor is the \(E\) term in the scaling laws below. Training uses teacher forcing: with a causal mask, all \(T\) positions of a sequence are predicted in parallel from the ground-truth prefixes. One forward pass therefore gives \(T\) training signals. This is a big reason why pre-training is so efficient on GPUs.
Practical details interviewers sometimes check:
- Packing. Documents are concatenated with an end-of-sequence token and cut into fixed-length windows (e.g. 4k or 8k tokens), so no compute goes to padding. Some recipes add an intra-document attention mask so tokens can't attend across document boundaries. Llama 3 did this, and it matters more at long context Grattafiori+ 2024.
- Loss in nats or bits. Frameworks report natural-log loss (nats per token). Divide by \(\ln 2\) to get bits.
- Multi-token prediction (MTP). Extra heads also predict tokens \(t{+}2, t{+}3,\dots\). This densifies the signal and allows speculative decoding later. DeepSeek-V3 used an MTP auxiliary objective DeepSeek-AI 2024. It is an add-on to next-token prediction, not a replacement.
Perplexity and bits-per-byte
Perplexity is \(\mathrm{PPL} = \exp(\mathcal{L})\): the effective number of equally likely choices per token. A loss of 2.0 nats per token gives PPL ≈ 7.4. Perplexity is tokenizer-dependent. A tokenizer with a larger vocabulary packs more text into each token, so per-token loss goes up even if the model is just as good per character. You cannot compare Llama and GPT perplexities directly.
Bits-per-byte (BPB) fixes this by normalising total information by the raw UTF-8 byte count:
$$\mathrm{BPB} = \frac{\sum_t \text{NLL}_t \;(\text{nats})}{\ln 2 \cdot \#\text{bytes}} = \frac{\mathcal{L}_{\text{token}}}{\ln 2}\cdot\frac{\#\text{tokens}}{\#\text{bytes}}.$$Example: loss 2.0 nats per token with about 4 bytes per token (typical for English with a BPE vocabulary of roughly 100k) gives \(2.0/0.693/4 \approx 0.72\) BPB. The Pile popularised BPB as a tokenizer-agnostic metric Gao+ 2020. BPB is also a direct compression rate: a model at 0.7 BPB, combined with arithmetic coding, would compress that text to about 9% of its raw size.
Prediction is compression, and compression needs understanding. To push down the loss on "The capital of Australia is ___", the model has to store a fact. To push it down on the next line of a proof or a function body, it has to model the logic. Once the easy statistical patterns (spelling, syntax) are learned, most of the remaining reducible loss sits in tokens that need world knowledge, multi-step reasoning or long-range coherence. That is why driving the loss lower keeps producing more general capability, and why the same objective yields translation, coding and arithmetic without any task-specific labels.
Why the objective produces general capabilities, and its limits
- Implicit multitask learning. Web text contains Q&A, translations, tutorials, code with comments and worked problems. Predicting all of it means learning all of those tasks. GPT-3 showed that this yields in-context few-shot learning at scale Brown+ 2020.
- Smooth loss, uneven skills. Validation loss improves as a smooth power law. Individual benchmarks can look flat and then jump (see the emergence debate in §3).
- Limits. The objective rewards imitating the data distribution, including its errors. It has no notion of helpfulness, truthfulness or following instructions. Those come from post-training (A4). It also spends equal effort on every token, whether trivial or crucial. That motivates data selection, weighting tokens by importance, and MTP-style objectives.
"Why is perplexity a poor metric to compare two models?" A strong answer covers three things. (1) It is tokenizer-dependent, so use BPB or a shared tokenizer. (2) It depends on the eval distribution and on contamination. (3) Lower perplexity does not always mean better downstream or instruction-following behaviour, since post-trained models often have worse perplexity on raw web text. Bonus: mention that loss on carefully chosen held-out domains (code, math) predicts downstream ability better than generic web perplexity does.
- Brown+ 2020, GPT-3: the paper that made "scale the objective and get in-context learning" mainstream.
- Gao+ 2020, The Pile: dataset design plus the BPB evaluation convention.
- DeepSeek-V3 report: the multi-token prediction objective at production scale.
2. Data: sources, pipeline, mixtures
Since about 2023 the clearest consensus in open research has been that data quality and composition are worth more than most architecture tweaks. DCLM's controlled experiments found model-based filtering to be the single most important curation step. A 7B model trained on their filtered pool reached 64% MMLU on 2.6T tokens, roughly matching Llama 3 8B on MMLU with about 6.6× less compute Li+ 2024.
Sources
| Source | What it gives the model | Typical issues | Examples |
|---|---|---|---|
| Web crawl (Common Crawl and in-house crawls) | Breadth: world knowledge, many styles and domains. The bulk of tokens. | Boilerplate, spam, SEO junk, near-duplicates, toxicity, PII, machine-translated junk. | Common Crawl (monthly dumps since 2008), RefinedWeb, FineWeb, DCLM-Pool |
| Code | Programming. Widely believed to help structured reasoning and long-range dependencies. | Licences, auto-generated or minified files, secrets, duplication across forks. | GitHub-derived corpora (The Stack family), Q&A sites |
| Math / science | Symbolic reasoning, LaTeX, step-by-step derivations. | HTML extraction often destroys math markup. | arXiv, math-heavy web subsets (e.g. OpenWebMath-style corpora), textbooks |
| Books / long-form | Long coherent context, narrative, rich vocabulary. | Copyright and legal exposure. Limited supply. | Public-domain books (Project Gutenberg) |
| Reference / curated | High fact density. | Small; easy to over-repeat. | Wikipedia, StackExchange, papers |
| Multilingual | Non-English ability, translation. | Low-resource languages have little clean data. Language ID is noisy on short text. | CC subsets filtered by language. Qwen3 covers 119 languages and dialects Yang+ 2025 |
| Synthetic | Targeted skills (reasoning, textbook-style explanation), rephrased web. | Collapse in diversity, inherited errors, cost to generate. | phi "textbook" data, rephrased CC (Nemotron-CC, WRAP) |
The filtering pipeline, step by step
1. Extraction
Common Crawl ships raw HTML (WARC) plus pre-extracted text (WET). WET text keeps lots of menus and boilerplate. FineWeb found that running its own extraction on WARC with trafilatura produced noticeably better models than using WET files, even though it costs more compute Penedo+ 2024. Math and code need special care: naive extraction destroys LaTeX and indentation.
2. Language identification
A fastText classifier assigns a language and a confidence to each document. Documents below a confidence threshold are dropped. FineWeb kept English documents with score above about 0.65 Penedo+ 2024. Multilingual pipelines route each language to its own filters, because English-tuned heuristics such as stop-word lists over-filter other languages.
3. Heuristic filters
These are cheap, document-level rules built from statistics of "good" text. The canonical sets are the Gopher rules Rae+ 2021 and the C4 rules Raffel+ 2019:
- Document length bounds (e.g. 50 to 100k words) and mean word length in a sane range.
- Symbol-to-word ratio (too many
#or…), and fraction of lines starting with bullets or ending in ellipses. - Must contain some common stop words (catches word salad and keyword lists).
- Repetition: fraction of duplicated lines, paragraphs and n-grams inside the document (catches templated spam).
- C4-style: keep only lines ending in terminal punctuation, drop pages with "lorem ipsum" or curly braces, drop pages that hit a bad-word list. The bad-word rule is now known to remove a disproportionate amount of dialectal and minority-community text.
RefinedWeb showed that web data alone, if filtered and deduplicated hard enough, can match curated corpora Penedo+ 2023.
4. Deduplication
Duplicates waste compute, increase memorisation and regurgitation, and can leak eval data. Deduplicating made models better and reduced verbatim memorised output Lee+ 2021. There are three levels:
| Method | Catches | How it works | Cost |
|---|---|---|---|
| Exact (hash) | Byte-identical docs or lines | Hash the normalised document (or each line or paragraph), keep the first occurrence. | O(n), trivial to shard |
| MinHash + LSH | Near-duplicates (templates, boilerplate variants, mirrored pages) | Turn each doc into a set of word n-gram shingles. Compute \(k\) MinHash values. Split them into \(b\) bands of \(r\) rows each. Docs that collide in any band become candidates (then optionally verify Jaccard). | O(n·k), embarrassingly parallel, plus clustering |
| Suffix array (exact substring) | Long repeated spans within otherwise different docs (licence text, quoted articles) | Build a suffix array over the whole corpus and remove any substring of at least ~50 tokens that appears more than once Lee+ 2021. | Memory-heavy, needs a big machine or sharding (reference code) |
MinHash mechanics. For a random hash function \(h\), the probability that two sets have the same minimum hash equals their Jaccard similarity \(s = |A\cap B|/|A\cup B|\). With \(b\) bands of \(r\) rows, the chance that two docs become candidates is
$$P(\text{candidate}) = 1 - (1 - s^{r})^{b},$$which is an S-curve with threshold near \((1/b)^{1/r}\). FineWeb used 5-gram shingles and 112 hashes split into 14 buckets of 8 (threshold ≈ 0.75 similarity) Penedo+ 2024. You tune \(b, r\) to trade false positives against false negatives.
Assuming more dedup is always better. FineWeb found that deduplicating globally across all Common Crawl dumps did worse than deduplicating within each dump. The global pass left older dumps full of low-quality leftovers, while content duplicated across crawls tended to be higher quality Penedo+ 2024. Duplication count is partly a quality signal, so many labs now deduplicate and then deliberately upsample good data.
5. Model-based quality classifiers
This is the biggest single win of the 2024 era. You train a cheap classifier to score "educational" or "high-quality" text and keep the top slice:
- FineWeb-Edu. Llama-3-70B-Instruct rated a few hundred thousand web pages on a 0–5 educational-value scale. A small regression head on text embeddings was trained on those labels, and the corpus was filtered at score ≥ 3. That kept about 1.3T tokens out of FineWeb's ~15T, and gave large gains on knowledge and reasoning benchmarks such as MMLU and ARC Penedo+ 2024 (dataset).
- DCLM-Baseline. A fastText bigram classifier trained with positives from instruction-style data (OpenHermes-2.5 and ELI5 answers) and negatives from random web pages. It keeps roughly the top 10%. Surprisingly, this simple classifier beat fancier options such as perplexity filtering and embedding-based scorers in their tests Li+ 2024.
- Nemotron-CC. Uses an ensemble of classifiers and LLM rephrasing of lower-quality documents instead of discarding them, to keep more unique tokens for long training runs Su+ 2024.
- Llama 3 used Llama-2-based quality classifiers plus separate code and reasoning classifiers Grattafiori+ 2024.
The trade-off: aggressive classifier filtering raises benchmark scores per token but shrinks the pool and its diversity. It can also bias toward "textbook English" and away from conversational, creative or long-tail knowledge. When you plan to train for 15T+ tokens you cannot keep only the top 10%, so labs bucket data by quality and upsample the better buckets.
6. Toxicity, PII, and safety filtering
Toxicity classifiers and URL blocklists (adult sites, known malware or spam domains) remove harmful content. Regex and NER-based scrubbers mask emails, phone numbers, IPs and IDs. Over-filtering for toxicity can hurt the model's ability to recognise toxicity and can remove dialects, so some labs filter lightly and handle behaviour in post-training.
7. Decontamination
Remove training documents that overlap with benchmark test sets, usually by n-gram matching (GPT-3 used 13-gram overlap Brown+ 2020). It is never perfect. Paraphrased, translated or reformatted copies slip through, and benchmark answers spread across the web after release. For interviews, the takeaway is that published scores of models trained on recent web crawls are partly inflated by contamination. Strong evaluation therefore uses held-out, post-cutoff or private sets (A8).
"Design a pre-training data pipeline for a 10T-token run." Interviewers want stages in the right order, with cheap filters first and expensive model-based scoring last, on the survivors. Mention how each stage is tested: train small ablation models (say 1B parameters on 30–50B tokens) on each variant and compare on low-noise early-signal benchmarks. FineWeb and DCLM both work this way. Then mention the scale problem (CPU cost, sharded MinHash, clustering), multilingual handling, keeping provenance so data can be removed for legal or PII reasons, and decontamination. Strong candidates also bring up the tension between quality and quantity for long runs.
Data mixtures and mixture optimisation
After filtering you have domains (web, code, math, books, multilingual and so on) and must pick sampling weights. The weights do not have to match natural proportions. High-value domains are upsampled (seen for several epochs) and low-value ones downsampled. Llama 3 reports a final mix of roughly 50% general knowledge, 25% math and reasoning, 17% code, 8% multilingual tokens Grattafiori+ 2024.
Ways to choose weights:
- Hand-tuning plus ablations. Still the most common. Train small models on candidate mixes and compare.
- DoReMi. (1) Train a small reference model on a default mix. (2) Train a small proxy model with Group-DRO, which keeps raising the weight of domains where the proxy's loss is far above the reference's (high "excess loss" means the domain is learnable but not yet learned). (3) Average the domain weights over training and use them to train the big model. A 280M proxy produced weights that sped up an 8B model's training on The Pile by about 2.6× to baseline accuracy Xie+ 2023.
- RegMix / data-mixing laws. Train many tiny models on random mixtures, fit a regression (or a functional "mixing law") from mixture weights to validation loss, then pick the predicted best mix Liu+ 2024 Ye+ 2024.
Caveat: the best mix depends on the target you optimise, such as average validation loss versus a code benchmark. It also depends on scale and total token budget, because with more tokens small high-quality domains get repeated too often. Mixtures found with small proxies transfer imperfectly. This is an active research area.
Curriculum and annealing
Strict easy-to-hard curricula have mixed evidence for LLMs. What clearly works is end-of-training emphasis on high-quality data: when the learning rate decays toward zero, the model "settles" into whatever it is seeing, so you switch to your best data then.
- Llama 3 annealed on a small set of high-quality code and math while decaying the LR to zero. This lifted GSM8K and MATH noticeably for the 8B model (less for 405B). They also used short annealing runs as a cheap way to measure the value of a new dataset Grattafiori+ 2024.
- MiniCPM put high-quality and SFT-style data into the decay phase of its WSD schedule Hu+ 2024. OLMo 2 formalised a "mid-training" stage with a curated high-quality mix and averaged ("souped") several annealed runs OLMo+ 2024.
- Qwen3 used three stages: about 30T general tokens, then about 5T tokens of STEM, code, reasoning and synthetic data, then a long-context stage at 32k tokens Yang+ 2025.
Synthetic data in pre-training
- Textbook-style generation. phi-1 (1.3B) was trained on filtered code plus GPT-3.5-generated "textbooks" and exercises, and was very strong on HumanEval for its size Gunasekar+ 2023. The phi line kept leaning on synthetic data Abdin+ 2024.
- Rephrasing the web. An instruction model rewrites web documents in styles such as "like Wikipedia" or "as Q&A". Mixing rephrased and real data sped up pre-training by about 3× in WRAP's setup Maini+ 2024. Nemotron-CC and Kimi K2 use rephrasing to get more useful, unique tokens out of limited high-quality data Su+ 2024 Kimi 2025.
- Reasoning traces and code-execution data are increasingly mixed into mid-training, which blurs the line with post-training.
- Risks. Training recursively on model output alone causes "model collapse", where the tails of the distribution vanish Shumailov+ 2023. Collapse is largely avoided when synthetic data is added to real data rather than replacing it Gerstgrasser+ 2024. Other risks: low diversity, the generator's errors and biases, contamination (the generator may have memorised benchmarks), and licensing terms for generator outputs.
How much synthetic data frontier labs put into pre-training is mostly undisclosed and has been rising quickly through 2025–26. Treat specific percentages you hear as uncertain, and check recent technical reports.
Data repetition limits (epochs)
High-quality text is finite, so how much is repeated data worth? Muennighoff+ trained 400 models up to 9B parameters and 900B tokens. With a fixed compute budget, up to about 4 epochs of repeated data gave a negligible loss penalty compared with unique data. With more repetition the value of extra compute falls off, and it becomes nearly worthless around 16+ epochs (their fitted "half-life" of repeated data). They also found that adding code and loosening filters were reasonable ways to stretch scarce data Muennighoff+ 2023. In practice, small premium domains (Wikipedia, textbooks, math) are commonly seen for several epochs, while web data is seen about once.
A repeated token gives a smaller gradient "surprise" each time because the model has partly memorised it. The first few passes still teach generalisable structure. After that, extra passes mostly add memorisation. This is the "data wall" argument for synthetic data, rephrasing and multimodal tokens.
How many tokens modern models are trained on
| Model (year) | Params (total / active) | Pre-training tokens | Tokens per (active) param |
|---|---|---|---|
| GPT-3 (2020) | 175B dense | 300B Brown+ 2020 | ~1.7 |
| Chinchilla (2022) | 70B dense | 1.4T Hoffmann+ 2022 | 20 |
| LLaMA 1 (2023) | 7B–65B | 1.0–1.4T Touvron+ 2023 | ~20–140 |
| Llama 3 / 3.1 (2024) | 8B / 70B / 405B dense | ~15T (405B: 15.6T) Grattafiori+ 2024 | ~1,900 (8B) · ~210 (70B) · ~38 (405B) |
| DeepSeek-V3 (2024) | 671B / 37B MoE | 14.8T DeepSeek-AI 2024 | ~400 per active param |
| Qwen2.5 (2024) | 0.5B–72B | 18T Qwen 2024 | ~250 (72B) |
| Kimi K2 (2025) | 1T / 32B MoE | 15.5T Kimi 2025 | ~480 per active param |
| Qwen3 (2025) | 0.6B–235B (incl. MoE) | ~36T Yang+ 2025 | very high for small models |
| Llama 4 Scout / Maverick (2025) | 109B / 17B and 400B / 17B MoE | ~40T / ~22T multimodal tokens (model card) | >1,000 per active param |
| SmolLM2 1.7B (2025) | 1.7B dense | ~11T Allal+ 2025 | ~6,500 |
The pattern: frontier-scale open models sit in the 15–40T token range, and small models are trained far past Chinchilla-optimal. Closed frontier labs generally do not disclose token counts, so any numbers quoted for them are estimates.
Token counts grew about 10× from 2023 to 2025. Releases after mid-2026 may be well beyond 40T, often counting image, video and audio tokens. Check the latest reports before quoting "the" number.
- FineWeb blog post: the most readable end-to-end account of building a web corpus, with ablations for every step.
- DCLM (Li+ 2024): controlled benchmark for data curation; shows why model-based filtering wins.
- Scaling Data-Constrained LMs (Muennighoff+ 2023): the repetition and epoch results.
- datatrove: production-grade open pipeline code for filtering and dedup.
3. Scaling laws
Training compute: \(C \approx 6ND\)
For a dense transformer with \(N\) non-embedding parameters, the forward pass costs about \(2N\) FLOPs per token: each weight is used in one multiply and one add. The backward pass costs about twice the forward, because it computes gradients with respect to both activations and weights. That gives
$$C_{\text{train}} \approx 6\,N\,D \quad\text{FLOPs},\qquad C_{\text{inference}} \approx 2N \text{ per generated token}.$$This ignores attention-score FLOPs, roughly \(6 \cdot n_{\text{layers}} \cdot T \cdot d_{\text{model}}\) extra per token (counting forward and backward), where \(T\) is the context length. They are a few percent at 4k context for big models but become significant at long context or for small models Kaplan+ 2020. For MoE models use the active parameters per token: DeepSeek-V3 costs like a ~37B dense model per token, not 671B. Activation recomputation (checkpointing) adds another forward pass, about 8ND, but that counts toward hardware utilisation, not model FLOPs (see MFU in §8).
Kaplan et al. 2020
OpenAI fit power laws to models up to about 1.5B parameters Kaplan+ 2020:
$$L(N) \approx \left(\tfrac{N_c}{N}\right)^{\alpha_N},\; L(D) \approx \left(\tfrac{D_c}{D}\right)^{\alpha_D},\; L(C_{\min}) \propto C_{\min}^{-\alpha_C}$$with \(\alpha_N \approx\) 0.076, \(\alpha_D \approx\) 0.095 and \(\alpha_C \approx\) 0.050. Key claims:
- Loss depends mainly on scale (\(N\), \(D\), \(C\)). Shape choices such as depth versus width barely matter within a wide range.
- Big models are more sample-efficient. For a fixed compute budget you should train a very large model and stop well before convergence. The optimal size grows as \(N_{\text{opt}} \propto C^{\approx 0.73}\) and data only as \(D \propto C^{\approx 0.27}\).
This guided the GPT-3 era: 175B parameters trained on only 300B tokens.
Chinchilla (Hoffmann et al. 2022)
DeepMind trained over 400 models from 70M to 16B parameters, with the cosine schedule length matched to each run's token budget. They estimated the compute-optimal allocation three ways: (1) the lower envelope of training curves, (2) IsoFLOP profiles, and (3) fitting a parametric loss Hoffmann+ 2022:
$$L(N, D) = E + \frac{A}{N^{\alpha}} + \frac{B}{D^{\beta}}$$Here \(E\) is the irreducible entropy of text, the second term is the penalty for a finite model, and the third is the penalty for finite data. Their Approach-3 fit was \(E \approx 1.69,\ A \approx 406.4,\ B \approx 410.7,\ \alpha \approx 0.34,\ \beta \approx 0.28\) (values as published; see replication caveat below). Minimising \(L\) subject to \(C = 6ND\) gives
$$N_{\text{opt}} \propto C^{\frac{\beta}{\alpha+\beta}},\qquad D_{\text{opt}} \propto C^{\frac{\alpha}{\alpha+\beta}},$$and all three approaches put both exponents near 0.5: parameters and tokens should grow equally with compute. The memorable rule is about 20 tokens per parameter. Chinchilla (70B, 1.4T tokens) used the same compute as Gopher (280B, 300B tokens) and beat it across the board.
Why Kaplan and Chinchilla disagreed
| Factor | Effect |
|---|---|
| Counting non-embedding parameters only, and fitting at small scale | Pearce & Song show that this alone biases the exponent toward Kaplan's ~0.73. Including embeddings and extrapolating properly recovers roughly Chinchilla Pearce+ 2024. |
| Learning-rate schedule not matched to run length | A cosine schedule set for a long horizon, but evaluated early, makes short runs look worse, which penalises the "more data" side. |
| Last-layer compute, warmup length, optimiser tuning (AdamW β2 at small batch) | Porian+ reproduce Kaplan's result and then remove the gap step by step by fixing these Porian+ 2024. |
Replication caveat: Besiroglu+ refit Chinchilla's Approach 3 from data reconstructed out of the paper's figures. They found the published parametric fit was inconsistent with the other two approaches and had implausibly narrow confidence intervals. Their corrected fit agrees with ~20 tokens per parameter Besiroglu+ 2024. The headline rule stands, but don't over-trust the exact published constants. DeepSeek also found that the optimal ratio depends on data quality: higher-quality data shifts the optimum toward bigger models DeepSeek-AI 2024.
Saying "Chinchilla means you should train with 20 tokens per parameter." Chinchilla-optimal minimises training compute for a target loss. It says nothing about inference cost. Almost no production model is trained at 20:1 any more.
Inference-aware scaling: why modern models are "over-trained"
If a model will serve \(D_{\text{inf}}\) tokens over its lifetime, total cost is about \(6ND_{\text{train}} + 2ND_{\text{inf}}\). A smaller model trained longer reaches the same loss at a somewhat higher training cost but a much lower serving cost. Sardana+ formalised this. With large inference demand, the optimal choice is smaller-and-longer than Chinchilla, and quality kept improving up to extreme ratios of 10,000 tokens per parameter in their runs Sardana+ 2023.
Worked example with the Chinchilla fit (illustrative only):
| Model | Tokens | Tokens/param | \(A/N^\alpha\) | \(B/D^\beta\) | Predicted \(L\) | Train FLOPs |
|---|---|---|---|---|---|---|
| 70B | 1.4T | 20 | 0.084 | 0.163 | ≈1.94 | 5.9e23 |
| 8B | 15T | 1,875 | 0.175 | 0.084 | ≈1.95 | 7.2e23 |
About 1.2× the training compute buys a model that is 8.75× cheaper to serve at about the same loss. That is the Llama 3 8B strategy in one line Grattafiori+ 2024. The fit is known to be imprecise at extreme ratios (Sardana+ found standard fits misestimate the value of extra tokens there), so treat this as a direction rather than a precise prediction.
"You have budget C. How do you choose model size?" A strong answer: (1) start from Chinchilla, \(N \approx \sqrt{C/120}\), which follows from \(C = 6N \cdot 20N\). (2) Shift toward smaller \(N\) and more tokens according to expected inference volume and latency or memory targets. (3) Check that enough unique high-quality data exists, remembering the ~4-epoch rule. (4) Fit your own small-scale scaling laws on your data and architecture, because constants depend on both. (5) Leave margin for the long-context and annealing stages. Example: C = 1e24 FLOPs gives \(N \approx 91\)B and \(D \approx 1.8\)T for Chinchilla-optimal. An inference-heavy product might pick about 20B parameters on about 8T tokens instead.
Emergent abilities: real or a measurement artefact?
Wei+ catalogued tasks where performance stays near chance and then rises sharply past some scale, such as multi-digit arithmetic and some BIG-Bench tasks Wei+ 2022. Schaeffer+ argued that many such jumps come from discontinuous metrics like exact-match accuracy. If per-token accuracy improves smoothly, the probability of getting all \(k\) tokens of an answer right is \(p^k\), which looks like a sudden jump. Under continuous metrics (token edit distance, log-likelihood of the correct answer) the curves become smooth and predictable Schaeffer+ 2023.
Balanced view as of 2026: the underlying capability usually improves smoothly with loss, but usable task success can still cross thresholds sharply. That matters in practice: an agent that succeeds 30% of the time versus 90% is a different product. So "emergence" is partly a metric artefact and partly a genuine property of composite tasks. Labs predict downstream metrics by mapping compute → loss → task metric with a fitted sigmoid, rather than extrapolating accuracy directly. Llama 3 used this two-step method to predict the 405B model's ARC-Challenge score from much smaller runs Grattafiori+ 2024.
Fitting scaling laws and using them for hyperparameters
How it's done in practice:
- Pick an IsoFLOP grid. For about 5–8 compute budgets (e.g. 1e18 to 1e21 FLOPs), train several model sizes at each budget, with the LR schedule set to the run length. Note the minimum-loss model size at each budget, fit a parabola in \(\log N\), and fit a power law of \(N_{\text{opt}}\) against \(C\).
- Or fit the parametric form \(E + A N^{-\alpha} + B D^{-\beta}\) with a robust loss (Chinchilla used Huber on log-loss with L-BFGS). Bootstrap to get confidence intervals. These are often wider than people expect.
- Extrapolate 1–2 orders of magnitude, not 5. GPT-4's final loss was reportedly predicted from runs using about 1,000–10,000× less compute OpenAI 2023.
- WSD makes this cheaper. One long constant-LR run with branched cooldowns produces many token budgets from a single run, instead of one cosine run per budget Hägele+ 2024.
Hyperparameter scaling laws. DeepSeek LLM fit the optimal batch size and learning rate as power laws in compute, roughly \(B_{\text{opt}} \propto C^{0.33}\) and \(\eta_{\text{opt}} \propto C^{-0.125}\) (their fitted exponents). They found a broad near-optimal basin and used the fits to set the 7B and 67B runs DeepSeek-AI 2024. As compute grows, the optimal batch gets larger and the optimal LR gets smaller. This complements μP (§4), which makes the optimal LR invariant to width by construction.
Scaling-law research has moved toward laws for data mixtures, data repetition, precision (FP8/FP4), MoE sparsity, distillation, and post-training or test-time compute. The "pre-training scaling is slowing" debate of 2024–26 is contested, and lab claims are often not independently verifiable. Present the classic results as well established, and any statement about current frontier trends as uncertain.
- Hoffmann+ 2022 (Chinchilla): the three estimation approaches. Read §3.
- Porian+ 2024: step-by-step reconciliation of Kaplan with Chinchilla. Great for understanding methodology traps.
- Sardana+ 2023: inference-aware scaling, i.e. why modern small models train on 10T+ tokens.
- Schaeffer+ 2023: the metric-artefact view of emergence.
4. Optimisation
AdamW: mechanics and memory
For each parameter, with gradient \(g_t\) Kingma+ 2014 Loshchilov+ 2017:
$$m_t = \beta_1 m_{t-1} + (1-\beta_1) g_t,\qquad v_t = \beta_2 v_{t-1} + (1-\beta_2) g_t^2$$ $$\theta_t = \theta_{t-1} - \eta_t\left(\frac{\hat m_t}{\sqrt{\hat v_t}+\epsilon} + \lambda\,\theta_{t-1}\right),\quad \hat m_t = \tfrac{m_t}{1-\beta_1^t},\ \hat v_t = \tfrac{v_t}{1-\beta_2^t}.$$- First moment \(m\): momentum. It smooths noisy minibatch gradients.
- Second moment \(v\): a per-coordinate scale. Dividing by \(\sqrt{v}\) normalises each update to about unit size, so rarely updated or small-gradient parameters (embeddings of rare tokens) still move. This makes Adam far less sensitive to gradient scale than SGD.
- Decoupled weight decay (the "W"): decay is applied directly to the weights rather than added to the gradient. If it is added to the gradient, Adam's normalisation rescales it per coordinate, which is not what you want.
- Typical LLM settings: \(\beta_1 = 0.9\), \(\beta_2 = 0.95\) (lower than the 0.999 default, which helps stability at large scale because \(v\) adapts faster after a gradient-scale shift), \(\epsilon = 10^{-8}\), weight decay 0.1, gradient clip 1.0. Peak LR falls with model size: about \(3\times10^{-4}\) at 7B and about \(1.5\times10^{-4}\) at 70B for Llama 2, and 8e-5 for Llama 3 405B.
Memory. With BF16 mixed precision and Adam, a common accounting is 16 bytes per parameter: BF16 weights (2) + BF16 or FP32 grads (2–4) + FP32 master weights (4) + FP32 \(m\) (4) + FP32 \(v\) (4). For 70B that is about 1.1 TB before activations. No single GPU can hold it, so you need ZeRO/FSDP sharding (A3). Adam's two states are the motivation for memory-lean optimisers such as Adafactor (factored second moment), 8-bit Adam, and Muon (one state).
Learning-rate schedules: warmup, cosine, WSD
- Warmup (linear, typically a few hundred to a few thousand steps). At initialisation Adam's \(v\) estimates are poor and the loss landscape is sharp. A large LR then causes divergence or wasted early steps. Warmup lets statistics settle and the network move into a flatter region Goyal+ 2017.
- Cosine decay to about 10% of peak (or to zero) over the planned run. It is robust and was the standard through 2023. Its drawback: you must fix the total steps in advance, and an intermediate checkpoint is not a finished model because its LR is still high.
- WSD (warmup–stable–decay): hold a constant LR for most of training, then decay quickly (linear, cosine or "1−sqrt") over the last ~10–20%. It matches or beats cosine Hu+ 2024 Hägele+ 2024. Advantages: you can extend training indefinitely, branch cooldowns for scaling-law fits, and put high-quality data into the decay phase (annealing). DeepSeek-V3 used a closely related schedule: constant, then cosine decay, then a short low-LR tail DeepSeek-AI 2024.
A useful mental model is a "river valley". At a high constant LR the iterate moves quickly along the valley floor (real progress) while bouncing between the steep walls (noise that hides progress in the loss). Decaying the LR stops the bouncing and drops the iterate to the valley floor, so the loss falls sharply during cooldown. The progress was already made, and cooldown "cashes it in". That explains why WSD loss curves show a sudden drop at the end, and why a high LR for longer, followed by decay, beats decaying slowly the whole way.
Batch size and the critical batch size
Larger batches average away gradient noise, but past a point they stop reducing the number of steps needed. McCandlish+ model this with the gradient noise scale \(B_{\text{noise}} = \operatorname{tr}(\Sigma)/|G|^2\), the ratio of per-example gradient variance to squared true-gradient norm. Below \(B_{\text{crit}} \approx B_{\text{noise}}\), doubling the batch roughly halves the steps (perfect scaling). Above it, returns diminish and you waste compute McCandlish+ 2018. They describe the trade-off as
$$\left(\frac{S}{S_{\min}} - 1\right)\left(\frac{E}{E_{\min}} - 1\right) = 1,$$where \(S\) is steps and \(E\) is examples processed. Key facts:
- The noise scale grows as loss falls: near the optimum the true gradient is small relative to noise. So the critical batch size rises during training. Hence batch-size ramp-up: Llama 3 405B went from about 4M to 8M to 16M tokens per batch as training progressed (schedule as reported) Grattafiori+ 2024.
- Recent work finds critical batch size scales mainly with data size (training duration) rather than model size Zhang+ 2024.
- Batch size and LR interact: bigger batches allow bigger LRs, up to a limit. Systems constraints (enough data parallelism to keep 10k+ GPUs busy) push batches up. Optimisation efficiency pushes them down. Typical frontier batches are a few million to tens of millions of tokens.
μP and hyperparameter transfer
Under standard parametrisation (SP) the optimal LR shrinks as width grows, so a sweep at 100M parameters doesn't tell you the LR for 10B, and sweeping at 10B is unaffordable. Maximal Update Parametrisation (μP) rescales init, per-layer LR and output multipliers so that every layer's activations and updates stay \(\Theta(1)\) as width \(\to \infty\). The optimal hyperparameters then become (approximately) width-invariant Yang+ 2022.
| Quantity (Adam, width multiplier \(m = d/d_{\text{base}}\)) | Standard param. | μP (simplified) |
|---|---|---|
| Hidden weight init variance | \(\propto 1/\text{fan\_in}\) | \(\propto 1/\text{fan\_in}\) |
| Hidden weight LR | \(\eta\) | \(\eta / m\) |
| Output logits | \(W_U h\) | \(W_U h / m\) (or zero-init readout) |
| Embedding LR | \(\eta\) | \(\eta\) (unchanged) |
| Attention logit scale | \(1/\sqrt{d_h}\) | \(1/d_h\) |
Workflow ("μTransfer"): tune LR, init scale and multipliers on a narrow proxy (e.g. width 256 to 1024), then scale width with the same values. The paper transferred hyperparameters from a ~40M-parameter proxy to GPT-3 6.7B and beat the published model, with tuning costing only about 7% of a pre-training run Yang+ 2022 (mup library). Caveats: classic μP covers width, while depth transfer needs extra residual-branch scaling (an active research line). Transfer across token budget and batch size isn't guaranteed. Many labs combine μP-style parametrisation with empirical scaling-law fits for LR and batch size.
Newer optimisers: Muon, Shampoo, SOAP
Muon (MomentUm Orthogonalized by Newton–Schulz), introduced by Keller Jordan in 2024 (blog post). For each 2-D hidden weight matrix:
$$M_t = \mu M_{t-1} + G_t,\qquad O_t = \text{NS}_5(M_t) \approx U V^{\top}\ \ (\text{where } M_t = U S V^{\top}),\qquad W \leftarrow W - \eta\, O_t$$It replaces the momentum matrix with its nearest semi-orthogonal matrix, so all singular values are set to 1. This uses about 5 Newton–Schulz iterations of matmuls instead of an SVD. The intuition: raw gradient updates for transformer weights are dominated by a few directions (low effective rank). Orthogonalising boosts the "rare" directions, giving a better-conditioned update with steepest descent under the spectral norm. Embeddings, the LM head and gains/biases still use AdamW. Memory: one momentum buffer instead of two.
- Moonshot added weight decay and per-matrix update-RMS matching so that Muon "just works" with AdamW-tuned LRs. They reported about 2× compute efficiency over AdamW in their scaling experiments, and trained Moonlight (16B MoE, 5.7T tokens) with it Liu+ 2025.
- Kimi K2 (1T total parameters) used MuonClip, which is Muon plus "QK-clip": rescale the query and key projection weights when attention logits get too large. It pre-trained on 15.5T tokens with no loss spikes Kimi 2025. Essential AI also reported practical efficiency gains, with Muon staying data-efficient at large batch sizes Essential AI 2025.
- Counterpoint: a careful benchmark with equal tuning effort found that claimed 1.4–2× speedups over well-tuned AdamW shrink with scale, to about 1.1× at 1.2B parameters. The fastest optimisers (Muon, SOAP) were all matrix-preconditioned Wen+ 2025.
Shampoo keeps left and right preconditioners for each weight matrix, \(L = \sum G G^\top\) and \(R = \sum G^\top G\), and updates with \(L^{-1/4} G R^{-1/4}\). That is a Kronecker-factored approximation of full-matrix AdaGrad Gupta+ 2018. It is expensive: matrix inverse roots, computed only every \(k\) steps, plus extra memory. SOAP runs Adam in Shampoo's eigenbasis (updated occasionally). This is more stable, has fewer hyperparameters, and showed meaningful step and wall-clock savings over AdamW at the 360M–660M scale Vyas+ 2024. Muon can be seen as Shampoo with the preconditioner accumulation turned off.
| Optimiser | State per param | Extra compute | Status (late 2026) |
|---|---|---|---|
| AdamW | 2 (m, v) | Negligible | Default for most published frontier and open runs |
| Adafactor / 8-bit Adam | ~1 or quantised | Small | Memory-saving variants; T5/PaLM-era use |
| Lion | 1 (sign momentum) | Negligible | Tried; mixed results at scale |
| Muon / MuonClip | 1 (matrices) + Adam for the rest | ~5 NS matmul iterations per matrix per step (small % of a step) | Used at 1T-parameter MoE scale (Kimi K2). Growing open-source adoption. Gains vs tuned AdamW debated. |
| Shampoo / SOAP | Preconditioner matrices (+ Adam states for SOAP) | Periodic eigendecomposition / inverse root | Strong in benchmarks. Engineering-heavy; less public frontier use. |
Optimiser adoption is one of the fastest-moving areas. Muon's use in 2025–26 open models was rising quickly (several MoE releases and some 2026 technical reports describe Muon-based runs), and closed labs rarely say what they use. Check recent reports before claiming what "everyone" uses.
"Why AdamW and not SGD for transformers?" Talk about the heterogeneous curvature and gradient scales across parameter types (embeddings, LayerNorm gains, attention vs MLP), which per-coordinate normalisation handles; sparse token-embedding gradients; and SGD's known poor performance on transformers. Then mention the costs (2× state memory) and the current challengers. For Muon, be able to explain orthogonalised momentum in one sentence, and why it is applied only to 2-D hidden matrices.
- Keller Jordan, "Muon": the original derivation and the NanoGPT speedrun evidence.
- Yang+ 2022 (Tensor Programs V): μTransfer, with practical tables.
- Hägele+ 2024: WSD/cooldowns and cheaper scaling experiments.
- Wen+ 2025: how to fairly benchmark optimisers, and how big the gains really are.
5. Training stability
Large runs fail in two main ways: loss spikes, where the loss jumps by a lot over tens of steps and sometimes recovers, and divergence, where it never recovers. Both get more frequent with scale, higher LR and lower precision. The 175B OPT logbook is a frank record of dozens of restarts, LR cuts and hardware failures Zhang+ 2022. PaLM reported roughly 20 spikes despite gradient clipping. Its fix was to restart from a checkpoint about 100 steps before the spike and skip a few hundred data batches, which suggests spikes come from specific batches interacting with a particular parameter state Chowdhery+ 2022.
Known failure modes and their fixes
| Failure mode | Mechanism | Mitigation |
|---|---|---|
| Attention-logit growth | \(\|q\|\) and \(\|k\|\) grow, so the logits \(q\cdot k\) blow up. Softmax saturates to one-hot (entropy collapse) and gradients vanish or explode. | QK-norm: LayerNorm/RMSNorm on q and k before the dot product Dehghani+ 2023 Wortsman+ 2023. Or QK-clip (Kimi K2), or logit soft-capping. |
| Output-logit divergence | The softmax normaliser \(\log Z\) drifts far from 0. Large logits make BF16 round-off and gradients unstable. | z-loss: add \(10^{-4}\cdot \log^2 Z\) to push \(\log Z \to 0\) Chowdhery+ 2022. Wortsman+ showed it removes this instability at small-scale proxies. |
| Logit soft-capping | Bound the logits smoothly: \(z \leftarrow c\cdot\tanh(z/c)\). | Gemma 2 capped attention logits at 50 and final logits at 30 Gemma Team 2024. Gemma 3 replaced soft-capping with QK-norm Gemma Team 2025. Soft-capping does not play well with fused attention kernels. |
| Gradient explosions | One bad batch or a sharp region produces a huge update. | Global-norm gradient clipping (typically 1.0). Lower \(\beta_2\) (0.95) so \(v\) reacts faster. Monitor the grad norm as an early warning. |
| Init and residual growth | The residual stream variance grows with depth, so early layers or embeddings dominate. | Pre-norm (RMSNorm). Init std about 0.02 with output projections scaled by \(1/\sqrt{2L}\). Careful embedding scale Takase+ 2023. Remove biases (PaLM). |
| AdamW epsilon / LR-sensitivity at scale | As models grow, gradients shrink. If \(\epsilon\) is comparable to \(\sqrt{v}\), updates get damped unevenly. | Smaller \(\epsilon\) (1e-8 or less). Validate LR sensitivity with small-scale proxies Wortsman+ 2023. |
| Data-induced spikes | Long runs of repeated or garbage tokens (e.g. a broken shard). | Shuffle well, filter, and log which batches you skip. On a spike, roll back and skip. |
OLMo 2 is a good recent open case study. It tracked down spikes and fixed them with a combination of QK-norm, z-loss, changed norm placement, no weight decay on embeddings, and data filtering of repeated n-grams. The result was much smoother loss curves OLMo+ 2024. DeepSeek-V3 and Kimi K2 both report runs with no irrecoverable spikes and no rollbacks DeepSeek-AI 2024 Kimi 2025.
"Your 70B run's loss just spiked at step 120k. What do you do?" A strong answer is a runbook. (1) Check whether it is a hardware or numerics fault (NaNs, one bad rank, ECC errors) or genuine optimisation. (2) Look at grad norm, max attention logit, \(\log Z\) and per-layer activation RMS to localise it. (3) If it recovers within a few hundred steps, keep going and watch. (4) If not, roll back to the last good checkpoint, skip the offending batches, and possibly lower the LR temporarily. (5) For the next run, add QK-norm or z-loss and run small-scale proxies at high LR to reproduce the instability cheaply.
Mixed precision (overview)
- BF16 (8 exponent bits, 7 mantissa bits) has the same range as FP32, so unlike FP16 it needs no loss scaling. Matmuls run in BF16 with FP32 accumulation. Master weights, optimiser states and reductions stay in FP32. It has been the default since about 2021.
- FP8 (E4M3 for activations and weights, E5M2 for gradients) doubles tensor-core throughput on Hopper and later GPUs. It needs careful per-tensor or fine-grained scaling. DeepSeek-V3 trained with FP8 at scale using tile/block-wise scaling and higher-precision accumulation DeepSeek-AI 2024.
- FP4-class formats for training are emerging on Blackwell-generation hardware. They are early and less proven.
Details on precision, memory sharding and parallelism are on the next page (A3: Distributed training).
FP8 and FP4 pre-training recipes changed a lot through 2025–26. Check the latest framework docs and lab reports before claiming what precision a given frontier run used.
- Wortsman+ 2023: reproduce large-scale instabilities in small models. Very practical.
- OPT (Zhang+ 2022) and its logbook: what a troubled 175B run actually looks like.
- OLMo 2: a modern open recipe with stability ablations.
6. Stages of a modern training run
Why train short first, then extend?
Attention cost per token grows with sequence length, and most documents are short anyway. Training at 4k–8k for most tokens is far cheaper. Context is then extended near the end:
- RoPE scaling. Raise the rotary base frequency (Llama 3 used base 500,000), use position interpolation Chen+ 2023, or use YaRN's frequency-dependent interpolation Peng+ 2023. These keep the rotation angles the model saw in training in range at longer positions. See A1 for RoPE.
- Continued training on long data. Llama 3 extended from 8K to 128K in about six stages over roughly 800B tokens, moving to the next stage only once short-context performance had recovered and needle-in-a-haystack retrieval was solved at the current length Grattafiori+ 2024. Qwen3 used hundreds of billions of tokens at 32k Yang+ 2025.
- Data lessons. Code repositories and books are good natural long-document sources. Mixing in high-quality short data keeps short-task ability intact. Evaluate on real downstream long-context tasks, because perplexity and simple needle tests can be misleading Gao+ 2024.
Mid-training
"Mid-training" became a common term in 2024–25 for the phase between general pre-training and post-training. It is a lower-LR, quality-heavy stage that upweights math, code, reasoning traces, synthetic textbooks and sometimes instruction-formatted data. It raises benchmark scores far more per token than general pre-training and prepares the base model for RL. Examples: OLMo 2's mid-training mix plus checkpoint souping OLMo+ 2024, and Qwen3's ~5T-token reasoning stage Yang+ 2025. The trade-off: a heavy instruction-like mix can make "base" model comparisons unfair (benchmark-tuned base models), and contamination risk goes up.
The line between mid-training and post-training has kept shifting, especially as reasoning-focused RL has grown. Terminology varies by lab, so define your terms when you use them in an interview.
7. Checkpointing and evaluation during pre-training
Checkpointing
- What to save: model weights, optimiser states (\(m, v\) are about 2/3 of checkpoint size with Adam), LR scheduler step, data-loader position and shuffle seed, and RNG states. With all of these you get bit-for-bit or at least statistically identical resumption. A 70B checkpoint with Adam states is about 1 TB.
- Why it matters at scale: with thousands of GPUs, failures are routine. Llama 3 405B saw ~419 unexpected interruptions in a 54-day snapshot, most from GPU and HBM faults, which is roughly one every few hours Grattafiori+ 2024.
- How often: trade lost work against checkpoint overhead. The Young/Daly rule of thumb is interval \(\approx \sqrt{2 \cdot t_{\text{ckpt}} \cdot \text{MTBF}}\). Use sharded, asynchronous checkpoints: each rank writes its own shard to fast storage while training continues, then it is copied to object storage. Keep a sparse set of permanent milestones for analysis (Pythia released 154 checkpoints per model for research Biderman+ 2023).
Evaluation during the run
| Signal | What it tells you | Gotchas |
|---|---|---|
| Training loss / grad norm / throughput | Health: spikes, divergence, stragglers, MFU drops | Training loss is in-distribution and noisy. A drop at a data-mix change is not progress. |
| Held-out loss / BPB per domain (web, code, math, multilingual) | Smooth, low-noise progress tracking. Detects regressions on one domain. | Must be decontaminated from training data. Compare models with BPB if tokenizers differ. |
| Few-shot downstream evals (MMLU, HellaSwag, ARC, GSM8K, HumanEval…) | Capability | Noisy at small scale. Pick "early-signal" tasks that are monotonic, low-variance and above chance at small scale, as FineWeb did when choosing ablation benchmarks Penedo+ 2024. Cloze (log-likelihood ranking) vs multiple-choice letter formats behave differently at small scale. |
| Short cooldown evals (WSD) | The "true" quality of an intermediate checkpoint | Under a high LR, checkpoints underestimate final quality. Rankings between runs can flip after decay Wen+ 2025. |
| Checkpoint averaging (EMA/SWA) | Free improvement along the trajectory | Gains similar to a partial cooldown Hägele+ 2024. |
Comparing two runs at intermediate checkpoints, mid-cosine or in the WSD stable phase, and concluding one recipe is better. LR state dominates intermediate loss. Compare at matched schedules, or after a short cooldown.
8. Compute and cost back-of-envelope
The method
- FLOPs: \(C = 6ND\) (dense) or \(6N_{\text{active}}D\) (MoE). Add about 5–15% for attention at long context if needed.
- Effective throughput per GPU: peak dense FLOP/s at your precision × MFU. MFU (model FLOPs utilisation) is the model FLOPs actually achieved divided by peak. It excludes recomputation, unlike HFU (hardware FLOPs utilisation), which counts recompute. Well-tuned large dense runs reach about 35–45% BF16 MFU. Llama 3 405B reported 38–43% on 16k H100s Grattafiori+ 2024.
- GPU-hours = \(C\) / (effective FLOP/s × 3600).
- Wall-clock = GPU-hours / #GPUs, plus 10–20% for failures, restarts, evals and checkpoints.
- Dollars = GPU-hours × $/GPU-hour (owned clusters or long reservations are typically ~$2–3 per H100-hour, versus higher on-demand prices).
Worked example: 70B dense on 15T tokens on H100s at 40% MFU
Sanity checks against reported numbers:
- Llama 3.1 405B: \(6 \times 4.05\text{e}11 \times 1.56\text{e}13 \approx 3.8\text{e}25\) FLOPs. The reported cumulative compute was about 30.8M H100-hours (model card), which implies about 35% effective utilisation once overheads are included. That is consistent with the paper's 38–43% MFU during steady-state training.
- Llama 3.1 70B: the model card lists about 7.0M H100-hours, which implies a lower effective utilisation (about 25%). This shows that overheads, long-context stages and smaller-model parallelism inefficiencies matter. Treat 40% MFU as a good-case assumption.
- DeepSeek-V3: \(6 \times 3.7\text{e}10 \times 1.48\text{e}13 \approx 3.3\text{e}24\) FLOPs in 2.788M H800-hours DeepSeek-AI 2024. That works out to about 330 TFLOP/s effective per GPU, roughly a third of BF16 dense peak (lower relative to FP8 peak). The widely quoted "~$5.6M" figure is just GPU-hours × $2. It excludes research, ablations and prior runs.
Interviewers care that you show the method and state your assumptions: precision, MFU, whether you use active or total parameters for MoE, and overheads. Know the cheat-sheet numbers: 6ND; ~1e15 BF16 FLOP/s per H100 (≈2e15 FP8); ~40% MFU as a good case; ~16 bytes per parameter of training state; a GPU-hour costing a few dollars. Common traps: using sparse peak FLOPs, using total MoE parameters, forgetting that the backward pass is 2× the forward, and quoting the final run cost as the cost of the model.
GPU prices and peak FLOP/s change with each hardware generation (H200, B200/GB200 and later), and rental prices fell substantially through 2024–26. Redo the arithmetic with current peak specs and prices.
9. Small-model training tricks
Distillation-based pre-training
Instead of one-hot next-token targets, the student minimises KL divergence to a larger teacher's full next-token distribution Hinton+ 2015:
$$\mathcal{L}_{\text{KD}} = \sum_t \mathrm{KL}\big(p_{\text{teacher}}(\cdot\mid x_{<t}) \,\big\|\, p_{\text{student}}(\cdot\mid x_{<t})\big) \quad(\text{often mixed with the standard CE}).$$- Why it helps: the teacher's soft distribution carries much more information per token than a single sampled label, including which alternatives are plausible and how uncertain the prediction is. It acts like a lower-variance gradient. Small models trained this way beat same-size models trained on raw data with the same number of tokens.
- Who does it: Gemma 2's 2B and 9B were pre-trained with distillation from a larger model, on far more tokens than compute-optimal Gemma Team 2024. Gemma 3 extended distillation across its family Gemma Team 2025. NVIDIA's Minitron pruned a 15B model (width and depth) and recovered quality with distillation using up to about 40× fewer training tokens than training from scratch Muralidharan+ 2024.
- When it's worth it: Distillation Scaling Laws found distillation beats standard training mainly when the student's compute budget is limited, and when a teacher already exists or will be reused across many students. With enough student compute, supervised training catches up. An overly strong teacher is not always better (the "capacity gap") Busbridge+ 2025.
- Engineering cost: an online teacher adds a forward pass, about \(2N_{\text{teacher}}\) FLOPs per token, which can exceed the student's own training cost. Offline alternatives store the teacher's top-k logits per token. That costs storage, and the truncated distribution loses some signal.
Other tricks that matter for small models
- Heavy over-training on curated data. SmolLM2 (1.7B) trained on about 11T tokens with a multi-stage mix that rebalanced web, code, math and new curated datasets between stages Allal+ 2025. Qwen3 small models saw the full 36T-token pipeline.
- Tied input and output embeddings. With a 128k–256k vocabulary, embeddings can be a large share of a 1B model's parameters. Tying saves memory at a small quality cost.
- Pruning plus distillation from an existing larger model (Minitron, and reportedly Llama 3.2 1B/3B) is cheaper than training from scratch.
- Upcycling. Initialise an MoE from a dense checkpoint by copying the MLP into experts Komatsuzaki+ 2022, to reuse sunk compute.
- Synthetic and textbook-quality data gives a disproportionate benefit at small scale, since small models can't afford to memorise noise (phi series).
- Busbridge+ 2025, Distillation Scaling Laws: when distillation is compute-optimal.
- Minitron (Muralidharan+ 2024): practical pruning + distillation recipe.
- SmolLM2: a fully documented small-model data and training recipe.
- HF Ultra-Scale Playbook: hands-on guide to the systems side (bridges into A3).
Interview question bank
1. Write down the pre-training loss and explain why minimising it produces general capabilities.
\(\mathcal{L} = -\frac{1}{T}\sum_t \log p_\theta(x_t \mid x_{<t})\), the token-level cross-entropy, computed in parallel over all positions with teacher forcing and a causal mask. It equals KL(data‖model) plus the entropy of text, so it is bounded below by that irreducible entropy. Reducing loss beyond surface statistics needs modelling facts, syntax, code semantics and multi-step reasoning, because those determine many hard-to-predict tokens. Web data implicitly contains countless tasks (Q&A, translation, tutorials), so prediction is implicit multitask learning, and in-context learning emerges at scale. Limits: it imitates the data, including errors, and has no notion of helpfulness, which is post-training's job.
2. A model has loss 1.8 nats/token with a tokenizer averaging 4.2 bytes/token. What are its perplexity and bits-per-byte? Why prefer BPB?
PPL = e^1.8 ≈ 6.05. Bits per token = 1.8 / ln 2 ≈ 2.60. BPB = 2.60 / 4.2 ≈ 0.62. BPB normalises by raw bytes, so it is comparable across tokenizers. A model with a bigger vocabulary has higher per-token loss simply because each token carries more text. Perplexity is only comparable with the same tokenizer and the same eval set.
3. Walk through a web-data pipeline for pre-training, in order, and justify the order.
Text extraction from WARC (better than WET boilerplate) → URL blocklists → language ID (fastText, with a threshold) → cheap heuristic filters (Gopher/C4: length, symbol ratios, stop words, repetition) → deduplication (exact hashes, then MinHash-LSH for near-duplicates, optionally suffix-array substring dedup) → model-based quality classifier (FineWeb-Edu/DCLM style) → toxicity and PII scrubbing → decontamination against evals → tokenisation and mixing. Cheap, high-recall filters go first so the expensive steps (dedup clustering, neural classifiers) run on less data. Each stage is validated with small ablation models trained on 10–50B tokens, judged on early-signal benchmarks. Keep provenance so data can be removed later.
4. Explain MinHash-LSH. With 20 bands of 5 rows, what's the candidate probability for two docs with Jaccard 0.8? With 0.5?
Shingle each document into n-grams. For each of \(k = b \cdot r\) hash functions take the minimum hash. P(two docs share a minhash) = Jaccard \(s\). Group the hashes into \(b\) bands of \(r\); a pair is a candidate if any band matches exactly: \(P = 1-(1-s^r)^b\). For s = 0.8: \(0.8^5 = 0.328\), \(1-0.328 = 0.672\), \(0.672^{20} \approx 3.5\times10^{-4}\), so P ≈ 0.9996. For s = 0.5: \(0.5^5 = 0.03125\), \(0.96875^{20} \approx 0.53\), so P ≈ 0.47. The S-curve threshold is near \((1/20)^{1/5} \approx 0.55\). Increase \(r\) to sharpen the curve and raise the threshold, and increase \(b\) to raise recall.
5. What did FineWeb-Edu and DCLM show about quality classifiers, and what are the risks?
Both showed that model-based filtering is the highest-leverage curation step. FineWeb-Edu used Llama-3-70B annotations of educational value to train a small classifier, kept about 1.3T of ~15T tokens, and got big MMLU/ARC gains. DCLM's simple fastText classifier, with instruction-style positives and random web negatives, keeping about the top 10%, produced a 7B model at 64% MMLU on 2.6T tokens, competitive with models trained on far more compute. Risks: reduced diversity, bias toward "textbook" register and against dialects and long-tail knowledge, a pool too small for 15T+ token runs (forcing repetition), and benchmark overfitting, since the classifier's notion of quality may align with the evals. Mitigations: quality buckets with upsampling instead of a hard cut, rephrasing low-quality data instead of dropping it (Nemotron-CC), and broader evals.
6. You have 2T unique high-quality tokens and budget for 8T training tokens. What do you do?
Repeating to 4 epochs costs almost nothing in loss per Muennighoff+ 2023, so 4× repetition of the full 2T is acceptable. Better: repeat the best subsets more (premium sources) and the rest less. Add lower-quality but distinct data (relaxed filters), code (which helps even for natural-language tasks), and multilingual data. Consider rephrasing or synthetic augmentation to create "new" tokens from the same knowledge. Alternatively, reduce training tokens and train a larger model if inference cost allows. Monitor held-out loss for overfitting on repeated domains.
7. Derive \(C \approx 6ND\) and estimate the compute for a 7B model on 2T tokens.
Each parameter takes part in one multiply-accumulate (2 FLOPs) per token in the forward pass, giving 2N per token. The backward pass computes gradients for activations and for weights, about 2× forward (4N). Total ≈ 6N per token, so 6ND overall (ignoring attention-score FLOPs, which are small at short context). 7B × 2T: 6 × 7e9 × 2e12 = 8.4e22 FLOPs. At 400 TFLOP/s effective per H100, that is 2.1e8 GPU-seconds ≈ 58k GPU-hours, about 2.4 days on 1024 GPUs.
8. Contrast Kaplan 2020 and Chinchilla 2022. Why did they disagree?
Kaplan: \(N_{\text{opt}} \propto C^{0.73}\), so grow the model faster than the data and stop early (hence GPT-3's 175B on 300B tokens). Chinchilla: \(N, D \propto C^{0.5}\), about 20 tokens per parameter. Chinchilla 70B/1.4T beat Gopher 280B at the same compute. Reasons for the gap: Kaplan counted only non-embedding parameters and fitted at small scale (Pearce & Song); did not match the cosine schedule to each run's length; and had last-layer compute, warmup and AdamW tuning issues at small scale (Porian+). Later refits (Besiroglu+) corrected Chinchilla's parametric constants but kept the ~20:1 conclusion. The optimal ratio also depends on data quality and architecture, so fit your own.
9. Why is Llama 3 8B trained on 15T tokens (~1,900 tokens/param) if Chinchilla says 20?
Chinchilla minimises training compute only. Lifetime cost includes inference, about 2N FLOPs per served token, and a model that is popular in deployment serves vastly more tokens than it trained on. Loss keeps improving log-linearly past the Chinchilla point, so a small model trained 50–100× longer gets close to a much bigger Chinchilla-optimal model's quality at a small multiple of its training compute, and is far cheaper to serve. Memory and latency constraints (single-GPU or on-device) also fix N. Sardana+ formalised this inference-aware optimum and saw gains up to 10,000 tokens per parameter.
10. Given C = 3e23 FLOPs, what's the Chinchilla-optimal model size and token count?
With D = 20N, C = 6N · 20N = 120N², so N = sqrt(3e23/120) = sqrt(2.5e21) ≈ 5e10, i.e. 50B parameters, and D ≈ 1T tokens. Mention that you'd shift toward a smaller N (e.g. 15–20B on 2.5–3.3T tokens) if inference volume matters, that you'd verify unique-data availability, and that you'd check with your own IsoFLOP fit.
11. Are emergent abilities real?
Partly. Wei+ 2022 documented sharp jumps on some tasks. Schaeffer+ 2023 showed many jumps vanish under continuous metrics: exact-match on a k-token answer behaves like \(p^k\), which looks like a step even when the per-token probability p improves smoothly. The underlying capability is usually smooth in loss and compute, which is why labs predict downstream scores via compute → loss → task-metric fits. Thresholds in usable success are still real, especially for long multi-step tasks where per-step reliability compounds. So it is predictable in principle but can be surprising in practice.
12. Explain AdamW's state and the memory needed to train a 13B model with mixed precision (no sharding).
Adam keeps an EMA of gradients (m) and of squared gradients (v), normalises each coordinate's step by \(\sqrt{v}\), and applies weight decay decoupled from the adaptive scaling. Memory: BF16 weights (2 B) + gradients (2 B) + FP32 master weights (4 B) + m (4 B) + v (4 B) = 16 B/param, so 13B × 16 = 208 GB before activations. That exceeds one 80 GB GPU, so you need ZeRO/FSDP sharding across about 4+ GPUs plus room for activations (with recomputation). Typical hyperparameters: β1 = 0.9, β2 = 0.95, ε = 1e-8, wd = 0.1, clip = 1.0.
13. Why warmup? Compare cosine and WSD schedules.
Warmup: early in training Adam's second-moment estimates are unreliable and curvature is high, so a large LR destabilises training. A linear ramp over hundreds to thousands of steps avoids early divergence. Cosine decays smoothly to ~10% of peak over a pre-set horizon. It is robust, but you must fix the length in advance, and intermediate checkpoints aren't finished models. WSD holds a constant LR and then decays over the last ~10–20%. It matches cosine, allows extending runs, lets you branch cooldowns from one run to get many token budgets (cheap scaling laws), and pairs naturally with annealing on high-quality data. The loss drops sharply during the decay, which is consistent with the "river valley" picture: noise is suppressed and progress already made shows up.
14. What is the critical batch size and how does it shape training at 16k GPUs?
McCandlish+'s gradient noise scale \(B_{\text{noise}} = \operatorname{tr}(\Sigma)/|G|^2\). Below it, doubling the batch halves the steps (near-linear speedup). Above it, you burn more examples for little step reduction. It rises as loss falls, so batch ramp-up schedules are common (Llama 3 405B roughly 4M→8M→16M tokens), and recent work says it grows mainly with training duration. At 16k GPUs, data parallelism wants a huge global batch. If that exceeds the critical batch, you waste tokens. So you use tensor, pipeline and context parallelism to keep the per-replica batch reasonable, and accept some inefficiency early in training.
15. What problem does μP solve and how?
Under standard parametrisation the optimal LR (and init) shifts with width, so small-scale sweeps don't transfer and sweeping at large scale is unaffordable. μP scales initialisation, per-layer learning rates (hidden-layer Adam LR ∝ 1/width) and output multipliers (logits ÷ width, attention 1/d) so that updates and activations stay order-1 as width grows. The optimal hyperparameters then become roughly width-invariant. Tune on a ~10–100M proxy and transfer. Tensor Programs V transferred from a ~40M proxy to 6.7B GPT-3 and beat the original at a fraction of tuning cost. Caveats: depth, batch size and training-duration transfer need extra care. Many labs pair μP with fitted LR and batch-size scaling laws.
16. Explain Muon and its current status.
For each 2-D hidden weight matrix, Muon accumulates momentum, then replaces the momentum matrix with its orthogonalised version \(UV^\top\) (all singular values set to 1), computed with ~5 Newton–Schulz iterations of matmuls. This equalises update magnitude across directions, so dominant directions don't swamp rare but useful ones. It is steepest descent under a spectral-norm geometry. Embeddings, head and norms stay on AdamW. It stores one state per parameter versus two. Status: Moonshot reported ~2× efficiency with weight decay and RMS matching (Moonlight). Kimi K2 (1T MoE, 15.5T tokens) used MuonClip with QK-clip, with no loss spikes. A careful equal-tuning benchmark found gains over tuned AdamW shrink to ~1.1× at 1.2B parameters. So it is promising and in real use, but the size of the advantage at frontier scale is still debated.
17. Name the main stability techniques and what each targets.
Gradient clipping (global norm 1.0) limits single-step blowups. Lower β2 (0.95) makes Adam react faster to gradient-scale shifts. z-loss (\(10^{-4}\log^2 Z\)) keeps the softmax normaliser near 0 and prevents output-logit drift. QK-norm (normalise q and k) prevents attention-logit growth and entropy collapse. Logit soft-capping (\(c\tanh(z/c)\), as in Gemma 2) bounds logits smoothly. Init: pre-norm, small init with depth-scaled output projections, careful embedding scaling, no biases. Operations: roll back and skip batches on spikes (PaLM), filter repetitive data, and monitor grad norm, max logits and activation RMS. Use small-scale proxies at high LR to reproduce instabilities cheaply (Wortsman+).
18. Estimate GPU-hours, wall-clock and cost to train a 70B dense model on 15T tokens on H100s at 40% MFU.
C = 6 × 7e10 × 1.5e13 = 6.3e24 FLOPs. Effective per GPU = 0.4 × 989 TFLOP/s ≈ 3.96e14 FLOP/s. GPU-seconds = 1.59e10, so ≈ 4.4M GPU-hours. On 8,192 GPUs ≈ 540 h ≈ 22.5 days, about 26 days with 15% overhead. At $2–3/hr ≈ $9–13M for the final run only. Memory: 16 B/param ≈ 1.1 TB of weights and optimiser state, so shard with FSDP/ZeRO plus TP. State assumptions: dense BF16 peak, not sparse; MFU realistic at scale; no long-context stage counted. As a sanity check, Llama 3.1 70B's reported ~7M H100-hours implies lower effective utilisation once overheads are included.
19. Describe the stages of a modern pre-training run and why each exists.
(1) General pre-training on 10–40T tokens at 4–8k context with a broad mix. This is where most compute goes and where knowledge is built. (2) Mid-training/annealing: LR decays while the mix shifts to high-quality code, math, reasoning, synthetic and instruction-like data. It gives large benchmark gains per token and prepares the model for RL (Llama 3 annealing, Qwen3's 5T reasoning stage, OLMo 2's mid-training plus souping). (3) Long-context extension: raise the RoPE base or use YaRN and train on long documents mixed with short ones, staged up to 128k+. It is done late because attention cost grows with length. (4) Post-training: SFT, then preference optimisation and RL, for instruction following, reasoning and safety. Throughout: checkpointing, per-domain held-out BPB, and early-signal evals.
20. How would you decide whether a new dataset (say, 50B tokens of math web pages) is worth adding?
Cheap path: take a mid-run checkpoint and run a short annealing experiment, decaying the LR over e.g. 40–100B tokens with a mix that includes the new data versus a control mix. Compare on math evals (GSM8K, MATH), held-out math BPB, and general evals to check nothing regresses. Llama 3 used this method. Alternatively, train small ablation models from scratch with and without it. Check for contamination with math benchmarks (n-gram overlap) before trusting gains. Then decide its mixture weight and epoch count (≤4 epochs if it is small and high quality).
21. When does distillation-based pre-training make sense for a 2B model?
When a strong teacher already exists (its cost is sunk or shared across a model family) and the student's compute budget is limited relative to its size. The soft targets give a richer signal per token, so the student reaches better loss per token (Gemma 2/3, Minitron with pruning). Costs: a teacher forward pass per token (~2N_teacher FLOPs, which can dominate) or storage for offline top-k logits. A teacher far stronger than the student isn't always best (capacity gap). Distillation Scaling Laws show that with enough student compute, plain supervised training catches up. Also consider pruning the teacher to initialise the student.
22. What's the difference between MFU and HFU, and why can a run with high GPU utilisation still have low MFU?
MFU counts only the FLOPs the model mathematically needs (6ND-style) per second, divided by peak. HFU also counts extra work such as activation recomputation. "GPU utilisation" in nvidia-smi just means a kernel was running, not that tensor cores were busy. Low MFU despite busy GPUs comes from recompute, memory-bound kernels (norms, softmax, elementwise), communication not overlapped with compute, pipeline bubbles, small matmuls (small models or heavy TP), MoE load imbalance, and data-loading stalls. Typical good large dense runs reach 35–45% MFU in BF16.
23. Design: you must deliver the best possible 8B base model with 1e24 FLOPs and a 6-month deadline. Outline the plan.
Budget: 1e24 / (6 × 8e9) ≈ 21T tokens, which is far over-trained and intended. Data: build roughly 8–12T unique filtered tokens (FineWeb/DCLM-style classifier buckets, MinHash dedup, decontamination). Upsample premium domains to ≤4 epochs, add code and math, and use rephrasing to stretch high-quality data. Recipe: proven architecture (GQA, RoPE, SwiGLU, RMSNorm, QK-norm, z-loss), AdamW (or Muon if validated in proxies), WSD schedule, LR and batch from μP proxies plus small-scale scaling fits, batch ramp. Stages: ~85% general → ~10% annealing on code/math/reasoning → long-context extension to 64–128k. Process: weekly small ablations, early-signal evals, sharded async checkpoints, a spike runbook. Time: at 40% MFU, 1e24 FLOPs is about 700k H100-hours, i.e. about 15 days on 2k GPUs, which leaves slack for ablations and failures. Optionally distil from an existing larger model if one is available.