← Atelier
Lesson 01 · Probability and losses · ~35 min + code

Cross-entropy and KL, from surprise

A language model reads the cat sat on the and gives the word that actually came next a probability of 0.02. How bad is that? We want an answer that adds up over a dataset, has a known best-possible value, and can be differentiated.

That one requirement gets you all of this lesson: surprise, cross-entropy, entropy, KL, the softmax gradient, and why SFT and RL use KL in opposite directions.

1 · One prediction

How surprised should the model be?

Score a single prediction. The model gave probability \(q\) to the word that actually occurred, and we want a penalty \(s(q)\). Two requirements are obvious: \(s(1) = 0\), because certain and right costs nothing, and \(s\) grows as \(q\) shrinks.

The third requirement decides everything. Two independent tokens in a row have probability \(q_1 q_2\). Scoring the pair should give the same total as scoring each token and adding the two scores.

Which penalty satisfies \(s(q_1 q_2) = s(q_1) + s(q_2)\)?

Test it with two tokens at \(q = \tfrac13\) each, so the pair has probability \(\tfrac19\).

  • \(1-q\): the tokens cost \(\tfrac23 + \tfrac23 = \tfrac43\), but the pair costs \(\tfrac89\).
  • \(1/q\): the tokens cost \(3 + 3 = 6\), but the pair costs 9.
  • \((1-q)^2\): fails the same way.

Only a logarithm turns products into sums, \(\log(q_1q_2) = \log q_1 + \log q_2\), and the minus sign makes the penalty positive.

\[ s(q) = -\log_2 q = \log_2 \tfrac{1}{q} \]

With base 2 the unit is bits, and a bit has a literal meaning. Probability \(1/8\) is exactly as unlikely as three fair coins all landing heads, so it scores 3 bits. Drag the point, or use the slider. The hook's 0.02 is where it starts.

Surprise of the observed word
Model A lowers its probability for the correct token from 0.90 to 0.45. Model B lowers it from 0.02 to 0.01. Which change adds more surprise?

Exactly 1 bit each: \(\log_2(0.90/0.45) = \log_2(0.02/0.01) = 1\). Surprise depends on the ratio of probabilities, not the difference, so every halving costs one bit wherever it happens. Check it on the slider: each step of 1 on the log scale is one halving.

That is why log loss punishes confident mistakes. Going from 0.01 to 0.001 costs 3.3 bits, the same as going from 1.0 to 0.1.

2 · A whole dataset

From one prediction to a dataset

Now score a model on data. The context the cat sat on the appears 20 times in a corpus. What came next:

The model gives the same distribution \(q\) every time it sees this context: mat \(\tfrac12\), floor \(\tfrac14\), sofa \(\tfrac18\), moon \(\tfrac18\). So each mat occurrence costs 1 bit, each floor 2 bits, and each sofa or moon 3 bits.

What is this model's average surprise over the 20 occurrences?

\((10\cdot 1 + 5\cdot 2 + 4\cdot 3 + 1\cdot 3)/20 = 35/20 = 1.75\) bits. The average runs over occurrences, not over words, so each word's surprise counts as many times as the word occurs. Averaging the four surprises (2.25) would treat moon as if it were as common as mat.

Here is that computation drawn out. Each column is one occurrence, and its height is that occurrence's surprise. Now change the model and watch which columns move.

20 occurrences, each scored by the model

Grouping identical words turns counts into frequencies, \(10/20 = 0.50\) and so on:

Name the frequencies \(p(x)\) and you have written down cross-entropy, the model's surprise averaged over how often each outcome really happens:

\[ H(p, q) = -\sum_x p(x)\, \log_2 q(x) \]

In training code it is the mean of \(-\log q(y_n)\) over a batch: the same number, computed one example at a time.

3 · The floor

What is the best possible model?

You can move the model freely. How low can the average go?

Which model has the lowest average surprise on these 20 occurrences?
Modelmat ×10floor ×5sofa ×4moon ×1Average
A0.046.646.646.643.34
B1.002.002.324.321.68
C0.512.473.186.641.84
D2.002.002.002.002.00

