MuZero and Value-Equivalent Models

Michael BrenndoerferJuly 16, 202657 min read

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 t=0,1,2,…t = 0, 1, 2, \ldots. 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 sks^k is not the environment's state; it is a vector of the network's own invention. Throughout this chapter we distinguish:

  • the observation oto_t (what the agent sees at wall-clock time tt);
  • the action ata_t (what the agent chooses at wall-clock time tt);
  • the reward utu_t (what the environment returns at wall-clock time tt);
  • the policy π\pi (a distribution over actions);
  • the transition p(st+1∣st,at)p(s_{t+1}\mid s_t, a_t) (how the environment's state evolves);
  • the observation model p(ot∣st)p(o_t\mid s_t) (how observations are generated from states);
  • the latent state sks^k and unrolled horizon kk 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 kk counts steps taken inside the imagination of the network. The subscript tt counts steps taken in the world. In our tutorial, a single observation oto_t is encoded into the initial latent state s0s^0. 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 tt as oto_t (what the agent sees), the action as ata_t (what the agent chooses), the reward as utu_t (what the environment returns), and the latent state as sks^k (the compressed summary the network invents at model step kk). The superscript kk indexes steps after the encoded root, not wall-clock time. In the training equations below, KK is the replay-supervised unroll length; search depth is a separate choice. MuZero parameterizes three functions with a shared network θ\theta (layers within each function have their own parameters, conventionally grouped as θ\theta). 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:

s0=hθ(ot),(sk+1,rk+1)=gθ(sk,ak),(pk,vk)=fθ(sk).\begin{aligned} s^0 &= h_\theta(o_t), \\ (s^{k+1}, r^{k+1}) &= g_\theta(s^k, a^k), \\ (p^k, v^k) &= f_\theta(s^k). \end{aligned}

where:

  • oto_t: the observation the agent sees at time tt
  • hθh_\theta: the representation function, which encodes available history into an initial latent state s0s^0; the displayed hθ(ot)h_\theta(o_t) is our single-observation tutorial specialization
  • aka^k: the action taken at unrolled step kk (encoded as a one-hot or embedding)
  • gθg_\theta: the dynamics function, which maps a latent state and an action to the next latent state sk+1s^{k+1} and the reward rk+1r^{k+1} on the transition into it
  • sks^k: the latent state at unrolled step kk; the tutorial implementation below bounds it with a tanh activation
  • fθf_\theta: the prediction function, which reads a latent state sks^k and emits a policy pkp^k (a distribution over actions, used as a prior during search) and a value vkv^k (the expected discounted return from that latent state)
  • k∈{0,1,…,K}k \in \{0, 1, \ldots, K\}: the model-step index within the training unroll, not wall-clock time
  • KK: 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 sks^k 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.

Latent state

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 tanh⁡\tanh; 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 tt. We run the network forward from s0=hθ(ot)s^0 = h_\theta(o_t), unrolling KK steps with the actual actions the agent took:

s0→  at  (s1,r1)→  at+1  (s2,r2)→  at+2  ⋯→  at+K−1  (sK,rK).\begin{aligned} s^0 &\xrightarrow{\;a_t\;} (s^1, r^1) \\ &\xrightarrow{\;a_{t+1}\;} (s^2, r^2) \\ &\xrightarrow{\;a_{t+2}\;} \cdots \\ &\xrightarrow{\;a_{t+K-1}\;} (s^K, r^K). \end{aligned}

At the root and each unrolled state we predict value vkv^k and policy pkp^k; each modeled transition predicts reward rk+1r^{k+1}. The reward target is the observed reward on that transition. The policy target πt+k\pi_{t+k} comes from search at the corresponding real state. In the original MuZero training objective, the value target zt+kz_{t+k} is the final outcome for board games or an observed nn-step return bootstrapped from a later search value for Atari. With our reward indexing, a schematic loss is:

L(θ)=Eτ∼D[∑k=1Kℓr(ut+k,rk)+∑k=0K(ℓv(zt+k,vk)+ℓp(πt+k,pk))]+c∥θ∥22\mathcal{L}(\theta) = \mathbb{E}_{\tau \sim \mathcal{D}} \left[\sum_{k=1}^{K} \ell^{r}(u_{t+k}, r^k) + \sum_{k=0}^{K} \bigl(\ell^{v}(z_{t+k}, v^k) + \ell^{p}(\pi_{t+k}, p^k)\bigr)\right] + c \lVert \theta \rVert_2^2

where:

  • τ∼D\tau \sim \mathcal{D}: a trajectory sampled from the replay buffer D\mathcal{D}, from which the starting index tt is drawn
  • L(θ)\mathcal{L}(\theta): the total training loss, minimized over the shared parameters θ\theta
  • KK: the number of model steps in this training unroll
  • ut+ku_{t+k}: the observed reward on the kkth transition after the starting state, for k=1,…,Kk=1,\ldots,K
  • rkr^k: the reward predicted by dynamics on that transition; no root reward r0r^0 is predicted
  • zt+kz_{t+k}: the outcome or nn-step bootstrapped value target appropriate to the domain
  • vkv^k: the value predicted by the prediction function at unrolled step kk
  • πt+k\pi_{t+k}: the policy target produced by the search at time t+kt+k
  • pkp^k: the policy prior predicted by the prediction function at unrolled step kk
  • ℓr\ell^{r}: the reward loss, whose form depends on the domain and target representation
  • ℓv\ell^{v}: the value loss, likewise chosen for the domain and target representation
  • ℓp\ell^{p}: the policy loss, a cross-entropy between the search visit-count distribution πt+k\pi_{t+k} and the network prior pkp^k
  • cc: the coefficient of the L2L_2 weight-regularization term
  • ∥θ∥22\lVert \theta \rVert_2^2: the squared L2L_2 norm of all parameters in the shared network θ\theta, 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 s~k+1\tilde s^{k+1} for the state component of gθ(sk,ak)g_\theta(s^k,a^k); then:

ℓcons=∥ s~k+1−sg⁡ ⁣(hθ(ot+k+1)) ∥22\ell^{\text{cons}} = \big\lVert\, \tilde s^{k+1} - \operatorname{sg}\!\big(h_\theta(o_{t+k+1})\big) \,\big\rVert_2^2

where:

  • ℓcons\ell^{\text{cons}}: the latent consistency loss at unrolled step kk
  • s~k+1\tilde s^{k+1}: the state component of the dynamics prediction after applying aka^k at sks^k
  • sg⁡(⋅)\operatorname{sg}(\cdot): the stop-gradient operator, which treats its argument as a constant during backpropagation
  • hθ(ot+k+1)h_\theta(o_{t+k+1}): the latent state the representation function would have produced had it seen the observation k+1k+1 steps later
  • kk: the unrolled step at which the consistency loss is computed, ranging over 0,…,K−10, \ldots, K-1

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 tanh⁡\tanh. 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.

In[3]:
Code
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 4×44 \times 4 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.

In[4]:
Code
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.

In[5]:
Code
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.

In[6]:
Code
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)
    ]
)
Out[7]:
Console
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 γ=0.95\gamma = 0.95 weights earlier rewards more heavily. Let ss denote an environment state, rt+1r_{t+1} the reward received after moving from sts_t, and π\pi the behaviour policy. The expected discounted return is

