all_lessons/World Models/04 · Learned filterlesson 4 / 31

Learning the filter and the dynamics

Lesson 2 built the filter by hand, with the model written in; lesson 3 said what a state must keep. Neither says where the machine comes from. It comes from the stream: the next reading is the answer to a question asked one step earlier, so predicting it needs no labels. A linear predictor of the next reading from the last m readings rediscovers the Kalman filter, its error falling toward the filter's innovation variance from above. A recurrent network packs the window into a fixed-size state that is a belief, carries the ball through the curtain and, given the nudges as inputs, answers what-if questions; the miss between prediction and arrival is a surprise signal. Then the Courtyard's fork, where the same recipe, trained by squared error, answers with the average of two futures.

The thesis, here
Predicting the next observation is a complete training signal. Its best answer under squared error is the conditional mean, and in a linear-Gaussian world that is exactly the Kalman filter's prediction (Kalman, 1960), so a learner that predicts well has to rediscover the filter. A recurrent state is the learned, fixed-size version: it becomes a belief and, with the actions as inputs, a dynamics model. The recipe is also its own limit: where the future is not one point, the conditional mean is still what it learns, and that can be a place the world never visits.
Linear position
Forced by: A good state keeps what predicts the future that matters and drops the rest; reconstructing the observation is the wrong way to ask for it, and predicting in a latent space needs protection against collapse. We still lack the machine itself: something that turns a stream of observations and actions into that state and moves it forward. How is the filter learned from data?
New idea: the next reading is the teacher: a network trained to predict it has to become the filter, and a recurrent state is that filter learned from data into a fixed-size vector. Give it the actions as inputs and it is a dynamics model; compare its prediction with what arrives and the miss is a surprise signal.
Forces next: A recurrent state trained to predict the next observation becomes the belief, and with actions as inputs it becomes a dynamics model; in a deterministic world it is almost exact. But at the post the ball goes up or down depending on a difference smaller than the sensor can see, and a model trained to minimize squared error predicts the average, a ball that passes straight through the post. What should a model output when the future is not one point?
The plan
Six moves. (1) Say what the stream teaches and what floor sits under any predictor. (2) Fit a linear predictor on a window of readings and compare it with the Kalman filter. (3) Compress the window into a recurrent state and read what the state holds. (4) Put the actions in: a what-if machine. (5) Compare prediction with arrival: surprise, and the stochastic latent of the RSSM. (6) Run the recipe on the fork and watch what it answers.

1 · The stream is the teacher

Lesson 2's filter is exactly as good as the model written into it: the friction, the sensor noise, the doubt it adds at each step. A learner handed only the stream has neither the model nor labelled states, and it does have a teacher. Every reading is the answer to a question the world asked one step earlier, what will the sensor say next?, so the stream labels itself: the training pair is (everything seen so far, the next reading). The stream here is the readings (x, y), not pictures: a reading carries no nuisance and the target is the observation itself, fixed by the world, so the collapse of lesson 3 cannot happen. The best possible answer is fixed by one identity. For any guess c of a random reading z,

E[(z − c)² | past] = Var(z | past) + (c − E[z | past])²

The first term belongs to the world, the second to the guess, and it vanishes only at the conditional mean. A learner that minimises squared error on the next reading is pushed to E[zt+1 | past] and no further.