Surprise per occurrence, in bits.

B. Betting on the most frequent word maximizes accuracy but gives the worst model here: the 10 non-mat occurrences each cost 6.6 bits. Sharpening (C) makes the mat columns cheaper and every other column more expensive, and the other columns lose more than mat gains.

Watch each one in the chart above:

No \(q\) beats \(q = p\). Try it in the chart: you won't get below 1.68 bits. Cross-entropy is minimized by the true frequencies, and the minimum has its own name:

\[ H(p) = H(p, p) = -\sum_x p(x) \log_2 p(x) = 1.68 \text{ bits} \]

Entropy is the average surprise of a model that knows the true frequencies exactly. It isn't zero, because the data itself is uncertain: even a perfect model can't know whether this particular occurrence will be mat or floor. Entropy is the floor under the training loss.

Why q = p is the minimum (two short derivations)

Lagrange. Minimize \(-\sum_x p(x)\log q(x)\) subject to \(\sum_x q(x) = 1\). Setting the derivative of the Lagrangian to zero gives \(-p(x)/q(x) + \lambda = 0\), so \(q(x) = p(x)/\lambda\). Summing over \(x\) gives \(\lambda = 1\), so \(q = p\).

Jensen (this also proves the gap in section 4 is never negative). Log is concave, so

\[ \sum_x p(x)\log\frac{q(x)}{p(x)} \;\le\; \log \sum_x p(x)\frac{q(x)}{p(x)} = \log \sum_x q(x) = \log 1 = 0 \]

Equality needs \(q(x)/p(x)\) to be constant, which means \(q = p\).

4 · The gap

Where does a worse model pay?

The uniform model averages 2.00 bits and the perfect model 1.68. The uniform model is worse on average. Is it worse on every word?

On an occurrence of moon, which model is less surprised: uniform (q = 0.25) or perfect (q = 0.05)?

The uniform model: \(-\log_2 0.25 = 2\) bits against \(-\log_2 0.05 = 4.32\). On moon, and on sofa too, the uniform model beats the perfect one, because it gives them more probability than they deserve. That probability had to come from somewhere: it pays on mat, 2 bits instead of 1, ten times over.

The charts now show the uniform model. Each word has a dashed tick at \(-\log_2 p(x)\), the perfect model's surprise on that word. Hatched means the model pays more than the perfect model; a dashed outline means it pays less.

Model score = floor + gap
extra vs. perfect modelcheaper than perfect modelperfect model's surprise, per wordaverage H(p,q)floor H(p)

Average the difference between the model's surprise and the perfect model's over all occurrences. That average has a name:

\[ \mathrm{KL}(p \,\|\, q) = H(p, q) - H(p) = \sum_x p(x) \log_2 \frac{p(x)}{q(x)} \]

Individual terms can be negative, like moon's. Their \(p\)-weighted sum never is: \(\mathrm{KL}(p\|q) \ge 0\), and it is 0 only when \(q = p\). Probability is conserved, so every word you over-predict forces you to under-predict another, and the words you under-predict are the ones that occur more.

Two consequences you'll use constantly:

Perplexity is the same quantity on another scale: \(2^{H(p,q)}\) in bits (or \(e^{\text{loss}}\) in nats), the number of equally likely options that would be as surprising. Uniform over 4 words gives 4. The perfect model here gives \(2^{1.68} = 3.2\).

5 · Play it

Bet on the next word

Same question, with money on it. Each round you split 20 chips across the four words, a word is drawn, and you are paid 4 × the share you put on it. Ten chips on the winner doubles your bankroll. Two chips cut it to 40%. Zero chips mean you lose everything.

You know mat comes up 50% of the time, floor 25%, sofa 20%, moon 5%. Which split grows your bankroll fastest over many rounds?

10 / 5 / 4 / 1. Money multiplies, so take logs: each round adds \(\log_2(4q)\) bits to your log-bankroll.