Vπ(s)=Eπ ⁣[∑i=0∞γi rt+i+1 ∣ st=s]V^\pi(s) = \mathbb{E}_\pi\!\left[\sum_{i=0}^{\infty} \gamma^{i}\, r_{t+i+1} \,\Big|\, s_t = s\right]

where:

  • Vπ(s)V^\pi(s): the expected discounted return when starting in state ss and following π\pi thereafter
  • Eπ[⋅]\mathbb{E}_\pi[\cdot]: the expectation over actions drawn from π\pi and over the environment's transitions
  • rt+i+1r_{t+i+1}: the reward received at step t+i+1t+i+1
  • γ∈[0,1)\gamma \in [0, 1): the discount factor, weighting rewards ii steps away by γi\gamma^i

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 VπV^\pi targets for late-episode transitions that have not reached the goal. The V⋆V^\star 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 tanh⁡\tanh 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.

In[8]:
Code
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 hθh_\theta, the next_state method implements gθg_\theta and splits its output into a latent half and a scalar reward half, and the predict method implements fθf_\theta 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.

In[9]:
Code
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
)
Out[10]:
Visualization
Log-scale line chart: reward and value losses decline, policy loss stays near 1.23, and latent-consistency loss varies.
Training losses for the one-step decision-target model over 200 epochs. Reward and value losses decrease; policy loss plateaus near 1.23 and latent consistency is nonmonotone. The objective contains no observation-reconstruction term.

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.

In[11]:
Code
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()
    )
Out[12]:
Visualization
Overlaid histograms compare trained and untrained latent-prediction errors, with substantial overlap.
Per-transition squared latent-consistency errors for the trained model and one coherent untrained model. This is an error in learned coordinates, not raw observation space; raw distractor variance is not a lower bound.

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.

