Chapter 4: Pretraining
Understand how LLMs are trained from scratch on massive text corpora. Learn about language modeling objectives (causal vs masked), the datasets and data pipelines that fuel pretraining, the distributed training infrastructure required, and the compute-optimal strategies that determine how to allocate your training budget.
Chapter Overview
Pretraining is the most expensive and foundational phase of building an LLM. During pretraining, the model learns language structure, factual knowledge, reasoning patterns, and coding ability by predicting tokens across trillions of training examples. The cost can range from thousands to hundreds of millions of dollars.
The key insight behind modern pretraining is remarkably simple: train a large transformer to predict the next token in a sequence, and let scale do the rest. This simple objective, applied to enough data with enough parameters, produces models capable of translation, summarization, code generation, mathematical reasoning, and more---none of which were explicitly trained for.
However, the simplicity of the objective belies the complexity of the engineering. Pretraining requires: carefully curated and deduplicated datasets, distributed training across thousands of GPUs, sophisticated optimization strategies, and careful monitoring for training instabilities. This chapter covers all these aspects.
This chapter covers:
- Language Modeling Objectives: How next-token prediction and masked prediction differ
- Causal vs Masked LM: When to use each, and why causal dominates for generation
- Training Data: Where it comes from, how it's cleaned, and why data quality matters more than quantity
- Training Infrastructure: Distributed training, mixed precision, and parallelism strategies
- Scaling Laws & Compute: Optimizing the allocation of your training budget
Chapter Roadmap
Click any topic to jump in
LM Objectives
Cross-entropy loss, perplexity, and teacher forcing — the training objectives that drive next-token prediction in every LLM.
Objectives determine architecture and data
Causal vs Masked LM
Autoregressive vs bidirectional objectives — why causal modeling dominates for generation and how BERT's masking works.
Training Data
Datasets, deduplication, and data mixtures — the fuel that determines what an LLM knows and how well it reasons.
Infrastructure and theory for training at scale
Training Infrastructure
Data parallelism, model parallelism, and mixed precision — the distributed systems engineering behind pretraining.
Scaling Laws
Chinchilla-optimal compute allocation — predicting loss from parameters, data, and FLOPs before training begins.
A transformer block is only an architecture. Before training, its billions of weights are random and it predicts nothing. To learn from raw text with no human labels, we need a task that the text itself can grade, and a single number that says how wrong the model is at every position. The choice of that task decides what the model can later do, and the choice of that number decides how smoothly it learns.
The previous chapter, Transformer Architecture, built the network. This chapter, Pretraining, covers how that network is trained on trillions of tokens, starting with the objective.
We start with cross-entropy loss, the negative log probability of the true next token. Then we turn loss into perplexity, a number that is easier to interpret. Next comes teacher forcing, the trick that lets one forward pass train every position at once, and the exposure bias it creates. We finish by contrasting token-level objectives, used in pretraining, with sequence-level ones, used later in RLHF.
Definition
A language modeling objective is the self-supervised training loss that turns raw text into a learning signal. For an autoregressive model it is the average cross-entropy between the model's predicted distribution over the vocabulary and the actual next token, taken over every position. Minimising it maximises the likelihood the model assigns to the training text, and its exponential is perplexity.
In this topic
Cross-Entropy Loss
At each position the model outputs a probability distribution over the vocabulary, given the earlier tokens and weights . The loss for that position is , the negative log of the probability given to the token that actually came next. Averaging over all positions gives the formula. A confident correct prediction costs almost nothing, and a confident wrong one costs a lot, because grows without bound as goes to 0. Loss is measured in nats with the natural log. Its value depends on the tokenizer, so losses from different vocabularies are not directly comparable.
The loss is a maximum likelihood estimator — minimizing it is equivalent to minimizing . With a vocabulary of tokens and context length , the model must assign probability mass across possible sequences. A loss of 2.0 means the model's average per-token surprise is 2 nats, equivalent to choosing among equally likely tokens. Each 0.1 reduction in loss multiplies the probability assigned to correct tokens by — a 10% improvement in confidence.
A model assigns probability 0.8 to the correct next token. What is the per-token loss?
Perplexity
Perplexity is the exponential of the average cross-entropy, as in the formula. If a model's loss is nats, its perplexity is . It reads as an effective branching factor: a perplexity of 10 means the model is as uncertain, on average, as if it were choosing uniformly among 10 tokens. Lower is better, and a perfect model would reach 1. Because it exponentiates the loss, small loss changes become visible ratios. GPT-3 reached a zero-shot perplexity of about 20 on Penn Treebank. Perplexity depends on the tokenizer and dataset, so compare it only between models that share both.
Perplexity converts nats to an effective vocabulary size. If PPL = 20, the model's uncertainty at each step is equivalent to a uniform distribution over 20 tokens. The key insight: perplexity is multiplicative — reducing PPL from 20 to 10 is the same relative improvement as 10 to 5. This is why we compare log-perplexity (cross-entropy) for arithmetic differences. For a vocabulary of , random guessing gives PPL = 50,000; a well-trained model achieving PPL = 15 has reduced uncertainty by a factor of 3,333.
Model A has perplexity 15 on a test set. Model B has perplexity 10. How much better is B?
Teacher Forcing
During training, the input at every position is the true text, not the model's own earlier predictions. This is teacher forcing. Because every position's input is known in advance, one forward pass with a causal mask computes the loss at all positions in parallel, which is what makes pretraining on trillions of tokens feasible. The cost is a mismatch with generation, where the model feeds on its own samples. It never practised recovering from its own mistakes, so one bad token can push it into text unlike anything it saw. This is called exposure bias, and sampling strategies partly compensate for it.
Teacher forcing computes all positions in parallel: the loss gradient provides a signal at every token simultaneously, giving gradient information per sequence versus for sequence-level methods. The train-test mismatch (exposure bias) grows with generation length: after autoregressive steps, errors compound as in the worst case, where is the single-step error rate. For a 512-token generation, even a 0.1% per-token error rate compounds to ~40% chance of at least one error propagating.
Training sequence: "The cat sat on the mat." If the model predicts "dog" instead of "cat" at position 2, what does it see at position 3?
Token-Level vs Sequence-Level Objectives
A token-level objective scores each next-token prediction separately against the known text, so every position gives a gradient and the signal is dense and low-variance. A sequence-level objective scores a whole generated output, for example with a reward model, a human preference, or a test that runs generated code. That matches what users care about, but the model must first generate text, the score is a single number for many tokens, and the gradient must be estimated with methods such as REINFORCE or PPO. So the standard recipe is token-level pretraining followed by sequence-level fine-tuning with RLHF.
Token-level loss provides gradient signals per sequence (one per position), while sequence-level objectives provide just 1 signal for the entire sequence. The variance of policy gradient estimators scales as for sequence-level rewards (exponential in length), versus for per-token cross-entropy. This is why pretraining uses token-level loss: with and , the search space is — far too large for sequence-level exploration. Sequence-level objectives become practical only after pretraining narrows the distribution to a manageable subspace.
Why not use sequence-level objectives during pretraining?
Theory Exercise
Problem:
A model is pretrained with cross-entropy loss and achieves perplexity 12 on a held-out test set. However, when used as a chatbot, it gives poor responses. Why might a low perplexity not translate to good conversational ability?
Hints:
- Think about what the training data looks like vs what a chatbot should do
- Consider the distribution of text on the internet
- Think about the difference between predicting text and following instructions
Related Problems on PixelBank
Raw text can be turned into a prediction task in more than one way. You can hide the next word and ask the model to guess it from the words before, or hide a few words in the middle and let it use context on both sides. The first matches how text is written and generated. The second gives richer context for understanding each word. Each choice fixes the attention mask, the fraction of tokens that give training signal, and whether the model can generate at all.
The previous topic, Language Modeling Objectives, defined the loss for predicting a token. This topic compares which tokens are predicted and what each prediction is allowed to see.
We start with causal language modeling, the left-to-right objective of GPT and Llama. Then we cover masked language modeling, BERT's fill-in-the-blank objective. Next comes prefix language modeling, a hybrid with a bidirectional prefix and a generated continuation. We finish with why causal models came to dominate, and how much of that was the objective and how much was scale.
Definition
Causal language modeling trains a model to predict each token from the tokens before it, using a triangular attention mask, so the model can generate text. Masked language modeling hides a random subset of tokens, usually 15 percent, and predicts them from context on both sides. Prefix language modeling sees a prefix bidirectionally and predicts the continuation causally.
In this topic
Causal Language Modeling (CLM)
Causal language modeling uses the chain rule of probability, as in the formula: the probability of a sequence is the product of each token's probability given all earlier tokens. The model predicts token from to , and a causal mask stops it from seeing later positions. Every position except the first is a training target, so almost all the text gives signal. The same factorisation gives generation: sample a token, append it, and repeat. GPT, Llama, Mistral, and PaLM are all causal models. The price is that each token's representation sees only its left context.
CLM factorizes using the chain rule — this is exact, not an approximation. The causal mask in self-attention zeroes out entries where : for , giving a lower-triangular attention matrix. This means position attends to tokens, and the total computation is pairwise interactions — exactly half of full bidirectional attention's .
How many prediction tasks does CLM create from a 100-token sequence?
Masked Language Modeling (MLM)
Masked language modeling, introduced with BERT, picks 15 percent of positions at random, the set , and asks the model to recover them from everything else, , as in the formula. There is no causal mask, so each prediction uses context on both sides. Of the chosen positions, 80 percent become a [MASK] token, 10 percent a random token, and 10 percent stay unchanged, so the model cannot rely on seeing [MASK], which never appears in real use. The representations are strong for classification and extraction. But only the masked positions give a loss, and the model cannot generate text naturally.
MLM approximates where typically 15% of tokens are masked. The training signal density is 0.15 tokens per position versus 1.0 for CLM — so MLM needs roughly more data to see the same number of prediction tasks. However, each MLM prediction conditions on bidirectional context ( tokens), while CLM position only sees tokens (average ). This richer context per prediction partially compensates for the lower density.
100-token sequence, 15% masking rate. How many prediction tasks? Compare to CLM.
Prefix Language Modeling
Prefix language modeling splits a sequence into a prefix, to , and a continuation. Inside the prefix the attention is bidirectional, so every prefix token sees the whole prefix. The continuation is generated causally, and the loss covers only continuation tokens, which gives the conditional probability in the formula. That suits tasks with a clear input and output, such as translation or question answering. UL2 and the prefix-LM variants in T5's study use it. Encoder-decoder models such as T5 get a similar effect with separate stacks. Wang et al. (2022) found the best choice depends on whether multitask fine-tuning follows.
Prefix LM splits the sequence into a bidirectional prefix of length and a causal suffix of length . The attention mask is: full attention for positions (bidirectional), causal attention for positions . The total attention computation is — between the of full bidirectional and the of pure causal. The prefix ratio controls the trade-off between understanding and generation.
For the task "Translate English to French: The cat is black. =>", what is the prefix and what is generated?
Why Causal LM Dominates
Masked models led language understanding benchmarks in 2018 and 2019, yet today's large models are almost all causal decoders. Four reasons stand out. Generation is the main use of an LLM, and a causal model generates naturally. Every token is a training target instead of 15 percent. One decoder stack with one objective is simple to scale and to serve with a key-value cache. And at scale, causal models learn tasks from prompts through in-context learning, which removes the need for a task-specific head. Wang et al. (2022) found causal decoders the best zero-shot models after pretraining alone.
The key advantage of causal LM is that generation is trivially parallel-free: each maps directly to a single forward pass position. For MLM, generating text requires iterative refinement — predict masks, fill them in, re-mask, repeat — with no guaranteed convergence. Autoregressive generation has sequential steps but each is a single forward pass. MLM generation needs iterations for full sequence generation, each requiring a full forward pass. Scaling laws also favor causal LM: the same compute budget produces lower perplexity because every position contributes to the loss.
BERT (340M, MLM) was state-of-the-art for NLU in 2018. GPT-3 (175B, CLM) beat it in 2020. Was it the objective or the scale?
Theory Exercise
Problem:
You need a model for two tasks: (A) sentiment analysis of product reviews, (B) generating product descriptions. Would you pretrain with CLM or MLM, and why?
Hints:
- Think about which objective supports both tasks
- Consider the cost of training two separate models
- Think about in-context learning capabilities
Related Problems on PixelBank
A model can only know what its data contains, and web data is mostly unusable as it comes. Common Crawl holds petabytes of pages: navigation menus, spam, machine-generated text, and the same article copied thousands of times. Train on it raw and the model spends capacity memorising duplicates and boilerplate. Filter too hard and you run out of tokens. Much of the difference between strong and weak models of the same size comes from data choices that papers often describe only briefly.
The previous topic, Causal vs Masked LM, settled what the model predicts. This topic covers what it reads, and how raw text becomes a training set of trillions of clean tokens.
We start with the common datasets and what goes into them. Then we walk through a cleaning pipeline, from language detection to quality filtering. Next comes deduplication, with evidence that repeated text causes memorisation. We finish with the data mixture, the proportions of web, code, books, and papers, and with curricula that change those proportions during training.
Definition
LLM pretraining data is a large corpus of text tokens, typically hundreds of billions to trillions, assembled from web crawls, books, code, papers, and reference sources. It passes through language identification, quality filtering, deduplication, and safety filtering, then is sampled according to a mixture of sources. The composition largely determines the model's knowledge, skills, and biases.
In this topic
Common Training Datasets
Most pretraining corpora start from Common Crawl, a public archive of monthly web crawls, and add curated sources. The Pile (Gao et al., 2020) is an 825 GiB English corpus built from 22 subsets, including PubMed, arXiv, GitHub, and books. RefinedWeb (Penedo et al., 2023) showed that carefully filtered web data alone can match curated mixtures, and it extracted five trillion tokens from Common Crawl. RedPajama reproduced the LLaMA recipe in the open, and FineWeb documented its filtering choices. Proprietary models add licensed and synthetic data. A dataset's token count depends on the tokenizer, so cross-paper figures are approximate.
The scale of modern training data is staggering: Llama 2 trained on 2 trillion tokens, roughly tokens. At an average of 4 characters per token and 5 characters per English word, that is approximately 1.6 trillion words — equivalent to reading 16 million books. The dataset composition directly affects model capabilities: code data (typically 5-15% of the mix) is responsible for much of the model's reasoning ability, as code requires precise logical thinking that transfers to natural language tasks.
Llama 2 trains on 2T tokens. If average English word is 1.3 tokens, how many words is that?
Data Cleaning Pipeline
Raw crawl data goes through a sequence of filters, each removing a large fraction. First, text extraction strips HTML, menus, and boilerplate. Language identification then keeps the target languages. Heuristic quality rules drop pages with too few words, too many symbols, or repeated lines, and a classifier may score pages by how much they resemble reference text such as Wikipedia. Deduplication removes copies. PII scrubbing and toxicity filters come last. Each filter trades quantity for quality, and aggressive quality classifiers can remove dialects and minority topics along with the spam.
Data deduplication reduces training compute waste: if fraction of the data is duplicates, then FLOPs are spent learning redundant information where is total compute. Empirically, web crawls contain 30-50% near-duplicates. MinHash with hash functions detects duplicates in time versus the naive pairwise comparison. Quality filtering thresholds create a precision-recall trade-off: strict filtering (high precision) may discard 80% of data but the remaining 20% produces better models than training on all of it.
Starting with 100TB of raw Common Crawl, how much data remains after cleaning?
Data Deduplication
Web text is full of repeats: syndicated news, licence text, templates, and scraped copies. Exact deduplication hashes documents or long substrings. Near-duplicate detection uses MinHash, which estimates the Jaccard similarity of two documents' sets of word n-grams, so pages differing by a few words still match. Lee et al. (2021) found that more than 1 percent of an LM's unprompted output was copied verbatim from training data, and one 61-word sentence appeared over 60,000 times in C4. Deduplicated models emitted memorised text ten times less often and needed fewer steps for the same accuracy. Too loose a similarity threshold can merge genuinely different documents.
Near-duplicate detection using MinHash approximates Jaccard similarity between document shingle sets. With hash functions, the collision probability equals the Jaccard similarity: for all-match, giving exponentially sharper discrimination. For hashes, two documents with have a 99.97% detection rate, while documents with have only a 0.0000...003% false positive rate. The memory requirement is where is the number of documents — feasible for billions of documents.
A news article is copied by 500 websites with minor changes. What happens if we don't deduplicate?
Data Mixture and Curriculum
Training draws from several sources at chosen sampling weights, and those weights shape the skills. The LLaMA mix was 67 percent Common Crawl, 15 percent C4, 4.5 percent each of GitHub, Wikipedia, and books, 2.5 percent arXiv, and 2 percent StackExchange. High-quality sources are often upsampled and seen more than once per run. More code tends to help code generation and structured reasoning, and more math helps arithmetic, but every percent given to one source is taken from another. Some runs use a curriculum, shifting toward cleaner or more specialised data near the end of training. Good weights are usually found with small proxy models.
Data mixing follows a weighted sampling distribution: if domain has weight and tokens, the effective epochs for domain are where is the total training tokens. Setting for upweights smaller high-quality domains (e.g., academic papers) relative to larger noisy domains (e.g., web crawl). The Doremi algorithm optimizes by training a small proxy model first and measuring per-domain loss improvements — essentially solving .
A model trained on 90% web text and 10% code scores 30% on HumanEval. Retraining with 50% code gets 60%. Why?
Theory Exercise
Problem:
You are building an LLM for a medical institution. The model should understand medical literature and generate clinical reports. Design your training data mixture and explain your choices.
Hints:
- Think about the domain-specific knowledge needed
- Consider general language ability vs specialized knowledge
- Think about the risks of medical data in training
One GPU cannot train a large language model. A 70B-parameter model trained with Adam in mixed precision needs about 1.1 TB for its weights, gradients, and optimizer state before any activations. An 80 GB accelerator holds a small fraction of that. Even a model that does fit would take decades to train on one device. The work must be split across thousands of GPUs, and every split adds communication that can leave expensive hardware waiting.
The previous topic, Training Data, prepared trillions of tokens. This topic covers the systems that push those tokens through the model across a cluster, and the memory arithmetic behind every design choice.
We start with data parallelism, which copies the model and splits the batch. Then come tensor and pipeline parallelism, which split the model itself when it no longer fits. Next is mixed precision, which moves most arithmetic to 16-bit formats. We finish with optimizers, where Adam's two extra states per parameter dominate memory and ZeRO-style sharding and lighter optimizers push back.
Definition
Training infrastructure is the combination of parallelism strategies, numeric formats, and optimizer design that lets a model train across many accelerators. Data parallelism replicates the model and splits the batch, tensor parallelism splits individual weight matrices, and pipeline parallelism splits the layers. Mixed precision and sharded optimizer states such as ZeRO cut the memory needed per device.
In this topic
Data Parallelism
In data parallelism every GPU holds a full copy of the model and processes a different slice of the batch, called its micro-batch. After the backward pass, an all-reduce averages the gradients across GPUs, so every copy applies the same update and the copies stay identical. Gradient accumulation adds several micro-batches' gradients before each update, so the effective batch in the formula can exceed memory limits. Throughput scales almost linearly while communication keeps up. The limit is memory, because each GPU still stores the whole model and optimizer state. ZeRO addresses this by sharding those states across the data-parallel GPUs.
Data parallelism replicates the model across GPUs, each processing a micro-batch of samples. The gradient is , synchronized via all-reduce in time where is the number of parameters (independent of for ring all-reduce). The speedup is : for a 7B model with 2048-token sequences, the compute-to-communication ratio is approximately , so communication overhead is negligible.
16 GPUs, micro-batch = 8, gradient accumulation = 4. What is the effective batch size?
Model Parallelism (Tensor & Pipeline)
When the model does not fit on one GPU, it is split. Tensor parallelism (Megatron-LM) splits individual matrices. The first feed-forward matrix is split by columns and the second by rows, so each GPU computes part of the layer and one all-reduce combines the results. That needs fast links, so it usually stays inside one 8-GPU node. Pipeline parallelism assigns groups of layers to different GPUs and passes activations forward. Micro-batches keep the stages busy, but GPUs sit idle while the pipeline fills and drains. Large runs combine tensor, pipeline, and data parallelism.
Tensor parallelism splits each weight matrix across GPUs along one dimension, so each GPU stores parameters. A matrix multiply is split as and requires one all-reduce per layer. Pipeline parallelism assigns entire layers to different GPUs, creating a -stage pipeline with bubble overhead of where is the number of micro-batches. With , the bubble overhead drops below 6% — acceptable for most training runs.
A 70B parameter model in FP16 requires 140GB. An A100 has 80GB. How many GPUs minimum?
Mixed Precision Training
Mixed precision runs the forward and backward passes in a 16-bit format and keeps an FP32 master copy of the weights for the update, as in the formula. Sixteen-bit matrix multiplies run several times faster on tensor cores, and 16-bit activations halve activation memory. The master copy matters because tiny updates, about the learning rate times the gradient, would round to nothing in 16 bits. BF16 keeps FP32's 8 exponent bits and gives up mantissa precision, so it has FP32's range and rarely overflows. FP16 has more precision but a narrow range, so it needs loss scaling to stop small gradients underflowing to zero.
FP16 uses 2 bytes per parameter versus FP32's 4 bytes — halving memory for activations and enabling throughput on Tensor Cores. The risk is that FP16 has range versus FP32's . Loss scaling multiplies the loss by (typically 1024-65536) before backward pass, then divides gradients by after — shifting values into FP16's representable range. BF16 preserves FP32's exponent range (8 bits) while using FP16's total bits (16), sacrificing mantissa precision (7 bits vs 10) for dynamic range — making loss scaling unnecessary.
A 7B model in FP32 uses 28GB for weights. How much with BF16 mixed precision?
Optimizer Choices
AdamW is the default optimizer for LLM pretraining. It keeps two running averages per parameter, as in the formula: the momentum of the gradients and the average of their squares . It updates each weight by , giving each parameter its own step size, and it applies weight decay directly to the weights. The two states cost 8 bytes per parameter in FP32. Adafactor stores factored second moments, Lion keeps only momentum, and 8-bit Adam quantises both states. The schedule is usually a linear warmup followed by cosine decay.
Adam maintains per-parameter first moment and second moment , requiring additional floats of memory (8 bytes per parameter in FP32). For a 70B model, optimizer states alone need 560 GB. AdamW adds weight decay directly to the update (not the gradient), which is equivalent to L2 regularization only for SGD — for Adam, the two differ by a factor of , making AdamW the correct formulation.
AdamW for a 70B model requires how much optimizer state memory?
Theory Exercise
Problem:
You have a cluster of 64 A100 GPUs (80GB each) and need to train a 30B parameter model. Design the parallelism strategy (data parallel, tensor parallel, pipeline parallel dimensions). What is the maximum batch size?
Hints:
- Calculate the memory requirements per GPU first
- Consider that tensor parallelism works best within a node (8 GPUs)
- Pipeline parallelism introduces bubble overhead
Related Problems on PixelBank
A large training run costs millions of dollars, and you get one attempt. Before starting, you must decide how many parameters the model should have and how many tokens to train it on. Both choices use up the same compute budget. Too big a model on too little data wastes compute on parameters that never get trained properly, and too small a model saturates early. Guessing wrong is expensive, and you only find out at the end.
The previous topic, Training Infrastructure, turned a cluster into usable FLOPs. This topic covers how to spend those FLOPs, using empirical scaling laws fitted on small runs and extrapolated to large ones.
We start with the compute budget, the rule that converts parameters and tokens into FLOPs and GPU-hours. Then we predict loss from compute with a power law. Next comes compute-optimal allocation, where Chinchilla revised the earlier advice and put tokens and parameters on equal footing. We finish with inference-aware scaling, which explains why Llama models train on far more data than Chinchilla recommends.
Definition
Scaling laws are empirical power-law relationships between a model's loss and its parameter count , training tokens , and compute FLOPs. They are fitted on many small training runs and used to predict the loss of larger ones. They also give the compute-optimal split of a fixed budget between model size and data, about 20 tokens per parameter in Chinchilla's analysis.
In this topic
The Compute Budget
Training compute is approximately FLOPs, as in the formula, where is the parameter count and the number of training tokens. The forward pass costs about FLOPs per token, one multiply and one add per weight. The backward pass costs about twice that, , because it computes gradients with respect to both activations and weights. That gives 6 FLOPs per parameter per token. The rule ignores attention's sequence-length term, which is small next to for typical context lengths. Real clusters sustain 30 to 50 percent of peak FLOPS, so divide by utilisation when converting to time.
The compute budget is measured in FLOPs: where is parameters and is training tokens. The factor 6 comes from: each token requires a forward pass ( multiply-adds for matrix multiplications) and a backward pass ( multiply-adds — gradient computation plus weight update). For Llama 2 70B trained on 2T tokens: FLOPs. On A100 GPUs at 312 TFLOPS each, this requires GPU-seconds, or ~1000 GPUs for 31 days.
Llama 2 7B trained on 2T tokens. An A100 does 312 TFLOPS. How many GPU-hours?
Loss Prediction from Compute
Kaplan et al. (2020) found that test loss falls as a power law in compute across several orders of magnitude, as in the formula. is a fitted constant, and the exponent is about 0.05, so each 10 times more compute multiplies the reducible part of the loss by . is the irreducible loss, the entropy of the text that no model can remove. Teams fit these curves on runs costing about a thousandth of the target and extrapolate to choose the final configuration. Downstream abilities are less predictable, since some benchmarks jump suddenly as loss falls smoothly.
The scaling law predicts loss as a power law in compute. Empirically, for language models, meaning each 10 increase in compute reduces loss by a factor of — about an 11% reduction. The irreducible loss represents the entropy of natural language itself (approximately 1.0-1.5 nats for English). This power law holds over 6+ orders of magnitude of compute, making it remarkably reliable for planning training runs.
A model trained with FLOPs achieves loss 2.5. Predict loss at FLOPs.
Compute-Optimal Allocation
For a fixed compute budget , a larger model sees fewer tokens. Kaplan et al. (2020) advised putting most extra compute into parameters, and GPT-3 used 175B parameters with only 300B tokens, about 1.7 tokens per parameter. Hoffmann et al. (2022) trained over 400 models and found that parameters and tokens should grow together, each roughly as , which works out to about 20 tokens per parameter. Their 70B Chinchilla, trained on 1.4T tokens, beat the 280B Gopher at the same compute. Epoch AI's replication later refitted the loss formula but supported roughly the same optimal ratio.
Chinchilla's key finding: the optimal allocation splits the compute budget equally between parameters and data in the loss decomposition . Minimizing subject to gives and . With , both and scale as — a 10 compute increase should use more parameters and more data.
Budget: FLOPs. What are the compute-optimal model size and data?
Beyond Chinchilla: Inference-Aware Scaling
Chinchilla minimises training loss for a given training budget, but a deployed model keeps costing compute on every token it serves, about FLOPs per token. If a model will serve trillions of tokens, a smaller model trained on more data can be cheaper overall, even though it reaches a given loss less efficiently during training. LLaMA made this choice explicitly, training a 7B model on 1T tokens. Llama 2 7B saw 2T tokens, about 286 per parameter, and Llama 3 8B saw 15T, about 1,875. Over-training has diminishing returns, but it has not hit a hard wall yet.
Chinchilla-optimal training minimizes loss for a fixed training budget , but total lifetime cost includes inference: where is total inference tokens. When , the optimal strategy shifts toward smaller models trained longer: . For a model serving 1 billion inference tokens per day over 2 years (), a 7B model trained on 4T tokens may be more cost-effective than a 70B model trained on 2T tokens — even if the 70B achieves lower loss.
Model A: 70B params, 1.4T tokens (Chinchilla-optimal). Model B: 7B params, 14T tokens (same compute). Which is better for deployment?
Theory Exercise
Problem:
Your company has a budget of 2/hour and achieve 312 TFLOPS. What is the largest compute-optimal model you can train? How many tokens?
Hints:
- Convert budget to GPU-hours, then to FLOPs
- Use C = 6ND and the Chinchilla ratio D/N = 20
- Remember to account for utilization (typically 40-50% of peak FLOPS)