\[ \log_2\frac{\text{bankroll}_n}{\text{bankroll}_0} = \sum_{t=1}^{n}\log_2\big(4\,q(x_t)\big) = 2n - \sum_{t=1}^{n} \underbrace{\big(-\log_2 q(x_t)\big)}_{\text{surprise}} \]

Your growth per round is 2 bits minus your surprise, on average \(2 - H(p,q)\). You already know which \(q\) minimizes that: \(q = p\). Anyone betting differently falls behind by exactly \(\mathrm{KL}(p\|q)\) bits per round. All-on-mat goes bankrupt the first time another word comes up. This is Kelly betting (Cover & Thomas, ch. 6), and its score is the loss from section 2.

Now play it. Level 1 checks the answer you just gave. Levels 2 and 3 take away what you know.

6 · Training

How does training find q = p?

Real models don't set \(q\) directly. They output logits \(z\), and \(q = \mathrm{softmax}(z)\). From here on we use the natural log (nats), like PyTorch's cross_entropy: 1 bit = 0.693 nats, and the floor here is \(H(p) = 1.165\) nats.

\[ q_k = \frac{e^{z_k}}{\sum_j e^{z_j}}, \qquad L_n = -\ln q_{y_n} = -z_{y_n} + \ln \sum_j e^{z_j} \]
The model sees one occurrence of mat and takes a gradient step on \(-\ln q_{\text{mat}}\). What happens to the four logits?

Differentiate the second form of \(L_n\). The first term touches only the observed word, and the log-sum-exp gives back softmax:

\[ \frac{\partial L_n}{\partial z_k} = q_k - \mathbb{1}[k = y_n] \]

For mat the gradient is \(q_{\text{mat}} - 1 < 0\), so descent raises it. Every other logit has gradient \(+q_k\), so each goes down, in proportion to how much probability it currently holds. The words compete for a fixed budget.

Average that over all 20 occurrences: 10 mat, 5 floor, 4 sofa, 1 moon. What is the gradient \(\partial L / \partial z\)?

The one-hot vectors \(e_{y_n}\) average to the frequencies \(p\):

\[ \frac{\partial L}{\partial z} = \frac{1}{20}\sum_{n=1}^{20} \big(q - e_{y_n}\big) = q - p \]

The gradient is zero exactly when \(q = p\): the minimum from section 3, now reached from the training side.

Step it, starting from a model that knows nothing (all logits 0, so \(q\) is uniform):

Gradient descent on the 20 occurrences, learning rate η = 1
Which word's probability will take the longest to get close to its true frequency?

Moon. It starts 5× over-predicted (0.25 against 0.05), yet its gradient \(q - p\) is small because both numbers are small, and it carries only 1/20 of the loss. Press "10 steps" twice and compare: mat is within 1% while moon is still about 30% off. The long tail is learned last, and the average loss hides it. That's why you look at loss per slice.

7 · Direction

Forward vs reverse KL: which mistakes are expensive

KL isn't symmetric. The difference matters when the model can't represent \(p\) exactly, which in practice is always.

Target \(p\): the lengths of responses a reference model gives to one prompt. It answers in two styles, short (about 20 tokens) and long (about 80). Model \(q\): a single bell curve, which can't have two peaks. Which bump is "closest" to \(p\) depends on which way round you measure.

\[ \underbrace{\mathrm{KL}(p\|q) = \mathbb{E}_{x\sim p}\Big[\ln\frac{p(x)}{q(x)}\Big]}_{\text{forward: averaged where } p \text{ has mass}} \qquad \underbrace{\mathrm{KL}(q\|p) = \mathbb{E}_{x\sim q}\Big[\ln\frac{q(x)}{p(x)}\Big]}_{\text{reverse: averaged where } q \text{ has mass}} \]
Minimize forward KL(p‖q) over the bell curve's center and width. What do you get?

Forward KL averages \(\ln(p/q)\) over samples from \(p\). If \(q\) is tiny anywhere \(p\) has mass, that log ratio explodes, so \(q\) must cover both styles even if that means putting mass in the valley. For a Gaussian \(q\), the optimum matches \(p\)'s mean and variance: center 50, width 31. Check both directions with the buttons below.

