Pre-training LLMs: the loss function, decoded
Transcript: this stretch, timestamped
The first half of Part 2 built cross-entropy out of a coding problem: you optimised a code for distribution q, reality turned out to be p, and H(p,q) = Σᵢ pᵢ·log₂(1/qᵢ) counted the bits you now spend per symbol. P11 showed the 2002 language-trees result as a scruffy empirical instance of that same question. This page is where the series cashes its cheque. We take a language model — a black box with billions of knobs — and derive its training objective from scratch, and the formula that falls out is character-for-character the cross-entropy from the coding story. Grant deliberately arrives at it from the other direction, as "average information per token", so that the coincidence lands as a surprise. Our job here is to write out every step he gestures at on screen, especially the one-hot collapse, and then to convert the resulting number into bytes so the identity is not a slogan but an arithmetic fact. P13 then asks the harder question this page leaves open: why the logarithm, and not some other decreasing function.
Outline, with timestamps
- 14:55 — The turn: from "how different are two languages" to "how different is a model's grasp of language from the real thing, as represented by training data".
- 15:36 — The object being trained: text splits into tokens; the model is a function from a token sequence to a probability distribution over every possible next token.
- 16:36 — Everything downhill is machinery: gradient descent and backpropagation are solved; define a good loss and you are essentially done.
- 17:08 — The actual design question: what function of the model's outputs has "minimising it" mean the same thing as "getting better"?
- 17:40 — The construction: for every prefix, read off the probability the model assigned to the token that really followed. All of them come out of a single forward pass.
- 18:12 — Take −log of each of those probabilities: information content, i.e. surprise. A model that follows the plot is rarely surprised.
- 18:42 — Why the shape is right — the vertical asymptote punishes confident wrongness brutally — and why "right shape" is not yet an argument for the logarithm specifically.
- 19:13 — Units: ML uses natural log, not log base 2. A constant factor, absorbed into the learning rate — but not absorbed if you want to read the loss as bits.
- 19:45 — "That's really all there is to pre-training": average the negative logs over every token in the corpus. Batching and optimisers are engineering, not objective.
- 20:17 — The floor: minimise this and you push the loss down toward the entropy of language itself. And the name nobody explained yet — cross-entropy loss.
The model is a function from prefix to distribution
Strip away the architecture — Grant does exactly this at 15:05, swapping the transformer animation (credited in the video description to Clayton Rabideau) for a featureless black box with a parameter dial on it. What survives the stripping is a type signature. Fix a vocabulary of V tokens — sub-word pieces, roughly 50,000 of them in a typical BPE tokeniser. Then a language model with parameters θ is a function
f_θ : (t₁, t₂, …, t_{k}) ⟼ z ∈ ℝⱽ ("logits", raw unnormalised scores)
q = softmax(z), qᵢ = exp(zᵢ) / Σⱼ exp(zⱼ)
so qᵢ ≥ 0 and Σᵢ qᵢ = 1 (a genuine distribution over the vocabulary)
Two things about the softmax are worth pausing on, because they matter later. First, it is the reason logits are unbounded: the network can output any real numbers it likes and still hand back a legal probability distribution, so the optimiser never has to fight a constraint. Second, softmax is shift-invariant — adding the same constant to every logit leaves q unchanged — which is why implementations subtract the max logit before exponentiating and why you should never write log(softmax(z)) when log_softmax(z) exists.
Now the detail Grant mentions in passing at 17:40 — that "the models are actually specifically designed to give you all of these probabilities on a single pass" — and which is worth spelling out, because it is the whole reason pre-training is economically possible. (This paragraph is my elaboration; Grant states the fact but not the mechanism.) A causal transformer applies an attention mask that forbids position k from attending to any position after k. So the output vector sitting above token k is a prediction made from the prefix t₁…t_k and nothing later. Feed in a window of T tokens and you do not get one training example — you get T of them, one per position, all for a single forward and backward pass. That is the difference between a loss you can evaluate on trillions of tokens and one you cannot.
The one-hot collapse, written out
Here is the step everyone nods past, and the reason this page exists.
At position k the model produces q, a distribution over all V tokens. What is the target? Not another distribution handed to us by an oracle. The corpus gives us exactly one fact: the token that actually came next. Call its index c. As a distribution over the vocabulary, that fact is a one-hot vector — probability 1 on the observed token, 0 on all V−1 others:
p = (0, 0, …, 0, 1, 0, …, 0) pᵢ = 1 if i = c, else 0
↑
position c
Substitute that p into the cross-entropy from the first half of the video and watch what happens:
H(p, q) = Σᵢ pᵢ · log(1/qᵢ)
= Σᵢ pᵢ · (−log qᵢ)
= 0·(−log q₁) + 0·(−log q₂) + … + 1·(−log q_c) + … + 0·(−log q_V)
└──────────────── every term with pᵢ = 0 vanishes ─────────────┘
= −log q_c
Every weight in the weighted sum is zero except one, and that one is 1. A sum over a fifty-thousand-token vocabulary collapses to a single lookup: the negative log of the probability the model assigned to the token that actually appeared. Nothing is approximated here; it is an exact identity, and it is why the loss you implement never looks like a sum over the vocabulary even though it is one.
This is also why the video can build the loss from the "average information per token" story at 18:12 without ever writing H(p,q) on screen. The two descriptions are the same object; the cross-entropy form is simply the one that has not yet been simplified. Grant flags, right at the edge of this chapter, that reciting the collapse is an unsatisfying justification for the name — if the formula evaporates, why invoke it? That objection is the engine of P13, and it is a good one. Hold it.
Averaging: from one token to empirical cross-entropy
One token gives one number. The loss is their mean over the corpus. With N tokens total and c_k the token observed at position k:
L(θ) = (1/N) · Σ_{k=1..N} −log q_θ( c_k | t₁ … t_{k−1} )
Grant is blunt that this is the entire objective, batching and optimiser choice being engineering rather than mathematics:
"And in some sense, that's it. That's really all there is to pre-training. All you're doing is looking through each token, taking the negative log of the probability the model gives to it, and averaging that over every token you see."— Grant Sanderson, 19:45
Now the statistical reading, which the video implies at 20:17 but does not formalise. (The next two paragraphs are standard material I am adding; they are what makes the word "cross-entropy" literally rather than loosely correct.) Treat the corpus as a sample drawn from some true distribution P — the actual statistics of the language you are trying to model. Then the sum above is a Monte Carlo estimate of an expectation:
L(θ) ≈ 𝔼_{x ~ P} [ −log q_θ(x) ] = H(P, q_θ) the cross-entropy of q_θ relative to P
So the empirical loss is an estimator of a true cross-entropy, converging to it as the corpus grows, by the law of large numbers. And cross-entropy decomposes — this is Gibbs' inequality, the key property the first half of the video spent its time on:
H(P, q) = H(P) + D(P‖q) with D(P‖q) ≥ 0, equality iff q = P ⟹ L(θ) ≥ H(P) the loss can never drop below the entropy of language itself
That inequality is the floor Grant means by "hopefully getting it to approach something like the entropy of language". Training does not drive the loss to zero and should not be expected to; it drives the excess term D(P‖q_θ) — the KL divergence, the subject of the video's closing footnote — toward zero. What is left over is irreducible: the genuine unpredictability of text. Every reported pre-training loss you have ever seen is a measurement of "entropy of language plus how wrong my model still is", with no way, from the number alone, to tell you how much of it is which.
Nats, bits, and perplexity
At 19:13 Grant notes that machine learning uses the natural log rather than log base 2, calls it a constant factor, and moves on — correctly, because for optimisation it is genuinely irrelevant. Scaling a loss by a constant scales every gradient by that same constant, and the learning rate absorbs it. (There is also a cleanliness argument: with softmax outputs and a natural log, the gradient of the loss with respect to the logits is exactly ∂L/∂z = q − p. Base 2 gives you the same vector divided by ln 2 — correct, just noisier to write. This gradient identity is my addition, and it is the practical reason frameworks fuse softmax and the log into one op.)
But the constant is not irrelevant for interpretation, and this is the unit bug that makes people misread the next section. A loss in nats becomes a loss in bits by dividing by ln 2 ≈ 0.6931:
bits per token = nats per token / ln 2 = nats × 1.4427
Perplexity is the third dialect: PPL = exp(L) for a loss in nats, equivalently 2^(bits per token). The reading is "the effective number of equally likely choices the model feels it is facing at each step". A perplexity of 7.4 does not mean the model believes there are 7.4 possible next tokens; it means its uncertainty is as costly as a uniform choice among 7.4 options. That reading is exact, because a uniform distribution over n outcomes has entropy log n on the nose.
A useful sanity check sits at the top of this table. A freshly initialised model outputs near-uniform logits, so its loss should start at about ln V. For a 50,257-token vocabulary that is 10.825 nats — and a training run whose first step does not land near that number has a bug, not a bad initialisation.
| loss (nats) | bits / token | perplexity | bits / char (≈4 ch·tok⁻¹) | 1T-token corpus |
|---|---|---|---|---|
| 10.825 = ln V | 15.617 | 50,257 | 3.90 | 1,952 GB |
| 4.0 | 5.771 | 54.6 | 1.44 | 721 GB |
| 3.0 | 4.328 | 20.1 | 1.08 | 541 GB |
| 2.5 | 3.607 | 12.2 | 0.90 | 451 GB |
| 2.0 | 2.885 | 7.39 | 0.72 | 361 GB |
| 1.5 | 2.164 | 4.48 | 0.54 | 271 GB |
The "bits per character" column uses the common rule of thumb that English averages roughly 4 characters per token under a BPE tokeniser — treat it as an order-of-magnitude conversion, not a constant of nature, since it drifts with tokeniser and domain.
The loss is a file size — with the arithmetic
Take the row in bold and do the conversion by hand, because doing it yourself is what turns the claim from a slogan into a fact.
loss = 2.0 nats / token
bits per token = 2.0 / 0.693147 = 2.885390 bits
tokens in corpus = 1 × 10¹²
total bits = 2.885390 × 10¹² = 2.885 Tbit
total bytes = 2.885390 × 10¹² / 8 = 3.6067 × 10¹¹ bytes
≈ 360.7 GB (decimal, 10⁹ bytes)
≈ 335.9 GiB (binary, 2³⁰ bytes)
So: a trillion-token corpus, encoded under a model whose pre-training loss is 2.0 nats, occupies about 361 GB. Not "is analogous to" — occupies, to within a rounding error we can name. The bridge is arithmetic coding, which Part 3 builds and which achieves the information content of a message to within fewer than 2 bits for the entire stream, not per token. Spread over 10¹² tokens that overhead is around 2 × 10⁻¹² bits per token: unmeasurable. Sum-of-information is the file size.
Push the comparison one step further and it gets vivid. At the same ~4 characters per token, that corpus is about 4 × 10¹² characters — call it 4 TB of raw UTF-8. Compressing it to 361 GB is roughly an 11× reduction, against the 3–4× that general-purpose gzip manages on English prose. This is exactly the gap that P11's language-trees experiment was operating in the shallow end of: gzip is a very weak estimator of cross-entropy, so it is remarkable that it recovered anything about linguistic lineage at all. A trained language model is a very strong one.
Two independent yardsticks say the 0.72 bits-per-character figure is in the right neighbourhood, not fantasy. Shannon's 1951 experiments with human subjects guessing the next letter of English put the entropy of printed English somewhere in the range of roughly 0.6 to 1.3 bits per character. And the Hutter Prize — Marcus Hutter's standing competition to losslessly compress enwik9, a fixed 1 GB extract of English Wikipedia — currently sits at a record of 110,793,128 bytes, which is 0.886 bits per character. The premise stated on that page is the thesis of this whole video series in one line: being able to compress well is closely related to intelligence.
Six lines where the collapse is visible
The identity is hardest to disbelieve when you can see it in the code. Both halves below compute the same number; the second is the first with the fused kernel unrolled so the vanished sum is legible.
import torch, torch.nn.functional as F logits = model(tokens[:, :-1]) # (B, T, V) raw scores, one row per position targets = tokens[:, 1:] # (B, T) the token that ACTUALLY came next # The target distribution p at each position is one-hot on targets[b, t]: # H(p, q) = sum_i p_i * (-log q_i) # = 0 * (-log q_0) + ... + 1 * (-log q_c) + ... + 0 * (-log q_V) # = -log q_c <-- every p_i = 0 term is annihilated # F.cross_entropy IS that collapse, fused with the softmax, reported in NATS: loss = F.cross_entropy(logits.reshape(-1, V), targets.reshape(-1)) # The same thing, unrolled, so you can see the sum that is not there: logq = torch.log_softmax(logits, dim=-1) # log q -- (B, T, V) picked = logq.gather(-1, targets.unsqueeze(-1)) # log q_c -- the Sigma already collapsed manual = -picked.mean() # nats per token assert torch.allclose(loss, manual, atol=1e-5) bits_per_token = manual.item() / 0.6931471805599453 gigabytes = bits_per_token * 1e12 / 8 / 1e9 # if the corpus were 1e12 tokens
Note that gather — an index lookup — is standing in for a sum over 50,257 vocabulary entries. That substitution is legal only because of the algebra in the section above, and it is the reason the loss costs essentially nothing to evaluate once you have the logits. Note too that the framework never sees a one-hot vector: it takes an integer index. The one-hot p is a mathematical fiction that exists only to make the formula come out right, which is exactly the state of affairs that makes the "cross-entropy" name feel unearned — and exactly what P13 repairs.
Where people get stuck
"Which distribution goes in which slot?" H(p,q) is not symmetric and the asymmetry is not cosmetic. The data distribution p supplies the weights — the widths of the bars in Grant's diagram — and the model distribution q supplies the code lengths, the heights. Mnemonic: you weight by what reality does, and you pay by what your model believed. Swapping them here is not merely a different number, it is undefined: H(q,p) with a one-hot p asks for log 0 at every token that did not occur, which is every token but one.
"If the formula collapses to −log q, why call it cross-entropy at all?" This is the right objection and Grant raises it himself at the boundary of this chapter. The answer is not in the per-token algebra — there the name really is vacuous. It is in what happens when you average over many occurrences of the same context: the empirical frequencies of the different continuations reappear as genuine weights, and the aggregate loss is a real cross-entropy against a real non-degenerate distribution. That argument, plus a constrained-optimisation proof that the logarithm is forced, is P13.
"The loss went from 3.2 to 3.1. Is that a lot?" In nats, a 3% relative move that looks like noise. Translate: perplexity falls from 24.5 to 22.2, and bits per token from 4.617 to 4.472 — the compressed corpus shrinks by exactly 3.125%. On a 1 TB archive that is 31 GB you no longer have to store. Log-scale losses compress the visual drama out of real gains; converting to bits or perplexity puts it back. This is also why loss curves are traditionally plotted on log axes.
"A model has 2.885 bits per token — but bits are integers." They are not, for the same reason a fair three-sided die carries log₂3 ≈ 1.585 bits. No individual token is written with 2.885 bits; the stream averages that, because arithmetic coding does not assign a codeword per symbol at all — it narrows a single interval over the whole message and emits the binary expansion of a point inside it. Fractional information is the norm, and integer-length codes (Huffman) are the special case that rounds up.
Going deeper, verified
- Language Modeling Is Compression — Delétang, Ruoss, Duquenne, Catt, Genewein, Mattern, Grau-Moya, Wenliang, Aitchison, Orseau, Hutter & Veness (2023) · The paper that makes this page's claim empirical: Chinchilla 70B used as a general-purpose compressor beats PNG on ImageNet patches (43.4% vs 58.5%) and FLAC on LibriSpeech audio (16.4% vs 30.3%). Marcus Hutter is a co-author, which is a pleasing closing of the loop.
- The Hutter Prize for Lossless Compression of Human Knowledge — Marcus Hutter (launched 2006 at 50,000 €; expanded 2020 to 500,000 €) · Compress enwik9, a 1 GB extract of English Wikipedia, below the standing record of 110,793,128 bytes. Scoring is compressor size plus archive size, which is the honest version of "the loss is a file size".
- Scaling Laws for Neural Language Models — Kaplan, McCandlish, Henighan, Brown, Chess, Child, Gray, Radford, Wu & Amodei (2020) · The canonical source for loss-vs-compute power laws. Read every y-axis in it as bits per token and the paper becomes a study of compression returns on capital.
- nanoGPT — Andrej Karpathy · The whole pre-training objective of this page is one call to F.cross_entropy in model.py. Worth reading precisely because there is so little of it. (The transformer animation in this stretch of the video is credited in the description as a nanoGPT animation by Clayton Rabideau.)
- torch.nn.functional.cross_entropy — PyTorch documentation · Confirms the two things people get wrong: it takes raw logits, not probabilities, and it takes integer class indices, not one-hot vectors. The docs also give the general weighted form that distillation uses.
Exercises
- Do the conversion cold — A model reports a pre-training loss of 1.8 nats per token on a 300-billion-token corpus. Compute, without looking back: bits per token, perplexity, the size of the compressed corpus in GB, and the implied bits per character at 4 characters per token. A good answer states the unit convention it used for "GB" and lands near 2.597 bits/token, perplexity ≈ 6.05, ≈ 97.4 GB, ≈ 0.65 bits/char — and notices that this is below the Hutter Prize record's 0.886 bits/char, then explains why that is not a contradiction (different corpus, and the model weights are not counted).
- Prove the collapse to yourself in code — Build a random (4, 8, 1000) logit tensor and random integer targets. Verify that F.cross_entropy equals the manual log_softmax-then-gather version, and also equals the fully explicit version that materialises the one-hot p with F.one_hot and computes −(p * log_softmax(z)).sum(-1).mean(). Then backpropagate and check numerically that ∂L/∂z = (q − p)/(B·T). A good answer reports all three losses agreeing to 1e-6 and explains where the B·T divisor comes from.
- Find the floor empirically — Take any short English text and estimate its entropy two ways: (a) a zeroth-order character model — count letter frequencies and compute H = Σ p·log₂(1/p); (b) gzip the file and report actual bits per character. A good answer lands around 4.0–4.5 bits/char for (a) — depending on how you treat case and punctuation, and comparable to Shannon's letter-frequency figure of 4.14 bits/char for a 27-symbol alphabet — then something like 2–3 bits/char for (b) on a file large enough for gzip's window to matter, notes that both are far above Shannon's 0.6–1.3 bits/char estimate for English, and says precisely what each method fails to model that a transformer does not.