In[13]:
Code
@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)))
Out[14]:
Console
Decision-target model value RMSE: 0.104
Empirical behaviour-value RMSE: 0.105
Out[15]:
Visualization
Grouped bar chart comparing optimal, behaviour-policy, and learned model values for 16 grid cells.
Value predictions by grid cell, sorted by optimal value. On this run, the learned model's RMSE against infinite-horizon optimal values is marginally lower than the empirical 20-step behaviour-return baseline's; this comparison includes truncation and sampling effects.

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 gθg_\theta 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 R8\mathbb{R}^8 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 zz, the reward that the dynamics function predicts on the transition into that node, a prior probability p(a)p(a) supplied by the prediction network, and running statistics. We write N(z,a)N(z, a) for the number of times action aa has been tried at a node holding zz, W(z,a)W(z, a) for the sum of values backed up through those trials, and Q(z,a)=W(z,a)/N(z,a)Q(z, a) = W(z, a) / N(z, a) for the resulting mean.

When the notation is clear from context we drop the argument zz and write NN, WW, and QQ for a single node. The full form N(z,a)N(z, a) is used in the PUCT rule below, where zz identifies the node and aa 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:

a⋆=arg⁡max⁡a[Q(z,a)+cpuct⋅p(a)⋅∑bN(z,b)1+N(z,a)]a^\star = \arg\max_a \left[ Q(z, a) + c_{\text{puct}} \cdot p(a) \cdot \frac{\sqrt{\sum_b N(z, b)}}{1 + N(z, a)} \right]

where:

  • Q(z,a)Q(z, a): the mean value of taking action aa at latent state zz, equal to W(z,a)/N(z,a)W(z, a) / N(z, a)
  • N(z,a)N(z, a): the visit count for action aa at zz
  • p(a)p(a): the prior probability the prediction network assigns to action aa
  • ∑bN(z,b)\sum_b N(z, b): the total visits to zz across all actions
  • cpuctc_{\text{puct}}: a constant controlling exploration strength

The displayed rule is schematic. Its edge mean Q(z,a)=W(z,a)/N(z,a)Q(z,a)=W(z,a)/N(z,a) 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 p(a)p(a). At fixed total node visits, its denominator 1+N(z,a)1 + N(z,a) 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 QQ.

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.

In[16]:
Code
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
]
Out[17]:
Visualization
Line chart of PUCT exploration bonus decaying with visit count for three priors.
PUCT exploration bonus as a function of an action's visit count, for three prior probabilities. The bonus shrinks rapidly as an action is tried and starts higher for larger priors, showing how the network's prior shapes early exploration within a fixed simulation budget.

The prior p(a)p(a) 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 zi+1z_{i+1} follows action aia_i from ziz_i, the state-return recursion is:

G(zi)=r(zi,ai)+γ G(zi+1),G(z0)=∑i=0H−1γir(zi,ai)+γHv(zH),\begin{aligned} G(z_i) &= r(z_i,a_i) + \gamma\,G(z_{i+1}), \\ G(z_0) &= \sum_{i=0}^{H-1}\gamma^i r(z_i,a_i) + \gamma^H v(z_H), \end{aligned}

where:

  • G(zi)G(z_i): the backed-up discounted return from the node holding latent state ziz_i
  • r(zi,ai)r(z_i,a_i): the reward predicted by dynamics on the edge from ziz_i to zi+1z_{i+1}
  • γ\gamma: the discount factor, which weights nearer rewards more heavily
  • HH: the number of simulated transitions before the value bootstrap
  • v(zH)v(z_H): 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 π(a)=N(a)/∑bN(b)\pi(a) = N(a) / \sum_b N(b) 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 nn-step returns.

Model-based policy improvement

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.

In[18]:
Code
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, policy

Reading 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, (3,3)(3,3). Six moves reach the absorbing goal, so the reward arrives on transition six and has discount exponent five: γ5=0.955≈0.774\gamma^5 = 0.95^5 \approx 0.774. We can see this by expanding the return:

G(z0)=r(z0,a0)+γ G(z1)=0+γ (0+γ (0+γ (0+γ (0+γ 1))))=γ5=0.955≈0.774.\begin{aligned} G(z_0) &= r(z_0,a_0) + \gamma\, G(z_1) \\ &= 0 + \gamma\,(0 + \gamma\,(0 + \gamma\,(0 + \gamma\,(0 + \gamma\,1)))) \\ &= \gamma^5 = 0.95^5 \approx 0.774. \end{aligned}

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.

In[19]:
Code
goal_steps = np.arange(0, 6)
discount_weights = GAMMA**goal_steps
Out[20]:
Visualization
Bar chart of discount factors gamma^0 through gamma^5; the final goal reward receives gamma^5.
Discount weights for the six transitions from the far corner. The only nonzero reward arrives on transition six, with exponent five and optimal return 0.774.
In[21]:
Code
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
)
Out[22]:
Console
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 (3,3)(3,3). 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, arg⁡max⁡aQ(a)\arg\max_a Q(a), with an exact optimal action and compute root-value RMSE against V⋆V^\star. 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.