One bell curve q fitted to a two-style target p
p, targetq, model

The lower two panels show where each KL collects its cost: the integrand at each length. Minimize reverse KL once starting from μ = 30 and once from μ = 70.

What the buttons show:

Averaged overExpensive mistakeBehaviorWhere you meet it
Forward KL(p‖q)data \(p\)\(q \approx 0\) where data existscovers every mode, hedgesMLE, pretraining, SFT, distillation on the teacher's distribution
Reverse KL(q‖p)the model's own samples\(q > 0\) where \(p \approx 0\)locks onto one mode, can drop others cheaplyvariational inference; the KL(πθ‖πref) penalty in RLHF-PPO and GRPO

Reverse KL on one mode costs about \(\ln 2 \approx 0.69\) here: \(q\) matches one component of \(p\), which holds half of \(p\)'s mass, so \(\ln(q/p) \approx \ln 2\) everywhere \(q\) lives. Forward KL for the same fit is 8 to 24 nats, because the other style's samples are nearly impossible under \(q\).

The 1-D picture only isolates the asymmetry. An LLM's distribution over sequences is far richer than one bell curve, but the pull is the same. SFT on data containing two answer styles pushes the model to keep both. An RL policy on a reverse-KL leash can settle on one style the reference model used and pay little for dropping the other.

8 · New cases, no hints

Apply it somewhere else

A click model scores a segment whose true click rate is 2%. Model A predicts 4% for every impression; model B predicts 1%. Both are off by 2×. Which has lower log loss?

B (0.1020 nats vs 0.1044; the floor \(H(p)\) is 0.0980). The ratio rule holds for one outcome, but cross-entropy averages over both outcomes. On clicks, A and B are off by the same factor in opposite directions: \(\pm 0.02 \ln 2\). The non-clicks decide it. A gives them 0.96 instead of 0.98, which costs \(0.98\ln(0.98/0.96) = +0.020\). B gives them 0.99, which is cheaper than the truth: \(-0.010\). KL: A 0.0064, B 0.0039.

A 4-class classifier is trained with label smoothing ε = 0.1: the target is 0.925 on the true class and 0.025 on each other class. What is the lowest loss a model can reach on one training example?

Cross-entropy is minimized at \(q = \) target, and the minimum is the target's entropy: \(-(0.925\ln 0.925 + 3\cdot 0.025\ln 0.025) \approx 0.35\). Predicting 1.0 now costs \(0.025 \cdot \ln(1/0) \cdot 3 = \infty\). The smoothed target is the floor, so the model is actively pushed to stay at 92.5% confidence. That's the whole regularization effect.

Your LM assigns exactly 0 probability to a token that appears once in 10,000 validation tokens. What is the validation loss?

\(-\log 0 = \infty\), and one infinite term makes the average infinite. This is forward KL's "cover everything" pressure in its purest form. Softmax never outputs an exact 0, but top-k or nucleus truncation does, so perplexity must be computed on the full untruncated distribution.

Two LMs with different tokenizers report validation loss 2.30 and 2.10 nats per token on the same text. Which one predicts the text better?

Per-token cross-entropy averages over different outcome spaces and different numbers of tokens. Convert to total surprise for the same text: multiply by the token count, then divide by bytes or characters. A tokenizer with coarser tokens has fewer, individually harder predictions.

9 · Implement, unaided

Write it without a framework

Plain editor, no completion, real Python with numpy, hidden tests. Treat it like a screen: state shapes and edge cases out loud before running. Tab indents, Shift+Tab dedents, ⌘/Ctrl+Enter runs. Your code is saved in this browser.

Exercise 1: stable cross-entropy and its gradient

Exercise 2: KL divergence with the zero conventions

10 · Staff-level follow-ups

Answer out loud, then check

Give yourself about a minute per question. Short answer first, then the reason. Reveal only after you've answered.

