Part of World Models Handbook
Explains how MuZero plans with learned latent dynamics and no reconstruction loss, plus a tutorial probing value equivalence and search targets.
Choose your expertise level to adjust how many terms are explained. Beginners see more tooltips, experts see fewer to maintain reading flow. Hover over underlined terms for instant definitions.
Article links
Make inline references clickable
MuZero and Value-Equivalent Models
Picture a video game. The screen is a torrent of pixels: explosions, particle effects, a scrolling scoreboard, a blinking cursor in the corner. Now picture the decision the game asks of you: move left, jump, or wait. Much of that pixel torrent may not matter for the decision. What matters is a task-relevant summary: where am I, where is the enemy, how much health do I have, is the platform solid? A game-playing agent must learn which details to retain for its decisions, and its training objective influences what its internal representation keeps.
A prominent visual-model lineage learns to reconstruct observations alongside its decision model. PlaNet and the Dreamer family, discussed in Part VIII: Decision-Centric Research Lineages, use a reconstruction objective to help train their latent states. That objective can spend capacity on observation details that do not change an action. This is one approach to model-based reinforcement learning, not its historical definition; classical control models often operate directly on compact physical states.
MuZero, introduced by Schrittwieser and colleagues in 2019 and refined in 2020, asked a sharper question. What if we skip the reconstruction entirely? What if the only requirement on the learned model is that it supports good decisions?
This decision-centric answer became highly influential. MuZero learns three connected functions:
- a representation function that turns an observation into a latent state,
- a dynamics function that advances a latent state under an action and predicts the reward,
- a prediction function that reads a policy and a value out of a latent state.
Here, a "latent state" means an internal representation the network learns, with no requirement that it reconstruct the raw observation. The reward is scalar feedback from the environment, and the policy is a distribution over actions. These three functions form MuZero's learned model; self-play data-collection workers, a replay buffer, and tree search are also essential parts of the agent.
There is no observation decoder or reconstruction loss in the original MuZero objective. Planning uses Monte Carlo tree search over the learned latent dynamics. Schrittwieser et al. reported strong results in Go, chess, shogi, and Atari without giving the model the games' transition rules. The paper's large-data Atari comparison used 20 billion frames for MuZero and 37.5 billion for R2D2; its 200-million-frame MuZero Reanalyze setting is a separate experiment. Neither result supports a blanket claim of a hundredfold or thousandfold sample-efficiency improvement.
This chapter builds a MuZero-style tutorial model on a small gridworld and tests latent prediction, tree search, and search-generated targets. The measurements also show limits to what these experiments establish: more search can worsen the planner's Q-greedy action agreement, while better value targets alone do not establish better learned-policy decisions. We will connect this to value equivalence, a formal criterion based on selected Bellman updates rather than pixel-level prediction.
Representation, Dynamics, and Prediction Networks
MuZero uses a learned latent model for planning instead of querying the environment during search. It does not plan in observation space. It plans in a representation learned from reward, value, and policy targets. This is a substantial shift in perspective. In a classical state-space model, the state variables may be supplied by the problem: the position of a chess piece, the reading of a sensor. In MuZero, the latent state is learned from data, and its usefulness must be evaluated through predictions and decisions. If planning fails, training and evaluation must diagnose and improve the model; MuZero does not automatically discard and relearn the entire state space.
Formally, in a Markov decision process the agent interacts with an environment over discrete time steps . At each step the environment is in some state, the agent selects an action, and the environment returns a reward and a new state. The learned model is called latent because its state is not the environment's state; it is a vector of the network's own invention. Throughout this chapter we distinguish:
- the observation (what the agent sees at wall-clock time );
- the action (what the agent chooses at wall-clock time );
- the reward (what the environment returns at wall-clock time );
- the policy (a distribution over actions);
- the transition (how the environment's state evolves);
- the observation model (how observations are generated from states);
- the latent state and unrolled horizon used inside the learned model.
Keeping these separate matters, because the central confusion you may run into is mixing up the unrolled time inside the learned model with the wall-clock time of the real environment. The superscript counts steps taken inside the imagination of the network. The subscript counts steps taken in the world. In our tutorial, a single observation is encoded into the initial latent state . Published MuZero's representation instead takes an observation-and-action history. After the root is encoded, search rolls forward through learned transitions rather than querying new environment transitions, though it may query the environment for legal actions at the root.
The three functions
Write the environment's observation at time as (what the agent sees), the action as (what the agent chooses), the reward as (what the environment returns), and the latent state as (the compressed summary the network invents at model step ). The superscript indexes steps after the encoded root, not wall-clock time. In the training equations below, is the replay-supervised unroll length; search depth is a separate choice. MuZero parameterizes three functions with a shared network (layers within each function have their own parameters, conventionally grouped as ). The representation function encodes available history into an initial latent state; our compact tutorial uses only the current observation. The dynamics function advances that latent state under an action while predicting the reward, and the prediction function reads a policy and a value out of a latent state:
where:
- : the observation the agent sees at time
- : the representation function, which encodes available history into an initial latent state ; the displayed is our single-observation tutorial specialization
- : the action taken at unrolled step (encoded as a one-hot or embedding)
- : the dynamics function, which maps a latent state and an action to the next latent state and the reward on the transition into it
- : the latent state at unrolled step ; the tutorial implementation below bounds it with a tanh activation
- : the prediction function, which reads a latent state and emits a policy (a distribution over actions, used as a prior during search) and a value (the expected discounted return from that latent state)
- : the model-step index within the training unroll, not wall-clock time
- : the number of model steps used for a replay-supervised training unroll; it is not the tree-search depth
Read the three equations top to bottom and the whole algorithm is visible. The first line says "look at the world and compress it." The second says "given a situation and an action, imagine the resulting situation and the reward it brings." The third says "given a situation, what would you do and how good is it?" Everything else in the chapter is a consequence of these three lines plus a search procedure.
Notice what is missing: there is no inverse dynamics function, no decoder, and no reconstruction head. The latent state is not constrained to be interpretable and is not compared to a real observation. The states produced by the dynamics function are trained to support reward, value, and policy predictions useful for planning. The model need not reproduce every observation detail, although its latent state can still retain irrelevant information.
The original MuZero dynamics function is deterministic: given a latent state and an action, it returns one next latent state rather than a sampled distribution. Our gridworld's position dynamics are deterministic, although its observation includes independent noise. A deterministic model can still predict useful expected values in a stochastic environment, but a single imagined branch does not represent the full outcome distribution. Search quality must be checked against real-environment outcomes.
In this chapter a latent state is a vector produced by the representation or dynamics network. It is not a belief state: as discussed in Part II: Inference, Decisions, and Control, a belief state is a distribution over environment states. Our tutorial bounds its latent coordinates with ; this is an implementation choice, not a property of all MuZero models or a guarantee of useful planning.
The training objective
Suppose we sample a trajectory from the replay buffer and pick a starting index . We run the network forward from , unrolling steps with the actual actions the agent took:
At the root and each unrolled state we predict value and policy ; each modeled transition predicts reward . The reward target is the observed reward on that transition. The policy target comes from search at the corresponding real state. In the original MuZero training objective, the value target is the final outcome for board games or an observed -step return bootstrapped from a later search value for Atari. With our reward indexing, a schematic loss is:
where:
- : a trajectory sampled from the replay buffer , from which the starting index is drawn
- : the total training loss, minimized over the shared parameters
- : the number of model steps in this training unroll
- : the observed reward on the th transition after the starting state, for
- : the reward predicted by dynamics on that transition; no root reward is predicted
- : the outcome or -step bootstrapped value target appropriate to the domain
- : the value predicted by the prediction function at unrolled step
- : the policy target produced by the search at time
- : the policy prior predicted by the prediction function at unrolled step
- : the reward loss, whose form depends on the domain and target representation
- : the value loss, likewise chosen for the domain and target representation
- : the policy loss, a cross-entropy between the search visit-count distribution and the network prior
- : the coefficient of the weight-regularization term
- : the squared norm of all parameters in the shared network , which penalizes large weights
In published MuZero, board games use a squared-error value loss and, because those experiments have no intermediate rewards, omit the reward-prediction loss. Atari uses categorical-support cross-entropy for transformed reward and value targets. The policy loss compares search visit counts with the network prior. The tutorial below uses simple scalar mean-squared reward and value losses for ease of inspection, not the complete published objective.
The expectation is over trajectories in replay. In published MuZero, actors repeatedly generate new trajectories using MCTS; replay is refreshed, not a permanently fixed dataset. Our tutorial instead uses a fixed exploratory dataset so that the small experiment is reproducible.
Search provides policy targets at replayed states; the value target also incorporates actual outcomes or later bootstrapped search estimates. This can improve decision targets, but it does not guarantee that a model trained on them will improve, especially outside the replay distribution. Our fixed-data experiment should be read as a controlled demonstration, not as a reproduction of MuZero's self-play loop.
This objective is unusual and worth examining closely. It does not ask the latent to reconstruct the next observation or resemble the environment state. Reward, value, and policy losses instead constrain the quantities used for decisions. Many different latent encodings can satisfy those targets, and stable learning depends on data, architecture, optimization, and target quality; the losses alone do not guarantee it.
The latent consistency term, and why it matters here
The original MuZero objective supervises unrolled reward, value, and policy predictions, but does not include an explicit loss matching a predicted latent to the representation of the next observation. Correct predictions constrain the dynamics indirectly, while multiple latent representations may remain equivalent for the supervised targets.
An auxiliary consistency objective can encourage the predicted state and the encoding of the next observation to agree. It is an additional inductive bias, not a required component of MuZero or a proof that search will be reliable.
For example, EfficientZero uses a self-supervised temporal consistency component. Our toy model uses a simpler squared-distance variant. Write for the state component of ; then:
where:
- : the latent consistency loss at unrolled step
- : the state component of the dynamics prediction after applying at
- : the stop-gradient operator, which treats its argument as a constant during backpropagation
- : the latent state the representation function would have produced had it seen the observation steps later
- : the unrolled step at which the consistency loss is computed, ranging over
This term compares the predicted state with a stopped-gradient next-observation encoding. It is self-predictive rather than reconstructive: nothing decodes pixels. Stop-gradient blocks the direct gradient through the target occurrence of the encoder. The encoder can still change through the online branch and the reward, value, and policy losses. The squared distance is a tutorial choice, not a latent probability model.
Without a stopped target, both branches can move toward one another. Stopping one branch changes the gradient direction but does not by itself prevent representational collapse: a constant zero latent satisfies the consistency term and survives . Decision losses and training design must preserve useful information.
The toy experiment asks whether this combination of losses reduces the representation of four unpredictable, decision-irrelevant observation bits. We will measure the result rather than assume that the bits disappear.
Why the latent state is allowed to forget
Consider the kinds of information a world model might be asked to represent:
- Decision-relevant, predictable. Where the agent is, where the goal is, how much health remains. When these factors affect reward or action quality, the learned model needs enough information to predict relevant rewards, values, and policy targets. It need not store each variable literally in the latent; a fixed goal, for example, can be implicit in the network parameters.
- Decision-relevant, unpredictable. An opponent's future move in a stochastic game may affect the return, but its realized outcome is not known at the current state. The representation may need information that predicts the distribution of possible moves or their expected effect on value; a scalar value target alone does not guarantee that it identifies the full distribution.
- Decision-irrelevant, predictable. The scrolling scoreboard animation. The dynamics can predict it, and nothing needs it. It may survive or not; nothing punishes its presence, and nothing rewards it.
- Decision-irrelevant, unpredictable. A random distractor bit. Consistency may discourage encoding it, but finite-capacity training can still leave measurable information in the latent.
These two axes distinguish decision relevance from predictability. A decision-relevant but unpredictable variable may call for uncertainty-aware values; a predictable but irrelevant variable need not be retained. An unpredictable irrelevant variable is a plausible candidate to discard, not something the objective is guaranteed to remove.
A reconstruction loss can reward encoding distractor detail, depending on its weighting and the model's capacity. MuZero's original decision-target objective lacks that direct incentive, but it can still retain distractors incidentally. The later probe measures how much remains in this particular learned latent.
Let us now build the smallest possible environment in which this effect is visible.
import copy
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
SEED = 0
np.random.seed(SEED)
torch.manual_seed(SEED)We use a gridworld. The agent starts somewhere other than the goal, moves in one of four directions with wall clamping, and the goal cell is absorbing. Reaching the goal yields a reward of 1; every other transition yields 0. Episodes run for a fixed horizon of 20 steps, so there is no termination signal to model. The choice of a tiny gridworld is deliberate: it is small enough that we can compute the exact optimal value function and compare against it, and large enough that a search of depth four cannot trivially exhaust it. You should think of it as a magnifying glass under which the mechanism of MuZero becomes visible in a way it never is on Atari.
The observation is deliberately redundant: a 16-dimensional one-hot of the agent's position, concatenated with 4 bits of pure uniform noise resampled independently at every step. The noise is unpredictable and decision-irrelevant. A model that plans well has no reason to represent it. The four noise channels play the role of the blinking cursor from the opening paragraph. They are always present, always random, and never useful. If a model wastes latent capacity on them, you will be able to see it directly.
N = 4 # grid side length
N_STATES = N * N # 16 cells
N_ACTIONS = 4 # up, right, down, left
GOAL = (0, 0)
HORIZON = 20
GAMMA = 0.95
N_DISTRACT = 4 # pure-noise observation channels
OBS_DIM = N_STATES + N_DISTRACT
LATENT_DIM = 8
def one_hot_position(pos):
v = np.zeros(N_STATES, dtype=np.float32)
v[pos[0] * N + pos[1]] = 1.0
return v
def observe(pos, rng):
"""Position one-hot plus independent uniform noise bits."""
noise = rng.integers(0, 2, size=N_DISTRACT).astype(np.float32)
return np.concatenate([one_hot_position(pos), noise])
def transition(pos, action):
"""Advance one grid step and reward first entry into the absorbing goal."""
if pos == GOAL:
return GOAL, 0.0
row, col = pos
if action == 0:
row = max(row - 1, 0)
elif action == 1:
col = min(col + 1, N - 1)
elif action == 2:
row = min(row + 1, N - 1)
else:
col = max(col - 1, 0)
next_pos = (row, col)
return next_pos, (1.0 if next_pos == GOAL else 0.0)The transition function used everywhere below is the deterministic gridworld step: it moves the agent in the requested direction if the destination is inside the grid, otherwise it stays put, and it returns a reward of 1 exactly when the agent reaches the goal. Because this function is a plain Python routine, we can use it both to generate the dataset and to compute the exact optimum by value iteration, which gives us a ground truth against which every learned quantity can be measured.
We collect data with a behaviour policy that, at non-goal cells, with probability 0.6 takes a step reducing Manhattan distance to the goal (breaking ties randomly) and otherwise acts uniformly at random. At the absorbing goal it acts uniformly. Its empirical returns leave a gap to the exact optimal values. That gap is an opportunity for search, not evidence that this particular learned-model search will close it.
def behaviour_action(pos, rng, bias=0.6):
if pos != GOAL and rng.random() < bias:
candidates = []
for a in range(N_ACTIONS):
nxt, _ = transition(pos, a)
d = abs(nxt[0] - GOAL[0]) + abs(nxt[1] - GOAL[1])
candidates.append((d, a))
best_d = min(d for d, _ in candidates)
return int(rng.choice([a for d, a in candidates if d == best_d]))
return int(rng.integers(N_ACTIONS))
def collect_dataset(n_episodes, rng, bias=0.6):
obs, nxt, act, rew, pos, lengths = [], [], [], [], [], []
for _ in range(n_episodes):
p = (int(rng.integers(N)), int(rng.integers(N)))
while p == GOAL:
p = (int(rng.integers(N)), int(rng.integers(N)))
start = len(obs)
current_obs = observe(p, rng)
for _ in range(HORIZON):
a = behaviour_action(p, rng, bias)
p2, r = transition(p, a)
next_obs = observe(p2, rng)
obs.append(current_obs)
act.append(a)
rew.append(r)
nxt.append(next_obs)
pos.append(p)
p = p2
current_obs = next_obs
lengths.append(len(obs) - start)
return (
np.asarray(obs, np.float32),
np.asarray(act, np.int64),
np.asarray(rew, np.float32),
np.asarray(nxt, np.float32),
np.asarray(pos, np.int64),
lengths,
)
rng = np.random.default_rng(SEED)
OBS, ACT, REW, NXT, POS, LENGTHS = collect_dataset(150, rng)
offset = 0
for episode_length in LENGTHS:
assert np.array_equal(
NXT[offset : offset + episode_length - 1],
OBS[offset + 1 : offset + episode_length],
)
offset += episode_length
assert offset == len(OBS)Because the grid is small and the dynamics are deterministic, we can compute the exact optimal value function by value iteration. That gives us a reference for measuring the learned model and planner. This reference is not the definition of value equivalence: Grimm et al. compare models by selected Bellman updates, which need not guarantee agreement on every optimal value. Value iteration repeatedly applies the Bellman optimality operator. The code below runs 300 sweeps, but from a zero initialization the reward signal reaches the far corner by sweep six in this deterministic grid. The result is the best possible discounted return from every cell.
def monte_carlo_returns(rewards, lengths, gamma):
out = np.zeros_like(rewards)
i = 0
for L in lengths:
g = 0.0
for t in range(L - 1, -1, -1):
g = rewards[i + t] + gamma * g
out[i + t] = g
i += L
return out
MC_RETURNS = monte_carlo_returns(REW, LENGTHS, GAMMA)
def value_iteration(n_sweeps=300):
V = np.zeros(N_STATES)
for _ in range(n_sweeps):
V_new = V.copy()
for s in range(N_STATES):
p = (s // N, s % N)
if p == GOAL:
V_new[s] = 0.0
continue
qs = []
for a in range(N_ACTIONS):
p2, r = transition(p, a)
qs.append(r + GAMMA * V[p2[0] * N + p2[1]])
V_new[s] = max(qs)
V = V_new
return V
def optimal_q(V):
Q = np.zeros((N_STATES, N_ACTIONS))
for s in range(N_STATES):
p = (s // N, s % N)
for a in range(N_ACTIONS):
p2, r = transition(p, a)
Q[s, a] = r + GAMMA * V[p2[0] * N + p2[1]]
return Q
V_STAR = value_iteration()
Q_STAR = optimal_q(V_STAR)
STATE_IDS = POS[:, 0] * N + POS[:, 1]
BEHAVIOUR_VALUES = np.array(
[
MC_RETURNS[STATE_IDS == s].mean() if np.any(STATE_IDS == s) else 0.0
for s in range(N_STATES)
]
)Transitions collected: 3000 Observation dimension: 20 (16 position + 4 noise) Mean logged-transition return-to-go under behaviour: 0.224 Optimal value at the furthest cell from the goal: 0.774 Behaviour value at that same cell: 0.700
The empirical behaviour value at the far corner is below the exact optimum. The undiscounted return-to-go is the sum of subsequent rewards; discounting with weights earlier rewards more heavily. Let denote an environment state, the reward received after moving from , and the behaviour policy. The expected discounted return is
where:
- : the expected discounted return when starting in state and following thereafter
- : the expectation over actions drawn from and over the environment's transitions
- : the reward received at step
- : the discount factor, weighting rewards steps away by
The Monte Carlo returns estimate this quantity from sampled 20-step episodes, with any reward beyond the cutoff set to zero. They are therefore truncated targets, not unbiased infinite-horizon targets for late-episode transitions that have not reached the goal. The reference is infinite-horizon, so RMSE gaps against these empirical returns combine policy quality with truncation and sampling error. Detours delay the goal reward and lower its discounted contribution. We will compare later searches and target changes against this baseline, without assuming they remove the gap.
Now the network. It is deliberately tiny: hidden layers of width 64, an 8-dimensional latent state, and squashing on the encoder and on the state component of dynamics. The squash bounds magnitudes; it does not prevent collapse, because the zero vector is still in its range. Reward, value, and policy supervision are needed to keep the representation useful.
class MuZeroNet(nn.Module):
def __init__(
self,
obs_dim=OBS_DIM,
n_actions=N_ACTIONS,
latent_dim=LATENT_DIM,
hidden=64,
):
super().__init__()
self.latent_dim = latent_dim
self.n_actions = n_actions
self.encoder = nn.Sequential(
nn.Linear(obs_dim, hidden),
nn.ReLU(),
nn.Linear(hidden, latent_dim),
)
self.dynamics = nn.Sequential(
nn.Linear(latent_dim + n_actions, hidden),
nn.ReLU(),
nn.Linear(hidden, latent_dim + 1),
)
self.prediction = nn.Sequential(
nn.Linear(latent_dim, hidden),
nn.ReLU(),
nn.Linear(hidden, n_actions + 1),
)
def encode(self, obs):
return torch.tanh(self.encoder(obs))
def next_state(self, z, action_onehot):
out = self.dynamics(torch.cat([z, action_onehot], dim=-1))
z_next = torch.tanh(out[..., : self.latent_dim])
reward = out[..., self.latent_dim]
return z_next, reward
def predict(self, z):
out = self.prediction(z)
return out[..., : self.n_actions], out[..., self.n_actions]Notice how the three heads map onto the three lines of the MuZero equations. The encode method implements , the next_state method implements and splits its output into a latent half and a scalar reward half, and the predict method implements and splits its output into action logits and a single value. Keeping this correspondence explicit makes the code much easier to read than if the heads were scattered across separate modules.
The training loop uses torch.no_grad() for the target branch of the auxiliary consistency loss. This blocks direct gradients through that occurrence of the encoder, but the encoder remains trainable through the online state and other losses. The loop is a one-step decision-target surrogate: it supervises the root value with behavior-policy Monte Carlo returns, the root policy with sampled behavior actions, and one predicted reward. It does not implement MuZero's multi-step unroll or search-generated policy targets.
def train_model(
model,
obs,
act,
rew,
nxt,
value_target,
policy_target,
epochs=200,
batch_size=256,
lr=1e-3,
seed=0,
use_consistency=True,
use_reconstruction=False,
):
torch.manual_seed(seed)
opt = torch.optim.Adam(model.parameters(), lr=lr)
obs_t = torch.tensor(obs)
nxt_t = torch.tensor(nxt)
act_t = torch.tensor(act)
rew_t = torch.tensor(rew)
val_t = torch.tensor(np.asarray(value_target, dtype=np.float32))
pol_t = torch.tensor(np.asarray(policy_target, dtype=np.float32))
n = obs_t.shape[0]
history = {
k: []
for k in (
"total",
"value",
"reward",
"consistency",
"policy",
"reconstruction",
)
}
for _ in range(epochs):
perm = torch.randperm(n)
sums = np.zeros(6)
n_batches = 0
for i in range(0, n, batch_size):
idx = perm[i : i + batch_size]
z = model.encode(obs_t[idx])
logits, v = model.predict(z)
z_next, r_hat = model.next_state(
z, F.one_hot(act_t[idx], N_ACTIONS).float()
)
with torch.no_grad():
z_target = model.encode(nxt_t[idx])
value_loss = F.mse_loss(v, val_t[idx])
reward_loss = F.mse_loss(r_hat, rew_t[idx])
consistency_loss = F.mse_loss(z_next, z_target)
policy_loss = (
-(pol_t[idx] * F.log_softmax(logits, dim=-1)).sum(-1).mean()
)
recon_loss = torch.tensor(0.0)
if use_reconstruction:
recon_loss = F.mse_loss(model.decode(z), obs_t[idx])
loss = value_loss + reward_loss + policy_loss
if use_consistency:
loss = loss + consistency_loss
if use_reconstruction:
loss = loss + recon_loss
opt.zero_grad()
loss.backward()
opt.step()
sums += np.array(
[
loss.item(),
value_loss.item(),
reward_loss.item(),
consistency_loss.item(),
policy_loss.item(),
recon_loss.item(),
]
)
n_batches += 1
sums /= n_batches
for j, key in enumerate(
(
"total",
"value",
"reward",
"consistency",
"policy",
"reconstruction",
)
):
history[key].append(sums[j])
return history
VE_MODEL = MuZeroNet()
BEHAVIOUR_POLICY = np.eye(N_ACTIONS, dtype=np.float32)[ACT]
hist_ve = train_model(
VE_MODEL, OBS, ACT, REW, NXT, MC_RETURNS, BEHAVIOUR_POLICY, epochs=200
)
The reward and value losses decrease, while the policy loss plateaus and the consistency loss varies and rises slightly late in training. A one-hot sampled action is a noisy target for the policy head; its plateau does not establish the entropy of the behavior policy. This plot alone cannot show whether the latent has discarded distractor bits; the probe below tests that separately.
The consistency loss can be inspected as a distribution rather than as a single scalar. If the dynamics has learned to match the encoder, the per-transition squared error should concentrate at a small value across the dataset.
with torch.no_grad():
z_online = VE_MODEL.encode(torch.tensor(OBS))
z_next_pred, _ = VE_MODEL.next_state(
z_online, F.one_hot(torch.tensor(ACT), N_ACTIONS).float()
)
z_next_target = VE_MODEL.encode(torch.tensor(NXT))
consistency_errors = (
((z_next_pred - z_next_target) ** 2).sum(dim=-1).numpy()
)
UNTRAINED_MODEL = MuZeroNet()
with torch.no_grad():
z_online_untrained = UNTRAINED_MODEL.encode(torch.tensor(OBS))
z_next_untrained, _ = UNTRAINED_MODEL.next_state(
z_online_untrained, F.one_hot(torch.tensor(ACT), N_ACTIONS).float()
)
z_target_untrained = UNTRAINED_MODEL.encode(torch.tensor(NXT))
consistency_errors_untrained = (
((z_next_untrained - z_target_untrained) ** 2).sum(dim=-1).numpy()
)
We evaluate the value head by encoding several noisy observations of each cell and averaging the predictions. Small within-cell spread would indicate that this head is relatively insensitive to the four distractor bits. That is a useful diagnostic, but it does not establish formal value equivalence or prove that the full latent is noise-invariant.
@torch.no_grad()
def predict_values_per_state(model, n_noise=8, seed=0):
rng_local = np.random.default_rng(seed)
obs_list, state_ids = [], []
for s in range(N_STATES):
p = (s // N, s % N)
for _ in range(n_noise):
obs_list.append(observe(p, rng_local))
state_ids.append(s)
stacked = torch.tensor(np.stack(obs_list))
z = model.encode(stacked)
_, v = model.predict(z)
v = v.numpy()
state_ids = np.array(state_ids)
means = np.array([v[state_ids == s].mean() for s in range(N_STATES)])
stds = np.array([v[state_ids == s].std() for s in range(N_STATES)])
return means, stds
def value_rmse(model, n_noise=8, seed=0):
means, _ = predict_values_per_state(model, n_noise=n_noise, seed=seed)
return float(np.sqrt(np.mean((means - V_STAR) ** 2)))
MODEL_VALUES, MODEL_VALUE_STD = predict_values_per_state(VE_MODEL)
ERR_VE = value_rmse(VE_MODEL)
ERR_BEHAVIOUR = float(np.sqrt(np.mean((BEHAVIOUR_VALUES - V_STAR) ** 2)))Decision-target model value RMSE: 0.104 Empirical behaviour-value RMSE: 0.105

The learned value head is trained against behavior-policy Monte Carlo returns. Its RMSE of 0.104 is only marginally below the empirical behavior baseline's 0.105 in this run; the difference does not establish a robust improvement. There is no multi-step unrolled loss in the tutorial training loop, so attributing value behavior to such a loss would be unsupported. Search may use the learned reward and dynamics predictions to change action rankings, which we test next.
The gap to the optimal curve gives room for search or revised targets to help, but their benefit must be measured; neither is guaranteed by the architecture.
Planning Without Reconstructing Observations
Training gives us a latent model. Now we need to use it. MuZero plans by running Monte Carlo tree search directly over latent states, with standing in for the environment's transition function.
This is the core of the method, so it is worth being explicit about what has changed relative to classical planning. In the search methods discussed in Part VII: Planning and Agency, search operates over states that mean something: board positions, robot configurations, belief distributions. Here, search operates over vectors in that have no agreed-upon meaning at all. The tree is built out of network activations. The only thing that makes it a planning tree rather than a random graph is that the reward and value heads read useful quantities off its nodes.
This may feel unusual at first. In a classical planner, if you inspect a node of the search tree, you can describe what it represents: this is the position after moving the knight to f3, this is the belief state after observing the red light. In MuZero, the analogous inspection yields eight floating-point numbers with no built-in board description. We can inspect predicted rewards, values, policies, or a separately trained probe, but any human interpretation of a latent node needs its own validation. The tree can support decisions through its action values and visit counts without each node having a readable label.
The search tree
A node in the tree stores a latent state , the reward that the dynamics function predicts on the transition into that node, a prior probability supplied by the prediction network, and running statistics. We write for the number of times action has been tried at a node holding , for the sum of values backed up through those trials, and for the resulting mean.
When the notation is clear from context we drop the argument and write , , and for a single node. The full form is used in the PUCT rule below, where identifies the node and selects an edge.
Each simulation walks down the tree from the root, choosing at each node the action that maximizes the PUCT score. PUCT (Predictor + Upper Confidence bounds applied to Trees), a rule from the AlphaZero lineage, balances exploiting actions with high estimated value against exploring actions with high prior probability:
where:
- : the mean value of taking action at latent state , equal to
- : the visit count for action at
- : the prior probability the prediction network assigns to action
- : the total visits to across all actions
- : a constant controlling exploration strength
The displayed rule is schematic. Its edge mean is defined after an edge has been visited. The tutorial stores a child-state estimate as child.Q, seeds it from the value head before that child is visited, and scores the parent edge with child.reward + GAMMA * child.Q. It also uses sqrt(node.N + 1) instead of the displayed square root of total action visits, so the exploration bonus is nonzero when the parent has no visits. These choices retain the prior/visit tradeoff without making the equation a line-by-line transcription of the code.
The first term favors actions with high estimated return. The second encourages exploration according to the network prior . At fixed total node visits, its denominator makes the bonus for a particular action fall as that action is tried; a larger prior raises that action's bonus. As total node visits grow, the square-root numerator increases, so the bonus is not constant. Its net behavior depends on how visits are distributed across actions and on the scale of .
The exploration bonus is a concrete, plottable quantity. Holding the total visit count fixed, the bonus for an action decreases as that action is tried and increases with the network's prior for it.
puct_c = 1.5
puct_total_visits = 20.0
puct_visits = np.arange(1, 21)
puct_priors = [0.25, 0.5, 0.75]
puct_bonus_curves = [
puct_c * prior * np.sqrt(puct_total_visits) / (1.0 + puct_visits)
for prior in puct_priors
]
The prior directs finite simulations toward actions favored by the network. A good prior can save search effort; a mistaken one can delay exploration of a strong action. Thus the learned policy shapes search, but mutual improvement must be measured rather than presumed.
When a simulation reaches a node that has never been expanded, we call the prediction network to get a value estimate, and we call the dynamics network once per action to materialize the children. Then the value is backed up along the path, with the immediate reward predicted by the dynamics on each transition folded in. If follows action from , the state-return recursion is:
where:
- : the backed-up discounted return from the node holding latent state
- : the reward predicted by dynamics on the edge from to
- : the discount factor, which weights nearer rewards more heavily
- : the number of simulated transitions before the value bootstrap
- : the prediction network's value estimate at the search horizon
The code stores each child's incoming-edge reward and its estimate of the return from the child state. A parent adds that reward to the discounted child-state return, matching the recursion above. This has the same discounted, additive form as a Monte Carlo return, but its rewards are predicted rather than observed, so reliability depends on those predictions and on the horizon value estimate.
After a fixed number of simulations, the tutorial search returns a root visit-count distribution and the largest estimated root action value. These become direct targets only in the later toy comparison. Original MuZero trains its policy head on search visit counts, while its value target incorporates outcomes or -step returns.
Exact policy improvement in a known tabular MDP can produce a non-worse policy under its assumptions. This finite-budget learned-model search is only an approximation: biased rewards, values, or latent transitions can make its chosen action worse. We therefore compare search decisions with the exact gridworld solution rather than presuming improvement.
Our fixed-horizon gridworld has an absorbing goal rather than a terminating episode, so the tutorial does not need a termination head. Original MuZero likewise does not use an explicit continuation head. Its search can proceed past a terminal state and relies on absorbing-state training and value predictions rather than learned termination. A model for environments with terminal states can predict continuation or discount explicitly, but neither this tutorial nor the cited MuZero results establish that such a head generally improves learning.
ACTION_ONEHOT = torch.eye(N_ACTIONS)
class MCTSNode:
__slots__ = ("z", "reward", "prior", "N", "W", "Q", "children", "expanded")
def __init__(self, z=None, reward=0.0, prior=1.0):
self.z = z
self.reward = reward
self.prior = prior
self.N = 0
self.W = 0.0
self.Q = 0.0
self.children = {}
self.expanded = False
@torch.no_grad()
def expand_node(node, model):
"""Materialize all children of a leaf using the learned dynamics."""
logits, value = model.predict(node.z)
priors = torch.softmax(logits, dim=-1).squeeze(0).numpy()
z_rep = node.z.expand(N_ACTIONS, -1)
z_next, rewards = model.next_state(z_rep, ACTION_ONEHOT)
_, child_values = model.predict(z_next)
for a in range(N_ACTIONS):
child = MCTSNode(
z=z_next[a : a + 1],
reward=float(rewards[a].item()),
prior=float(priors[a]),
)
child.Q = float(child_values[a].item())
node.children[a] = child
node.expanded = True
return float(value.item())
def select_action(node, c_puct):
total = np.sqrt(node.N + 1)
best_a, best_score = 0, -np.inf
for a, child in node.children.items():
q_edge = child.reward + GAMMA * child.Q
score = q_edge + c_puct * child.prior * total / (1 + child.N)
if score > best_score:
best_score, best_a = score, a
return best_a
_selection_probe = MCTSNode()
_selection_probe.N = 1
_selection_probe.children = {
0: MCTSNode(reward=1.0, prior=0.5),
1: MCTSNode(reward=0.0, prior=0.5),
}
_selection_probe.children[0].Q = 0.0
_selection_probe.children[1].Q = 0.5
assert select_action(_selection_probe, c_puct=0.1) == 0
def simulate(node, model, c_puct, depth, max_depth):
if not node.expanded:
value = expand_node(node, model)
node.N += 1
node.W += value
node.Q = node.W / node.N
return value
if depth >= max_depth:
value = node.Q
node.N += 1
node.W += value
node.Q = node.W / node.N
return value
a = select_action(node, c_puct)
child = node.children[a]
value = child.reward + GAMMA * simulate(
child, model, c_puct, depth + 1, max_depth
)
node.N += 1
node.W += value
node.Q = node.W / node.N
return value
_depth_cap_probe = MCTSNode()
_depth_cap_probe.expanded = True
_depth_cap_probe.N = 1
_depth_cap_probe.W = 0.5
_depth_cap_probe.Q = 0.5
assert np.isclose(simulate(_depth_cap_probe, None, 1.5, 4, 4), 0.5)
assert _depth_cap_probe.N == 2
assert np.isclose(_depth_cap_probe.W, 1.0)
@torch.no_grad()
def run_mcts(model, obs, n_sims=32, c_puct=1.5, max_depth=4):
"""Plan from a single observation and return (value, action values, policy)."""
z0 = model.encode(torch.tensor(obs, dtype=torch.float32).unsqueeze(0))
root = MCTSNode(z=z0)
expand_node(root, model)
for _ in range(n_sims):
simulate(root, model, c_puct, 0, max_depth)
visits = np.array(
[root.children[a].N for a in range(N_ACTIONS)], dtype=np.float64
)
q_values = np.array(
[
root.children[a].reward + GAMMA * root.children[a].Q
for a in range(N_ACTIONS)
]
)
policy = visits / max(visits.sum(), 1.0)
return float(q_values.max()), q_values, policyReading the search code top to bottom, expand_node applies prediction and then dynamics for every action in a batch. select_action ranks the parent-edge return child.reward + GAMMA * child.Q plus exploration; simulate backs up the same reward/value convention. This compact tutorial eagerly creates all children and initializes their values; it is not the exact optimized tree implementation used in the MuZero paper, which expands a selected leaf per simulation.
Worked example: planning from a single cell
Take the cell furthest from the goal, . Six moves reach the absorbing goal, so the reward arrives on transition six and has discount exponent five: . We can see this by expanding the return:
where the reward is 1 only on the final step into the goal and 0 on every other step. The behaviour policy takes a longer, less direct route, so its average discounted return from that cell is lower than the optimal value. Let us see what search produces.
The six transitions contribute rewards at discount exponents zero through five. Only the final transition has nonzero reward.
goal_steps = np.arange(0, 6)
discount_weights = GAMMA**goal_steps
rng_demo = np.random.default_rng(2)
demo_obs = observe((3, 3), rng_demo)
demo_value, demo_q, demo_policy = run_mcts(
VE_MODEL, demo_obs, n_sims=64, max_depth=4
)Search value from (3,3): 0.741
Optimal value from (3,3): 0.774
Behaviour value from (3,3): 0.700
Root action values: {'up': 0.741, 'right': 0.606, 'down': 0.653, 'left': 0.717}
Search policy: {'up': 0.219, 'right': 0.0, 'down': 0.016, 'left': 0.766}The printed visit-count policy chooses left most often, which is one optimal action from . Its estimated value of 0.741 is below the exact 0.774 but above this dataset's empirical behavior estimate of 0.700. This four-step search cannot reach the goal from the far corner: its estimate combines predicted rewards along searched edges with a learned value bootstrap, rather than observing the sixth-step goal reward. Search can select a good action while remaining imperfectly calibrated. The paired-budget evaluation below tests all grid cells.
Does search get better with more simulations?
The natural question is how this finite search behaves as its simulation budget grows. We compare the Q-greedy root action, , with an exact optimal action and compute root-value RMSE against . This diagnostic is distinct from selecting an action from MuZero's root visit-count policy. The same noisy observation for each state is reused across budgets, making their comparison paired. Neither metric is guaranteed to improve monotonically in an approximate learned model.
def evaluate_planning(
model, sims_list=(1, 4, 16, 64), n_noise=4, max_depth=4, seed=3
):
rng_local = np.random.default_rng(seed)
optimal_sets = [
set(np.flatnonzero(Q_STAR[s] >= Q_STAR[s].max() - 1e-9))
for s in range(N_STATES)
]
paired_inputs = [
(s, observe((s // N, s % N), rng_local))
for s in range(N_STATES)
for _ in range(n_noise)
]
results = {}
for n_sims in sims_list:
matches, values, targets = 0, [], []
for s, o in paired_inputs:
root_value, q_vals, _ = run_mcts(
model, o, n_sims=n_sims, max_depth=max_depth
)
matches += int(int(np.argmax(q_vals)) in optimal_sets[s])
values.append(root_value)
targets.append(V_STAR[s])
rmse = float(
np.sqrt(np.mean((np.array(values) - np.array(targets)) ** 2))
)
results[n_sims] = (matches / len(paired_inputs), rmse)
return results
PLANNING_RESULTS = evaluate_planning(VE_MODEL)
SEARCH_VALUE_RMSE = {k: metrics[1] for k, metrics in PLANNING_RESULTS.items()}
The bars show Q-greedy optimal-action agreement falling from about 0.92 at one simulation to about 0.91 at 64, while value RMSE improves from 0.087 to 0.075. The intermediate budgets are not monotonic: agreement peaks near 0.98 at four simulations, and RMSE is lowest at 16. More search improved one diagnostic but worsened the other in the one-to-64 comparison. Approximate model/value errors can change Q rankings even when value calibration improves; this does not by itself measure the action selected from root visit counts. These 64 observations are too few to establish a general budget law.
Value Equivalence and Task Sufficiency
The tutorial's decision-target model illustrates the motivation for value equivalence, but we have not proven that it satisfies a formal value-equivalence criterion. Here is the criterion and the distinction it makes.
In Grimm et al.'s definition, two models and are value equivalent relative to a selected policy set and function set when their Bellman operators agree on every selected pair: for all and . This is a statement about one-step Bellman updates on chosen test functions, not automatically equality of every policy value. Full model equivalence additionally requires agreement on transition kernels and reward functions.
Model equivalence asks for a faithful transition/reward description. Value equivalence asks for agreement only on selected decision-relevant Bellman tests. A model can differ in transitions to states that a chosen function set assigns the same value, yet agree on those Bellman updates. Whether that suffices for planning depends on the policies and functions included; a narrow test set gives a weaker guarantee.
Here is the expected immediate reward plus discounted expected at the next state under policy in model . If the selected function set is sufficiently rich and closed under relevant Bellman iterations, agreement can imply agreement of corresponding policy values. Equality of all is not the paper's Definition 1.
Grimm and colleagues use this criterion to study model classes sufficient for selected planning calculations. The lesson is not that prediction accuracy is irrelevant, but that the quantities needed by the planner determine which predictive distinctions matter. A model that reconstructs pixels accurately can still misrank actions; one that omits pixels can still support good decisions on its tested distribution.
Reconstruction versus decision-target representations
Consider two models of our gridworld:
- Model A has an observation decoder and is trained to reconstruct all 20 observation coordinates, including the four unpredictable bits.
- Model B produces an 8-dimensional latent with no interpretation whatsoever, and no decoder.
For a decoder of the current encoded observation, retaining the noise bits can reduce reconstruction loss. For a predictor of future independent bits, no deterministic estimate can eliminate expected squared error. These are distinct tasks; the plot below tests how much current-bit information each encoder retains, not whether future noise is forecast.
Model B has no direct reconstruction incentive, but its latent can still encode distractor bits incidentally. Consistency may reduce that information without guaranteeing its removal. We compare trained probes rather than infer what happened from the objectives alone.
This decoder comparison does not test formal value equivalence. For a Bellman example, imagine two gridworld models with the same position transitions and rewards but different distributions for the next observation's independent noise bits. Their full observation-transition kernels differ. Yet for policies and value functions that depend only on position, both models give the same immediate reward and expected next-position value, hence the same selected Bellman updates. That restricted agreement illustrates value equivalence without claiming that either trained network above satisfies the formal criterion.
Task sufficiency
The practical goal is task sufficiency: the latent should carry enough information for the policies and values the agent needs. In this fixed-goal gridworld, position matters, while the four independent distractor bits do not affect optimal actions. The goal location is constant, so it need not be separately encoded in every observation.
Sufficiency is relative to a task, state distribution, and decision rule. A representation adequate for navigation may omit information needed for contact-rich manipulation. A compact decision-centric latent can be efficient for its training objective, but transfer to a new reward or environment must be tested.
There is a relationship to bisimulation, discussed in Part III: Representing Agents and Worlds. Under usual MDP assumptions, bisimilar states have equal immediate rewards for each action and matching transition probabilities over equivalence classes. Such a quotient preserves values for policies compatible with the abstraction. Value equivalence asks only for selected Bellman tests and can therefore be weaker.
Measuring what the latent keeps
We can test all of this empirically. We train a second model, identical in architecture except that it also has a decoder and is trained with a reconstruction loss instead of the consistency loss. Then we ask two questions of each model:
- How large is the sampled between-position share of latent variance? A noise-invariant latent would vary little within a position, but eight sampled noise masks per cell also put sampling variation into this statistic. Decision sufficiency alone does not require noise invariance.
- How accurately can a linear probe read the noise bits out of the latent? If the latent is noise-invariant, a linear probe should be no better than chance.
These two diagnostics are complementary. The finite-sample variance share describes one aspect of latent organization, while the linear probe tests recoverable current-noise information. A model can score differently on the two measures, so we inspect both without treating either as a sufficiency test.
class ReconstructiveNet(MuZeroNet):
def __init__(
self,
obs_dim=OBS_DIM,
n_actions=N_ACTIONS,
latent_dim=LATENT_DIM,
hidden=64,
):
super().__init__(obs_dim, n_actions, latent_dim, hidden)
self.decoder = nn.Sequential(
nn.Linear(latent_dim, hidden),
nn.ReLU(),
nn.Linear(hidden, obs_dim),
)
def decode(self, z):
return self.decoder(z)
REC_MODEL = ReconstructiveNet()
hist_rec = train_model(
REC_MODEL,
OBS,
ACT,
REW,
NXT,
MC_RETURNS,
BEHAVIOUR_POLICY,
epochs=200,
use_consistency=False,
use_reconstruction=True,
)@torch.no_grad()
def latent_position_variance(model, n_noise=8, seed=7):
"""Finite-sample between-position share of latent variance."""
rng_local = np.random.default_rng(seed)
Z = np.zeros((N_STATES, n_noise, model.latent_dim))
for s in range(N_STATES):
p = (s // N, s % N)
for k in range(n_noise):
o = observe(p, rng_local)
Z[s, k] = model.encode(torch.tensor(o).unsqueeze(0)).numpy()[0]
grand_mean = Z.reshape(-1, model.latent_dim).mean(axis=0)
between = ((Z.mean(axis=1) - grand_mean) ** 2).sum(axis=1).mean()
within = ((Z - Z.mean(axis=1, keepdims=True)) ** 2).sum(axis=2).mean()
return float(between / (between + within))
def distractor_probe_accuracy(model, obs, n_train=2000, seed=0):
"""Linear probe accuracy for recovering the noise bits from the latent."""
with torch.no_grad():
Z = model.encode(torch.tensor(obs)).numpy()
Y = obs[:, N_STATES:]
perm = np.random.default_rng(seed).permutation(len(Z))
tr, te = perm[:n_train], perm[n_train:]
A = np.concatenate([Z[tr], np.ones((len(tr), 1))], axis=1)
W, *_ = np.linalg.lstsq(A, Y[tr], rcond=None)
A_te = np.concatenate([Z[te], np.ones((len(te), 1))], axis=1)
pred = (A_te @ W) > 0.5
return float((pred == (Y[te] > 0.5)).mean())
VAR_VE = latent_position_variance(VE_MODEL)
VAR_REC = latent_position_variance(REC_MODEL)
ACC_VE = distractor_probe_accuracy(VE_MODEL, OBS)
ACC_REC = distractor_probe_accuracy(REC_MODEL, OBS)

The decision-target model has a higher sampled between-position variance share, but its distractor probe remains above chance in this run. We drew eight independent noise masks per cell from the 16 possibilities, with repeats possible, so random mask variation can enter the between-cell term; this is not an unbiased estimate of variance caused by position alone. The reconstructive model retains still more current-bit information, consistent with its decoder objective. The comparator also changes decoder architecture and independently initializes the shared layers, so the gap cannot be attributed to the auxiliary objective alone. Probe test observations are held out from the linear probe fit, not from representation training. This comparison is suggestive, not a proof that consistency eliminates noise or that one objective dominates for every task.
This is not an argument that reconstruction is bad. It provides supervision when rewards are sparse and can preserve information for tasks not known during training. The two tutorial objectives emphasize different signals; whether either yields better control or transfer requires a matched downstream evaluation that this probe does not provide.
Search Targets and Reanalysis
Search consumes the learned model and supplies policy targets for later training. Better targets can improve the model, but this feedback loop is not guaranteed to converge or improve control under approximation.
Write for a search procedure that uses and to produce a policy distribution and a root-value estimate. MuZero trains its policy head against the search distribution, while its value target also uses observed outcomes or -step returns. Our later toy comparison directly trains on search root values as a separate experimental simplification.
Self-consistency, , would mean the value head matches the particular search estimates on tested states. It would not prove optimality. Classical tabular value iteration has contraction guarantees under specific assumptions; finite tree search plus a learned model and function approximation does not inherit them automatically.
MuZero actors run search while selecting real actions and store trajectories with search statistics in replay. Reusing old search-policy targets creates staleness as the network changes. Reanalyzing stored observations with newer parameters can refresh those policy targets, at additional compute cost.
Reanalysis
The MuZero Reanalyze variant searches stored states again with a newer network to provide fresh policy targets for a reported 80% of updates. Its value target uses observed -step rewards bootstrapped from a target network's later value estimate. It does not simply replace every value label with the current root search value.
Reanalysis trades compute for more current targets. Its benefit is empirical and depends on the data regime, target network, and model quality; it is not a standalone guarantee of sample efficiency.
Training a value head directly on its own search estimates can reinforce an overestimate. Observed rewards ground the sampled transitions, while target-network bootstrapping supplies another learned estimate that may stabilize targets; neither certifies unseen imagined actions or rules out search exploiting inaccurate off-data transitions. The original MuZero objective has no explicit latent-consistency term.
The toy experiment below intentionally tests the more direct search-value target. Treat any outcome as a property of this fixed dataset and optimizer, not as a reproduction of the paper's Reanalyze value-target construction.
An experiment: same data, different targets
We can compare target sources while holding data and optimization fixed. Take the trained decision-target model and fine-tune it twice, on exactly the same data, for exactly the same number of epochs:
- Control run. Fine-tune on Monte Carlo returns and behaviour-policy action targets.
- Reanalysis run. Fine-tune on search values and search visit-count policies.
The runs start from matched model copies and use the same data, optimizer settings, gradient steps, and seed. Value and policy targets change together, so this intervention cannot identify either target's separate effect. If the search-target model does not improve, that is a useful counterexample to assuming that a more accurate raw target is always internalized by this training loop.
rng_re = np.random.default_rng(11)
REANALYZE_IDX = rng_re.permutation(len(OBS))
SEARCH_VALUES = np.zeros(len(REANALYZE_IDX))
SEARCH_POLICIES = np.zeros((len(REANALYZE_IDX), N_ACTIONS))
for i, j in enumerate(REANALYZE_IDX):
v, _, pi = run_mcts(VE_MODEL, OBS[j], n_sims=24, max_depth=4)
SEARCH_VALUES[i] = v
SEARCH_POLICIES[i] = pi
MC_TARGETS = MC_RETURNS[REANALYZE_IDX]
POLICY_TARGETS = BEHAVIOUR_POLICY[REANALYZE_IDX]
SUB_OBS = OBS[REANALYZE_IDX]
SUB_ACT = ACT[REANALYZE_IDX]
SUB_REW = REW[REANALYZE_IDX]
SUB_NXT = NXT[REANALYZE_IDX]MODEL_CONTROL = copy.deepcopy(VE_MODEL)
hist_control = train_model(
MODEL_CONTROL,
SUB_OBS,
SUB_ACT,
SUB_REW,
SUB_NXT,
MC_TARGETS,
POLICY_TARGETS,
epochs=80,
seed=1,
)
MODEL_REANALYZED = copy.deepcopy(VE_MODEL)
hist_reanalysis = train_model(
MODEL_REANALYZED,
SUB_OBS,
SUB_ACT,
SUB_REW,
SUB_NXT,
SEARCH_VALUES,
SEARCH_POLICIES,
epochs=80,
seed=1,
)
ERR_CONTROL = value_rmse(MODEL_CONTROL)
ERR_REANALYSIS = value_rmse(MODEL_REANALYZED)
ERR_SEARCH_TARGETS = float(
np.sqrt(np.mean((SEARCH_VALUES - V_STAR[STATE_IDS[REANALYZE_IDX]]) ** 2))
)
with torch.no_grad():
replay_observations = torch.tensor(SUB_OBS)
_, replay_head_values = MODEL_REANALYZED.predict(
MODEL_REANALYZED.encode(replay_observations)
)
replay_reference = V_STAR[STATE_IDS[REANALYZE_IDX]]
ERR_HEAD_REPLAY = float(
np.sqrt(np.mean((replay_head_values.numpy() - replay_reference) ** 2))
)
REPLAY_GOAL_COUNT = int(np.count_nonzero(STATE_IDS[REANALYZE_IDX] == 0))
Raw search RMSE on replay observations: 0.040 Fine-tuned head RMSE on the same replay: 0.039 Absorbing-goal replay rows: 2197/3000 Behaviour-policy cell-mean value RMSE: 0.105
Under equal weighting of the 16 grid-cell means, search-target fine-tuning lowers the direct head's error to 0.072, compared with 0.095 for behaviour-target fine-tuning and 0.104 for the base model. The raw search values were instead computed on replay observations. On those same observations, raw search RMSE is 0.040 and the fine-tuned head RMSE is 0.039. Their ordering therefore cannot be inferred from the unequal-weighted cell-mean chart. The replay metric also weights the absorbing goal heavily (2,197 of 3,000 rows), so it should not stand in for uniform state coverage. This single-run joint value-and-policy target change improves cell-mean value accuracy, but is not a general convergence or action-quality guarantee.
This distinction separates the quality of a target generator from how well a parameterized network learns its targets. The experiment cannot justify dropping search at deployment; that would require direct action-quality evaluation without search.
Limitations and Impact
MuZero shows that observation reconstruction is not necessary for its reported game-playing results. The learned model still needs sufficiently accurate reward, value, and policy predictions on states relevant to search. The key parameters in this tutorial are:
- LATENT_DIM: The dimensionality of the learned latent state. More capacity may allow the model to retain more detail, but does not force it to do so; depending on architecture, data, and regularization, extra dimensions can also make dynamics learning more difficult.
- GAMMA: The discount factor, weighting near rewards more heavily than distant ones in both value targets and search backup.
- n_sims: The number of tree-search simulations per decision. More simulations cost more compute and can amplify model error; improvement is not guaranteed.
- c_puct: The exploration constant in the PUCT rule, balancing exploitation of high-value actions against exploration guided by the network prior.
- max_depth: The maximum unrolled planning depth, limiting how far the learned dynamics are trusted during search.
These parameters interact, but there is no general rule that increasing LATENT_DIM requires more simulations or that deeper search requires a lower c_puct. Our values were illustrative rather than tuned; practical choices should be validated against held-out environment outcomes.
The first limitation is interpretability. The 8-dimensional latent in this run is more strongly associated with position than the reconstructive baseline, but it also retains some distractor information. A larger model need not admit a simple semantic interpretation. A decoder can provide an observation-space diagnostic, though a plausible reconstruction is not proof of reliable decisions.
The second limitation is model exploitation. Search may favor imagined branches with inflated predicted rewards or values, especially away from replayed actions. Observed-reward supervision constrains sampled transitions but does not certify all branches; auxiliary consistency is neither part of original MuZero nor a cure-all. Conservative planning, shorter horizons, uncertainty estimates, and real-environment evaluation can help diagnose or limit exploitation.
The third limitation is compute. MuZero performs search when actors select environment actions; the Reanalyze variant also searches selected replay states, not every gradient update. EfficientZero adds a self-supervised component and a value-prefix design for the Atari 100k data regime. Sampled MuZero samples actions for larger or continuous action spaces, while Gumbel MuZero changes root action selection and analyzes policy improvement under its own assumptions. These methods address different costs and settings, not a single universal ranking.
The fourth limitation is stochasticity. Original MuZero uses a deterministic latent transition even when observations or outcomes may be stochastic; it can learn useful expected predictions without explicitly branching on chance, but cannot represent every outcome distribution in that transition. Stochastic MuZero introduces learned chance outcomes and a more involved search/training design.
The fifth limitation is generalization and reuse. A representation trained only for selected reward and value tests may omit information needed after the task changes. A reconstructive objective can encourage broader observation information, but it does not guarantee transfer either. New tasks may require retraining either representation; compare them empirically before assuming one transfers better.
MuZero's contribution is a demonstrated way to couple a learned latent model with online search and search-generated training targets without pixel reconstruction. Related methods explore different action spaces, data budgets, and uncertainty assumptions. The next chapter, TD-MPC and Control-Centric Representations, keeps the decision-centric latent objective but uses short-horizon model-predictive control rather than this tree-search design.
Summary
MuZero trains a model on planning-relevant reward, value, and policy targets and evaluates the resulting decisions. Its architecture has three functions: a representation that encodes available history into an initial latent state (our tutorial uses ), a dynamics that advances latent states under actions and predicts rewards via , and a prediction that reads a policy and a value out of a latent state via . There is no decoder and no reconstruction loss in the original objective.
- The original MuZero latent is not supervised to reconstruct observations. Reward, value, and policy targets constrain its decision-relevant predictions; the tutorial adds an auxiliary consistency loss and uses only one-step supervision.
- Latent consistency is optional. Original MuZero does not include it; EfficientZero and our tutorial use variants. The toy probe shows reduced but nonzero distractor information, so it did not eliminate all four noise bits.
- Planning happens in latent space. Search uses predicted edge rewards and child values alongside priors. In this learned model, increasing the simulation budget worsened Q-greedy action agreement while improving root-value accuracy; that diagnostic does not directly measure visit-policy action selection.
- Value equivalence is relative to selected tests. Grimm et al. define it through Bellman-operator agreement for selected policies and value functions, not blanket equality of all policy values. Reconstruction and decision-target objectives preserve different information.
- Search targets need careful use. Published MuZero uses search policy targets plus outcome or bootstrapped-return value targets; Reanalyze refreshes selected policy targets. In this toy experiment, jointly changing search-value and search-policy targets improves equally weighted cell-mean head RMSE from 0.104 to 0.072. A separate replay-matched check gives raw-search and tuned-head RMSE of 0.040 and 0.039, respectively; those replay rows overrepresent the goal.
- The failure modes are real. Uninterpretable latents, model exploitation through overestimation, expensive search, and an assumption of deterministic dynamics all require care. The lineage that followed MuZero, from EfficientZero to Sampled and Gumbel MuZero, is largely a catalogue of responses to these specific weaknesses.
Quiz
Ready to test your understanding? Take this quick quiz to reinforce what you've learned about MuZero, latent planning, value equivalence, and search-generated targets.
MuZero and Value-Equivalent Models Quiz
Reference
Citation details
Cite or share this article.
Continue with the full handbook
This chapter is part of World Models Handbook. Use the handbook page to browse the complete table of contents and continue reading in sequence.
Explore World Models HandbookStay up to date
Get articles, book updates, and news delivered to your inbox.
No spam, unsubscribe anytime.
Join the community
Sign in to remove popups, track your reading progress, and join the discussion.

Comments
No comments yet. Be the first to share your thoughts!