In[23]:
Code
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()}
Out[24]:
Visualization
Bar chart of Q-greedy optimal-action agreement across 1, 4, 16, and 64 simulations, with corresponding value RMSE labels.
Q-greedy optimal-action agreement and root-value RMSE over 64 paired noisy observations (four per state) at four search budgets. Agreement uses argmax root Q, not the visit-count policy; these are measured outcomes, not a monotonic-improvement guarantee.

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.

Value equivalence

In Grimm et al.'s definition, two models mm and m′m' are value equivalent relative to a selected policy set Π\Pi and function set V\mathcal V when their Bellman operators agree on every selected pair: Tmπv=Tm′πvT_m^\pi v = T_{m'}^\pi v for all π∈Π\pi\in\Pi and v∈Vv\in\mathcal V. 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 Tmπv(s)T_m^\pi v(s) is the expected immediate reward plus discounted expected vv at the next state under policy π\pi in model mm. 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 VπV^\pi 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:

  1. 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.
  2. 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.

In[25]:
Code
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,
)
In[26]:
Code
@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)
Out[27]:
Visualization
Bar chart comparing sampled between-position latent variance shares for two models.
Finite-sample between-position share of latent variance in two tutorial models, using eight sampled noise masks per cell. The decision-target model has a higher share; sampling noise can also contribute to this statistic.
Out[28]:
Visualization
Bar chart of noise-recovery probe accuracy for two models with a chance line.
Linear-probe accuracy for recovering four current-observation noise bits. Both models retain recoverable distractor information above the 0.5 chance line; the reconstructive model retains more in this run.

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 S\mathcal{S} for a search procedure that uses gθg_\theta and vθv_\theta 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 nn-step returns. Our later toy comparison directly trains on search root values as a separate experimental simplification.

Self-consistency, vθ≈S(vθ)v_\theta \approx \mathcal{S}(v_\theta), 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 nn-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.

In[29]:
Code
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]
In[30]:
Code
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))
Out[31]:
Visualization
Bar chart comparing value RMSE across three training conditions.
Equally weighted grid-cell mean value RMSE against the exact optimum: base 0.104, behaviour-target fine-tune 0.095, and search-target fine-tune 0.072. This plot does not compare the head with raw search values; those use a separate replay-weighted evaluation below.
Out[32]:
Console
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 hθh_\theta that encodes available history into an initial latent state (our tutorial uses s0=hθ(ot)s^0 = h_\theta(o_t)), a dynamics gθg_\theta that advances latent states under actions and predicts rewards via (sk+1,rk+1)=gθ(sk,ak)(s^{k+1}, r^{k+1}) = g_\theta(s^k, a^k), and a prediction fθf_\theta that reads a policy pkp^k and a value vkv^k out of a latent state via (pk,vk)=fθ(sk)(p^k, v^k) = f_\theta(s^k). 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

Question 1 of 70 of 7 completed
What does the original MuZero objective require the learned latent state to do?

Comments

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

Reference

Citation details

Cite or share this article.

BIBTEXAcademic
@misc{brenndoerfer2026muzerovalue, author = {Michael Brenndoerfer}, title = {MuZero and Value-Equivalent Models}, year = {2026}, url = {https://mbrenndoerfer.com/writing/muzero-value-equivalent-models-latent-planning}, organization = {mbrenndoerfer.com}, note = {Accessed: 2026-09-30} }
APAAcademic
Michael Brenndoerfer (2026). MuZero and Value-Equivalent Models. Retrieved from https://mbrenndoerfer.com/writing/muzero-value-equivalent-models-latent-planning
MLAAcademic
Michael Brenndoerfer. "MuZero and Value-Equivalent Models." 2026. Web. September 30, 2026. <https://mbrenndoerfer.com/writing/muzero-value-equivalent-models-latent-planning>.
CHICAGOAcademic
Michael Brenndoerfer. "MuZero and Value-Equivalent Models." Accessed September 30, 2026. https://mbrenndoerfer.com/writing/muzero-value-equivalent-models-latent-planning.
HARVARDAcademic
Michael Brenndoerfer (2026) 'MuZero and Value-Equivalent Models'. Available at: https://mbrenndoerfer.com/writing/muzero-value-equivalent-models-latent-planning (Accessed: September 30, 2026).
SimpleBasic
Michael Brenndoerfer (2026). MuZero and Value-Equivalent Models. https://mbrenndoerfer.com/writing/muzero-value-equivalent-models-latent-planning

About the author

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 Handbook
Newsletter

Stay up to date

Get articles, book updates, and news delivered to your inbox.

No spam, unsubscribe anytime.

or

Join the community

Sign in to remove popups, track your reading progress, and join the discussion.