Why is minimizing cross-entropy the same as maximum likelihood?
  • For i.i.d. data the likelihood is \(\prod_n q(y_n)\); its log is \(\sum_n \log q(y_n)\).
  • Mean cross-entropy is \(-\frac1N \sum_n \log q(y_n)\): negative log-likelihood per example.
  • Equivalently it is \(H(\hat p) + \mathrm{KL}(\hat p\|q)\) for the empirical distribution \(\hat p\), so MLE minimizes forward KL from data to model.
Validation loss plateaued at 1.9 nats/token. Is the model done?
  • Loss \(= H(p) + \mathrm{KL}\). You don't know \(H(p)\), so the plateau could be the entropy floor, or capacity, optimization, or data limits.
  • Evidence: does loss still fall with more parameters or data (scaling trend)? Train-val gap? Loss by slice (rare tokens, domains, long contexts)? Does the plateau move with LR schedule changes?
  • Decide with downstream evals too: equal loss can hide very different behavior on the slices you care about.
Which KL direction does SFT minimize, which one is the RLHF/GRPO penalty, and what behavior does each produce?
  • SFT/MLE: forward \(\mathrm{KL}(p_{\text{data}}\|\pi_\theta)\), an expectation over data. Mode-covering: \(\pi_\theta\) must put mass on everything in the data.
  • RL penalty: reverse \(\mathrm{KL}(\pi_\theta\|\pi_{\text{ref}})\), estimated on the policy's own samples. It punishes going where \(\pi_{\text{ref}}\) has little mass and barely punishes dropping some of \(\pi_{\text{ref}}\)'s modes.
  • The KL-regularized objective \(\mathbb{E}_{\pi}[r] - \beta\,\mathrm{KL}(\pi\|\pi_{\text{ref}})\) has optimum \(\pi^* \propto \pi_{\text{ref}}\, e^{r/\beta}\). DPO starts from this closed form.
How is the KL term computed in GRPO? Be specific about which variant.
  • In DeepSeekMath's GRPO the KL term is subtracted in the objective as \(\beta\,\mathrm{KL}\), not folded into the reward as in InstructGPT-style PPO, where a per-token KL penalty is part of the reward.
  • It's estimated per token on sampled completions with \(\frac{\pi_{\text{ref}}}{\pi_\theta} - \ln\frac{\pi_{\text{ref}}}{\pi_\theta} - 1\) (Schulman's "k3"). It's unbiased for \(\mathrm{KL}(\pi_\theta\|\pi_{\text{ref}})\) under samples from \(\pi_\theta\), and every sample is non-negative.
  • Implementations differ: some later variants (e.g. DAPO) drop the KL term entirely. Name the variant before arguing about behavior.
For a CTR model, when do you insist on log loss rather than AUC?
  • Log loss is a strictly proper scoring rule: in expectation it is minimized only by the true probability (section 3). It measures calibration and ranking together.
  • AUC is unchanged by any monotone transform of the scores, so it measures ranking only.
  • Insist on calibration when the value is used downstream: auctions (pCTR × bid), blending multiple objectives, thresholds, budget pacing. Check with reliability plots or ECE by slice.
Why does distillation train on the teacher's full distribution instead of its argmax?
  • The loss is cross-entropy to soft targets \(p_T\), which is \(\mathrm{KL}(p_T\|q) + H(p_T)\). The gradient on the logits is \(q - p_T\), so every class gets a signal about how plausible it is, not just the top one.
  • Temperature \(T\) raises the teacher's entropy so small probabilities carry signal. The loss is typically scaled by \(T^2\) to keep gradient magnitudes comparable.
  • Forward KL makes the student cover the teacher's modes; on-policy or reverse-KL distillation makes it mode-seeking, which can suit a small student that can't cover everything.

Sources: C. E. Shannon, A Mathematical Theory of Communication (1948); Cover & Thomas, Elements of Information Theory, ch. 2; Shao et al., DeepSeekMath (2024), for GRPO; J. Schulman, Approximating KL divergence (2020); Hinton, Vinyals & Dean, Distilling the Knowledge in a Neural Network (2015).