One future is a lie: distributions
A recurrent state trained to predict the next observation becomes a belief, and in a deterministic world it is almost exact. The Courtyard has a fork. At a diamond-shaped post the ball goes up or down depending on a difference smaller than the sensor can see, and a model trained to minimise squared error predicts the average of the two: a ball in the middle of the post. This lesson derives why squared error can only give the mean, makes the model output a distribution instead (a mixture of Gaussians whose size is the control), separates the uncertainty a model cannot remove from the uncertainty more data would, checks that the probabilities are honest, and then uses the distribution the only way a model of consequences can be used: by sampling it, and feeding the sample back in. Counting what that costs is the next lesson.
New idea: a world model outputs a distribution over what happens next: squared error gives the mean, which can be a state the world never visits, while a mixture trained by likelihood puts probability where the outcomes are. What it outputs is only useful if it is calibrated, and it is used by drawing samples.
Forces next: A model can now output a distribution and draw samples: several futures, each plausible. Using it means feeding each prediction back in as the next input, again and again, though the model was only ever trained on real inputs. A small error becomes the input of the next step. How fast do errors grow when a model predicts from its own predictions, and how far ahead can it be trusted?
1 · The fork, and what least squares does to it
The fork is a diamond-shaped post placed just after the curtain, its left vertex at x = 3.8 m (the grey diamond in the widget). A launcher fires the ball along the centre line, at the vertex, at a speed between 2.3 and 2.5 m/s that the model is not told (lesson 4's launches all had one speed). Its aim wobbles: the ball starts a height b above the line (below it if b < 0), where b is hidden and Gaussian with a standard deviation of 15 cm, and the sensor reads the ball's height once, with an error of 10 cm. The fork is deterministic (the earlier lessons' jolts are off): the ending is a function of b and the launch speed. Scan b and the map is nearly a step. Once the ball arrives more than about 4.3 cm off the axis (3.6 cm at the fastest launch, 5.4 at the slowest) it goes round the post by one face and, 4 s after launch, is more than a metre above or below the centre line. Closer than that it meets the rounded vertex, and its height at 4 s climbs steeply from one side to the other (lesson 4 quoted 4.4 cm for its 3.2 s). The world is mirror-symmetric, so the ball ends above the centre line exactly when b > 0; "up" will mean that. What the model gets is one reading z = 2.5 + b + noise, and its task is to say where the ball is 4 s after launch (lesson 4 asked at the frame where it came out).
Train a network on 800 launches to minimise squared error and ask what it expects for a reading z = 2.5, exactly on the centre line. It answers a height of 2.48 m, and, from a second output for the horizontal position, a point 0.36 m inside the diamond (measured along an axis from the nearest edge; the centre is 0.6 m inside). The nearest of 300 endings simulated for this reading is 0.66 m from that point. This is the fork's average: a ball that goes through the post.
It is not a bad network. It is the right answer to the question it was asked. For any single prediction c of a random height y given the input z,
E[(y − c)² | z] = Var(y | z) + (c − E[y | z])²
The first term belongs to the world and cannot be changed; the second is the model's, and it is zero only at c = E[y | z]. Squared error is minimised by the conditional mean, whatever shape the distribution has. Take a fork where the ball ends d above or d below the centre with probabilities π and 1 − π. The best prediction is c* = (2π − 1)d with loss 4π(1 − π)d². At π = ½ that is c* = 0, the middle, with a loss of 1.44 m² for d = 1.2 m; always predicting the up branch costs ½(2d)² = 2.88 m². The mean halves the loss of choosing a branch, and it is never right.
2 · Where the randomness comes from
Nothing in the fork's physics is random. The model is uncertain because it is not given b or the launch speed. By Bayes' rule, a Gaussian prior on b (standard deviation 15 cm) and a Gaussian reading with noise 10 cm combine by adding precisions: 1/0.15² + 1/0.1² = 144 m⁻². Given a reading that is z′ above the centre line, b is Gaussian with mean 0.69·z′ and standard deviation 8.3 cm, about as large as the 8.6 cm width (±4.3 cm at 4 s) of the band of aims that leave the ending undecided. One reading rarely settles it. The ball ends above the line exactly when b > 0, so the posterior probability of "up" is Φ(0.69·z′ / 0.083 m), with Φ the standard normal distribution function, and for 61 % of launches it lies between 10 % and 90 %. Average ten readings and the posterior narrows to 3.1 cm and the share left ambiguous falls to 21 %. It does not reach zero: b = 0 is a knife edge, and however narrow the posterior, the launches within a posterior width of it stay undecided.
This is one of two kinds of uncertainty, and models mix them up constantly.
| Aleatoric | Epistemic | |
|---|---|---|
| what it is | what is still unknown given the inputs the model has | what the model does not know because it has seen too little |
| in the fork | the spread of the ending given one reading | how much two networks trained on different launches disagree |
| more data | does not shrink it | shrinks it |
| a better input | shrinks it, partly | does not shrink it |
3 · Output a distribution
If the model has to say what may happen, its answer is a distribution p(y | z). A family wide enough for a fork is a mixture of Gaussians,
p(y | z) = Σk=1..K πk(z) · N(y; μk(z), σk(z)²)
and the network outputs 3K numbers: K logits (a softmax makes the weights πk), K means and K log standard deviations (so that σk > 0). It is trained by maximum likelihood: minimise the average negative log-likelihood (NLL) −ln p(y | z) over the data. Squared error is a special case: for one Gaussian with a fixed variance, −ln N(y; μ, σ²) = (y − μ)²/2σ² + const. Adding a learned variance (K = 1) adds an error bar and changes nothing else; its mean is still the conditional mean.
Likelihood is the right loss for a second reason: it is a proper scoring rule. Its expected value when the data come from p and the model is q is H(p) + KL(p‖q), so it is minimised exactly when q = p, in the whole shape and not only in the mean. The gradient is simple. With the responsibility of component k for a sample, rk = πkNk / Σj πjNj,
∂NLL/∂logitk = πk − rk, ∂NLL/∂μk = rk(μk − y)/σk², ∂NLL/∂ln σk = rk(1 − (y − μk)²/σk²)
Each sample pulls on every component in proportion to how much that component explains it: that is the code in the box below. On 2000 held-out launches (heights in units of 1.2 m, the NLL in nats) the NLL is 0.70 for K = 1, −0.01 for 2, −0.15 for 3 and −0.17 for 4. The first step is large: the density grows a second hump and the model stops averaging the branches. The second is the in-between endings of the glancing hits: at the centre reading one of the three components sits level with the post with weight 0.40, and a fraction 0.41 of the launches that give this reading have an aim inside the band around the axis. The fourth buys 0.02. The number of components the data justify is a count of the qualitatively different outcomes of the collision. Other ways to represent several futures (a histogram over binned heights, discrete latent variables, generative samplers) trade differently.
4 · Two checks an honest forecast must pass
Does an ensemble see the fork? An ensemble of M networks, each trained on a bootstrap resample of the data, is a standard estimate of epistemic uncertainty: where the data fix the function the members agree, and where they do not they differ (PETS, Chua et al. 2018, reads aleatoric uncertainty from a predicted variance and epistemic uncertainty from the ensemble's disagreement). Five squared-error networks, each trained on a bootstrap resample of the 800 launches, disagree, at the centre of the fork, by a standard deviation of 0.09 m. The real endings given that reading have a standard deviation of 0.96 m, which is 11 times more. Train the members on only 40 launches and they disagree by 0.25 m, which is epistemic and falls as data arrive; the real spread is still 3.8 times larger. The ambiguity is not in the function the networks are learning, it is in the data, so disagreement between point predictors is the wrong detector for a fork. The right one is a model whose output is a distribution, with an ensemble of those on top to add what the data have not yet settled.
Are the probabilities calibrated? A forecast is calibrated if events it gives probability p happen a fraction p of the time. Take the event "the ball ends above the centre line" (y > 2.5). A mixture assigns it P = Σk πk·P(N(μk, σk²) > 2.5). Sort the held-out launches into eight bins by that P and compare, bin by bin, the mean predicted probability with the observed frequency (a reliability diagram). The launch-weighted average gap (the expected calibration error) is 4.4 points for K = 1, 1.8 for 2 and 3.2 for 3; a perfectly calibrated forecaster would show 1.8 from sampling noise alone on 2000 launches, so K = 2 is as calibrated as the sample can tell. Calibration on one event is necessary and not sufficient. The single Gaussian is within a few points on this event and still places 49 % of its probability, at the centre reading, on heights level with the post (within 0.55 m of the centre line), where 22 % of the real endings are (19 % of the 300 drawn in the widget). The mixtures put 25 % (K = 2) and 27 % (K = 3) there, about as many as the world does: those real endings are the balls that stalled at the vertex. Checking an event's probability and checking the shape of the distribution catch different failures.
The widget
What to try. Leave the reading at 0 and K = 1. The red cross sits at 2.48 m, inside the diamond, while the cyan dots form two clouds above and below it with a thinner scatter between; the real endings level with the post are 19 % of the 300 dots (22 % in the long run), and the single Gaussian puts 49 % of its mass there. Raise K to 2: the amber curve grows two humps, the mass level with the post falls to 25 % and the NLL from 0.70 to −0.01. At K = 3 the amber curve fills in the middle: one component sits there with weight 0.40 (NLL −0.15), and K = 4 gains almost nothing (−0.17). At the centre reading the up, down and level-with-the-post shares of the 300 dots are 39 %, 42 % and 19 %; in the long run they are 39, 39 and 22 %. Now move the reading to +20 cm: the weight moves almost entirely to the up cluster; the three-component model says 94 % that the ball ends above the line, where the exact answer, Φ(0.69 × 0.2 / 0.083) from the posterior of §2, is 95 %. Switch the ensemble to 800 launches: the five purple ticks have a standard deviation of 0.09 m at reading 0, against a real spread of 0.96 m; switch it to 40 launches and the ticks scatter to 0.25 m. The samples-off-the-data readout is the average over readings of the share of samples farther than 0.1 m from every training outcome: 10.7 % for K = 1, 2.0 % for 2, 0.5 % for 3 and 0.2 % for 4.
5 · Using a distribution means sampling it
A model that outputs a distribution is used by drawing from it: one sample is one imagined future, many samples are many futures, and a rollout draws a sample, treats it as the next state, and draws again. That makes the probability of a bad sample the quantity that matters. Call a sample off the data if no training outcome has a height within 0.1 m of its height. Real endings, drawn afresh from the same readings, are off the data 0 % of the time; the mixture heads are not (averaged over seven readings from −30 to +30 cm): 10.7 % for K = 1, 2.0 % for 2, 0.5 % for 3. For the single Gaussian they are mostly its tails, at heights the 800 training launches never reached (below 1.1 m or above 3.9 m). A small number per draw is not a small number for a chain. If each self-fed step lands off the data with the same probability q, independently of the others, the chance that at least one of k steps has done so is 1 − (1 − q)k. For the three-component model that is 4.6 % after 10 steps and 13 % after 30; for the single Gaussian it is 68 % after 10. Real systems expose a knob for this: the world model of Ha and Schmidhuber (2018) predicts the next latent state with a mixture of five Gaussians and scales the sampling uncertainty with a temperature τ.
Common mistakes / failure modes
Checkpoint exercise
Where this points next
A model can now output a distribution: at the fork, a mixture with a component for each kind of ending, whose probabilities are close to calibrated on the event we checked and whose samples stay on the data nearly all of the time (0.5 % off the data per draw at K = 3). But a sample is only useful when it is used, and using a model of consequences means feeding each prediction back in as the next input, again and again, though the network was only ever trained on real inputs. Even at 0.5 % per draw, a chain of 30 draws leaves the data 13 % of the time if the draws are independent, and the single Gaussian does so 68 % of the time in ten. They are not independent: a draw that leaves the data hands the next step a state the network never saw, where its output is whatever it extrapolates. How fast do errors grow when a model predicts from its own predictions, and how far ahead can it be trusted?
Interview prompts
- Why does training with squared error predict the average of a bimodal outcome? (§1 — E[(y − c)²] = Var(y) + (c − E y)² is minimised at the conditional mean, whatever the shape.)
- The fork is deterministic. Why does its model still need a distribution? (§2 — the aim error is not in the model's input; given one reading it is a Gaussian of 8 cm, as large as the whole band of aims that leave the ending undecided.)
- Define aleatoric and epistemic uncertainty and say which one an ensemble measures. (§2, §4 — aleatoric is what remains given the inputs; epistemic is lack of data; the ensemble's disagreement is the second.)
- Write the loss of a mixture density network and the gradient with respect to a mean. (§3 — NLL = −ln Σ π_k N(y; μ_k, σ_k²); ∂/∂μ_k = r_k(μ_k − y)/σ_k², r_k the responsibility.)
- How would you choose the number of mixture components? (§3 — held-out NLL: it falls while a component captures a real outcome and flattens after; the count matches the kinds of ending.)
- What is calibration, and why is it not enough? (§4 — events given probability p happen a fraction p of the time; a single Gaussian is within a few points on an event and still puts half its mass level with the post, where a fifth of the real endings are.)
- Why does sampling from a model make a small off-data probability matter? (§5 — samples are fed back as inputs; the chance that at least one of k self-fed steps leaves the data is 1 − (1 − q)^k.)
Companion reads: Lesson 24 · Spreads, not averages (the same problem as a training recipe at scale), Robot Model Training · 04 People are not functions: distributions over actions (the policy-side version of the same failure), Generative Models · 11 VQ tokenizers (discrete latent variables) and Reinforcement Learning · 08 Exploration (epistemic uncertainty used to explore).