In free flight that mean has a closed form: along each axis the Courtyard is linear and Gaussian. This lesson switches on a small random jolt to each velocity at every step (without it the whole path is fixed by its start and the error keeps falling as the window grows, to 1.28 times the sensor's noise variance at m = 14 and 1.05 at 56: no window is long enough).

symbolmeaningvalue
(p, v)position and velocity along one axis; one step is 0.1 s: p′ = p + g v, v′ = d v, then the jolt; F is this step as a matrix
d, gvelocity kept per step, e−γΔt; position gained per unit velocity, (1 − d)/γ0.9656, 0.0983 s
σvjolt added to the velocity at each step (process noise, Q = diag(0, σv²))0.08 m/s
σO, H, Rsensor noise on the position; the sensor reads p only, H = (1 0); R = σO²0.10 m

The filter of lesson 2 carries the covariance of its doubt, which grows before each reading and shrinks after it. Iterate that recursion (the Riccati recursion) and it settles within a few dozen steps: before a reading the filter is uncertain about the position by 6.67 cm, and the next reading is Gaussian around the filter's prediction H x⁻ (x⁻ = F x, the state it expects before the reading) with variance

S∞ = H P⁻ Hᵀ + R = (6.67 cm)² + (10 cm)² = 0.0144 m², an rms miss of 12.0 cm

the innovation variance. Run the filter over a quarter of a million simulated steps and the variance of its innovations (reading minus prediction) agrees with S∞ to within 1 %. No predictor of the next reading, whatever its form or size, goes below it in this world. A predictor that sees only the latest reading misses by 1.77 times that floor (section 2), so what the filter has and this learner lacks is memory. The question is how a learner shown only readings can build it, and how much it needs.

2 · A window of readings learns the filter

The simplest learner is linear in a window of the last m readings, ẑt+1 = Σj=0..m−1 wj zt−j, with the weights fitted by least squares (ridge penalty 10⁻⁶, no intercept). The data are 1500 launches in free flight (a floor wide enough that no wall is in reach, speed 2–3 m/s, any direction), x and y pooled. Every window length is scored on the same windows, those that predict step 14 onward of each of 6000 fresh launches, as a multiple of S∞. With m = 1 the predictor sees one reading, which holds no velocity, and misses by 1.77 times the floor. Then the ratio falls: 1.71 at 2, 1.49 at 4, 1.24 at 6, 1.11 at 8, 1.04 at 10, 1.00 at 14. It falls toward the floor and, to within the noise of the test set (a quarter of a percent), never goes under it.

Why it must. In a linear-Gaussian world the conditional mean given the last m readings is a linear function of them, so a least-squares window searches the right family. Write the steady filter as xt = M xt−1 + K zt, with M = (I − K H) F and K the steady gain. Unrolled, xt = Σj Mj K zt−j, so the filter's prediction of the next reading, H F xt, is itself a linear window, of infinite length, with weights

hj = H F Mj K

A window of m readings is that filter started m readings ago knowing nothing about the ball, and its error Sm comes from the same recursion: 4.02 times the floor at m = 2 (two noisy readings are the least that show a velocity, and extrapolating from just those two misses by about four times the floor), 1.115 at 8, 1.004 at 14. The fitted curve lies under it for short windows, by 27 % at m = 3, because the fit has also learned that launches move at 2–3 m/s, and on it for long ones, by 1 % at 8.

The weights agree with the impulse response. For the Courtyard hj is 0.363, 0.284, 0.213 at lags 0, 1, 2: positive and falling, then slightly negative from lag 8 (−0.009), which is how the filter contrasts old readings with new ones to read a velocity out of positions. The weights fitted at m = 14 match hj to within 0.011 at the first eight lags, and the relative distance between the two vectors falls from 0.91 at m = 6 to 0.36 at 10 and 0.10 at 14. Nobody wrote F, the jolt or the sensor noise into the regression; least squares found the filter because the filter is the predictor that minimises squared error.

What the window costs. It needs all m readings in memory, two numbers per step for the two axes, and the curtain removes 5 in a row (the median over the 400 test launches). A window of 5 has an answer at 66.6 % of the steps of a launch, a window of 14 at 27.5 % (the first m − 1 steps have no full window either). And the floor is set by the noise, not by m: from 10 readings to 14 the error improves by only 3.7 %.

3 · Compress the window into a state

The window's costs say what a better memory needs: a summary of fixed size, updated one reading at a time, that a missing reading does not erase. A recurrent network keeps a fixed number of units and lets each reading update them: ht = tanh(Wx xt + Wh ht−1 + b), ẑt+1 = Wy ht + by. The input xt is the reading (zero behind the curtain), a flag saying whether there was one, and the nudge applied at that step; h has 16 units. The loss is the squared error on the next reading, counted only where that reading exists. Training is backpropagation through time (Adam, one episode per update) on 200 logged episodes of the plain Courtyard: walls, curtain, the jolts, random nudges.

After 60, 120, 180, 240 and 300 epochs its error on free-flight steps (eight readings before, no nudge, 0.6 m from any wall) is 1.86, 1.47, 1.32, 1.27 and 1.22 times the floor; the Kalman filter, which is told the model, sits at 1.01 on the same steps. Sixteen numbers and no model come within 21 % of the model-based filter's error.

Is the state a belief? Fit a linear probe, ridge regression from h to the true position and velocity, on 200 launches the network never trained on, and test it on 400 others. It reads the position to 6.6 cm and the velocity to 0.23 m/s. From the current reading alone the best linear answers are 10.0 cm and 0.35 m/s; the Kalman filter's own posterior is 5.5 cm and 0.17 m/s. The state holds a position better than any single reading and a velocity that no reading contains; it is close to the belief without equalling it. Nobody asked for either quantity; the next reading needed them.

Behind the curtain. With no reading the network keeps predicting from its state: dead reckoning. At the first reading after a clean crossing (195 test launches with a gap of three to five steps, no nudge, the ball between 1 and 4 m high) its error is 2.5 times the floor, against 2.1 for the Kalman filter, whose own covariance had claimed 2.0. Two readings later the network is back to 1.23, three later to 1.15. A window has no answer during the gap, nor for m − 1 steps after it.

4 · Actions as inputs: a dynamics model

The input already carries the nudge. While the data were logged the world applied a random impulse at 8 % of the steps (0.5–2.5 m/s, any direction) and the log recorded it, so the network learned pairs (history and nudge → next reading), and what if I nudge by a? has an answer. Run the network over the real readings up to step 18, put a in at step 18, and let it imagine five steps on its own predictions (each predicted reading is fed back as the next input); do the same with no nudge; the difference is its estimate of the nudge's effect. In free flight the simulator's answer is exact: a nudge adds a to the velocity and nothing else, so after k steps the position differs by g a (1 − dk)/(1 − d) and the jolts cancel. Take a nudge of 1.5 m/s along y (up on some launches, down on others) over the 125 launches that are in free flight at step 18. The true effect at k = 3 is 42.7 cm; the imagined effect is 1.04, 1.00 and 1.01 times the true one at k = 1, 3, 5; the ratio varies from launch to launch with a standard deviation of 0.10 at k = 3, and the error of the effect is 12 % of its size. The network never saw the simulator, only the logs.

Legitimate here, not in general. The logged nudges were random: independent of everything the network could see, so whatever changed after one was caused by it. The model of Ha and Schmidhuber (2018) is trained on 10,000 rollouts of a random policy, whose actions are independent of the state in the same way. Had the nudges been chosen by someone watching the ball, the same pairs would say something else.

5 · Surprise, and the stochastic latent

A prediction is a claim, so its miss is information. Define the surprise at a step as the squared miss in units of the miss the network makes in free flight, s = |zt+1 − ẑt+1|²/(2 ref) over the two axes (ref is the free-flight error per axis), so that its mean in free flight is 1. At the first reading after a clean crossing it is 2.1 on average. After a push of 3 m/s along y that the world feels and the log does not record, it is 3.1, 4.2, 3.2 and 2.5 at the four frames that follow, and still 1.6 at the sixth. It is modest because the push moves the ball 0.29 m in a step against a sensor that already jitters by 0.10 m: a surprise is measured against what the model believes is possible.

A probabilistic model has the same signal built in. Give the state a stochastic part st. Before the reading the model holds a prior p over it, what it expects; after the reading a posterior q, what it now believes; the Kullback–Leibler divergence between them says how far the reading moved the belief. For the Kalman filter everything is Gaussian and exact. With the innovation ν = z − H x⁻, its variance S and κ = H P⁻ Hᵀ/S, the share of the reading's variance that was doubt about the position,

KL(q ‖ p) = ½ [ κ (ν²/S − 1) − ln(1 − κ) ]

(for n states the covariance term of the Gaussian KL is n − κ, which cancels its −n; the mean moves by K ν; the determinants differ by 1 − κ). The first term averages to zero and swings with how surprising the value was; the second is what any reading is worth. In the Courtyard's steady state κ = 0.308 and a reading is worth 0.184 nats per axis; at the first reading after the curtain, when the prior has grown wide, it is worth 0.55, 3.0 times as much.

The stochastic latent (RSSM), in equations
The recurrent state-space model of PlaNet (Hafner et al., 2019) splits the state in two: a deterministic memory ht = f(ht−1, st−1, at−1) and a stochastic part with a prior p(st | ht) (predict, no observation) and a posterior q(st | ht, ot) (correct, with it). Both paths are needed: PlaNet's pure-recurrent and pure-stochastic variants are both worse. DreamerV3 (Hafner et al., 2023) reads the observation back out with a decoder p(ot | ht, st) and trains the whole with a prediction loss, a dynamics term KL[sg(q) ‖ p] that teaches the prior to predict the posterior, and a representation term KL[q ‖ sg(p)] that keeps the posterior close to what the prior could have said (sg = stop-gradient), each clipped below 1 nat. The KL is the surprise of this section, learned. The network of this lesson is the deterministic path alone: it trains no st and no KL term.

The widget

A window, a recurrent state, a nudge, a push and a fork
Top left: one launch (grey line = the ball, cyan dots = readings, amber crosses = the window predictor's guess of the next reading, purple rings = the recurrent net's). Top right: the window's error against its length as a multiple of the floor (amber), the filter started cold (dashed) and the trained net (purple). Bottom left: the weights (bars) against the filter's impulse response (line). Bottom right: the mode's experiment. The second slider is the launch (mode 1), the nudge or push in m/s (modes 2, 3) or the single reading in cm (mode 4).
window error, m readings
—
floor S∞
—
weights vs filter response
—
steps with a full window
—
recurrent net error
—
state read from h
—
error after the curtain
—
nudge effect ÷ true, k = 1, 3, 5
—
surprise, frames 0–3 after the push
—
fork: error at the coming-out frame vs elsewhere
—
net's point inside the diamond
—
real outcomes within 0.3 m of it
—
Show the core JS
DY.window = function (gm, m, lam) {
  ...
  for (i = 0; i < m; i++) { for (j = 0; j < m; j++) A[i * m + j] = gm.G[i * D + j]; A[i * m + i] += lam; b[i] = gm.G[i * D + gm.L]; }
  return CY.la.solve(A, b, m);
};
...
  var M = [(1 - K[0]), (1 - K[0]) * g, -K[1], -K[1] * g + d], h = [], v = [K[0], K[1]], j;
  for (j = 0; j < nh; j++) { h.push(v[0] + g * v[1]); v = [M[0] * v[0] + M[1] * v[1], M[2] * v[0] + M[3] * v[1]]; }
...
RNN.step = function (h, x) {
  ...
    var s = this.bh[j];
    for (i = 0; i < nI; i++) s += x[i] * this.Wx[i * nH + j];
    for (i = 0; i < nH; i++) s += h[i] * this.Wh[i * nH + j];
    hn[j] = Math.tanh(s);
...
  for (k = 0; k < K; k++) { r = net.step(h, x); h = r.h; out.push([r.y[0] * 4 + 4, r.y[1] * 4 + 2.5]); x = [r.y[0], r.y[1], 1, 0, 0]; }

What to try. In the first mode drag the window from 1 to 14: the amber curve gives 1.77 times the floor at m = 1, 1.49 at 4, 1.11 at 8 and 1.00 at 14, the dashed line is the filter started cold, and the bars climb onto the filter's impulse response (distance 0.10 at 14). The amber crosses in the arena vanish behind the curtain and for m steps after it, and the share of steps with a full window falls from 86.5 % to 27.5 %. Press the first button five times: the purple line drops to 1.86, 1.47, 1.32, 1.27 and 1.22, the rings carry on through the curtain, the probe reads 6.6 cm and 0.23 m/s, and the bars show the error after the curtain, 2.5 times the floor at the first reading and 1.15 three readings later. In the nudge mode (second slider at +1.5) the imagined effect is 1.04, 1.00 and 1.01 of the true one at k = 1, 3, 5. In the push mode (+3.0) the surprise is 3.1, 4.2 and 3.2 at the first three frames, against 1 in free flight. In the last mode press the second button five times and set the reading to 0; section 6 reads the result.

6 · The post: the same recipe on a fork

The recipe has one assumption in its loss: that the mean of the possible next readings is a good answer. Put the ball where that fails. The Courtyard's fork is a diamond post with its left vertex at x = 3.8 m, just beyond the usual curtain. The launcher's aim error b is hidden and Gaussian with a standard deviation of 15 cm, and the jolts of sections 1 to 5 are switched off: given b the whole path is fixed, and the only randomness is b and the sensor. A ball aimed more than about 4 cm to one side of the vertex slides up or down the post's face and leaves on that side; a closer one hangs on the vertex, and the closer it is the longer it hangs; within the 3.2 s of this stream the edge of that band is at ±4.4 cm. A difference smaller than the sensor's 10 cm decides the ending. To keep the readings from settling the ending, the curtain here runs from x = 0.6 m to 4.0 m and hides the approach and the first contact. The sensor reads the ball once at the launcher, z′ above the centre line, and next when it comes out, after 21 to 32 steps, on average 0.43 m above or below the centre line; 77 % of the launches come out within 32 steps. One reading leaves b Gaussian with a standard deviation of 8.3 cm around 0.69 z′ (precisions add: 1/0.15² + 1/0.10² = 144 m⁻²).

Train the same network with the same loss for 300 epochs on 240 such launches. At every step that has something to predict except the first reading after the post, its error is 0.013 m² per axis, 1.26 times the sensor's own noise variance (0.010 m²): in a world with no jolts it is almost exact. At the coming-out frame the best any predictor can do, from the exact posterior over b pushed through the real physics, has an expected error of 0.047 m² over the launches that come out; the exact predictor scores 0.046 on the 313 test launches that do. The network's is 0.045 m², the same within sampling noise, and its answers sit 8.4 cm (rms per axis) from the exact conditional mean. That error is 3.6 times the one it makes elsewhere. For the 52 launches whose one reading is within 5 cm of the centre line it is 0.128 against 0.115 attainable on average, 10 times the error elsewhere.

Take the cleanest case, z′ = 0: the ball comes out above or below with probability ½ each. The network answers (4.03, 2.50), inside the diamond, which holds the points with |x − 4.4| + |y − 2.5| below 0.6 m. Measured so, the answer is 23 cm inside its edge (the centre is 0.6 m inside). Of the real coming-out points drawn from the same posterior, 0 % are within 0.3 m of it. Where the reading does decide, at z′ = +10 cm, the network answers up, and 83 % of the real points are within 0.3 m of what it says. This is the identity of section 1 doing what it said: the mean of two branches is a point the ball cannot be at, and nothing in the loss asks for more. A bigger window, a larger network, more data or a longer run would put the network more exactly on the mean and could not move it off. The failure sits in the question the loss asks.

Road not taken · a bigger window
The tempting fix is more readings, less noise. From m = 10 to 14 the error goes from 1.04 to 1.00 times the floor, and no window goes below 1, because the floor is what the sensor and the jolts leave unpredictable. Each reading kept costs two numbers (§2).
Road not taken · tune a filter for each environment
Lesson 2's way: write the model in and let the Riccati recursion do the rest. It is the better tool when the model is known: on free-flight steps the Kalman filter reaches 1.01 times the floor against the network's 1.22. What the learned filter buys is that nothing was written. The same code and the same loss learned the Courtyard and the fork, where no linear-Gaussian filter knows what a post does.
Road not taken · predict the picture
The same recipe on pictures asks for the next frame. It needs the decoder that lesson 3 argued against: a loss in pixel units spends the model on what is loudest in the picture, and a ball is a few pixels of it.
What this lesson did not do
It worked from readings; pictures are lesson 11. It trained only the deterministic path of the RSSM, with no stochastic state and no KL term. Its what-if question is valid because the logged nudges were random (lesson 7). Predicting the next reading is its only signal, and that is not a sufficient statistic in general: in the float and reset example of Littman, Sutton and Singh (2001) a second, two-step prediction is needed, and how far ahead the state can be trusted is lesson 6. And every trained-network figure here is one run from one fixed initialisation. Over four initialisations of the same network on the same data (this one and three more) the free-flight error ends between 1.22 and 1.35 times the floor, the error at the curtain's edge between 1.9 and 5.1 times, the imagined nudge effect at three steps is 1.00 to 1.12 times the true one, and the fork's answer for z′ = 0 lies 16 to 23 cm inside the post with at most 0.5 % of the real outcomes within 0.3 m of it. The conclusions hold in each run; the third digit does not.

Common mistakes / failure modes

"the network learns the world"
It learns the conditional mean of the next reading. At the fork that mean is a point 23 cm inside the post (§6).
"a longer window always helps"
From 10 readings to 14 the error goes from 1.04 to 1.00 times the floor, and nothing goes below 1 (§2).
"a hidden state is an opaque code"
A linear probe reads position to 6.6 cm and velocity to 0.23 m/s, against 10.0 cm and 0.35 m/s from the reading alone (§3).
"a learned dynamics model answers any what-if"
It answers the ones its logs support: the nudges were random, and even then the effect is off by 12 % at three steps (§4).
"a surprise signal catches any unmodelled push at once"
A 3 m/s push moves the ball 0.29 m in a step, three times the sensor's 10 cm, and scores 3.1 times free flight at its first frame, 4.2 at the second, when the net has seen a reading of the new motion (§5).
"the fork failed because the net was too small"
At the coming-out frame its error is 0.045 m² against 0.047 attainable: it is optimal for squared error (§6).

Checkpoint exercise

Try it
A ball only drifts (one state, its position, F = 1) and a steady filter watches it with gain K = 0.3. (a) What weights does the filter put on the readings when it predicts the next one? (b) How long must a window be to hold 95 % of the total weight? Answer: (a) the update is xt = (1 − K) xt−1 + K zt, so hj = K (1 − K)j: 0.30, 0.21, 0.147, … They sum to 1. (b) The weight beyond m readings is (1 − K)m, which is 0.058 at m = 8 and 0.040 at 9, so 9 readings.

Where this points next

A recurrent state trained on the next reading is now the filter, the belief and, with the nudges as inputs, a dynamics model, and it carries the ball through the curtain. With the jolts of sections 1 to 5 it sits at 1.22 times a floor nothing can beat; on the fork, where nothing is random but the hidden aim and the sensor, it is almost exact away from the post (1.26 times the sensor's own noise variance). But the floor and the loss both assume that the future given the past is one blob with one centre. At the post it is two. The same network, trained the same way on the fork, answers 23 cm inside the post at the frame where the ball comes out, when its first reading was on the centre line, and none of the real outcomes is within 0.3 m of that answer, though its error there is the smallest squared error can reach. What should a model output when the future is not one point?

Takeaway
The next reading is the teacher. Squared error on it is minimised by the conditional mean, which in a linear-Gaussian world is the Kalman prediction, with a floor S∞ = 0.0144 m² that no predictor beats. A linear window of the last m readings approaches the floor from above (1.49 at 4, 1.00 at 14) with weights that converge to the filter's impulse response, at the price of m readings that must all exist. A recurrent state compresses the window into 16 numbers (1.22 times the floor after 300 epochs), approaches a belief that a linear probe reads (6.6 cm and 0.23 m/s, against 5.5 cm and 0.17 m/s for the Kalman posterior), and dead-reckons through the curtain. With the actions as inputs it answers what-ifs, legitimately because the logged actions were random. Its miss is a surprise, which the RSSM's KL term learns. And on the fork, a world with no jolts where it is otherwise almost exact, the same recipe outputs the average of two futures: a point 23 cm inside the post, at the best squared error attainable.

Interview prompts

Companion reads: Lesson 21 · From teacher forcing to rollout (the same next-step recipe as a training schedule at scale), Lesson 22 · How the handle gets in (actions as inputs, in practice), Reinforcement Learning · What the agent sees (the POMDP view of the same problem) and Computer Vision · 11 Keypoints, pose, tracking (Kalman filters as trackers).