Part of World Models Handbook
Explains how TD-MPC pairs short-horizon latent planning with a learned terminal value, why decoder-free control-centric representations help, and its limits.
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
TD-MPC and Control-Centric Representations
Picture a quadruped learning to walk in a simulator. The controller must choose an action now, but the consequence that matters may arrive well beyond the next few steps. A learned model can help compare immediate consequences; rolling that imperfect model far into the future can make the comparison unreliable. TD-MPC makes the boundary explicit: plan briefly with the model, then use a learned value estimate for the rest.
One alternative learns a policy or action-value function directly from interaction without an explicit transition model. Another learns dynamics and searches through predicted trajectories. Their relative data and compute costs depend on the task and algorithm; the important distinction here is where prediction ends and value estimation begins.
Longer open-loop rollouts may expose a planner to model error and out-of-distribution states, but error need not rise monotonically on every trajectory or metric. The useful planning horizon therefore has to be measured, not deduced from a general slogan about compounding error.
TD-MPC, short for Temporal Difference Learning for Model Predictive Control (Hansen, Su, and Wang, ICML 2022), combines short-horizon latent-space planning with a terminal action-value estimate. Its reported continuous-control benchmarks show strong sample efficiency against the baselines tested. The model and critic are learned together: training unrolls predicted latent states for reward, consistency, and value losses, while the released implementation computes TD bootstrap targets from encoded observed successor states.
Two commitments define the method, and both are worth stating plainly before any notation appears.
- The model should be control-centric. TD-MPC learns a decoder-free latent representation for reward prediction, short-horizon dynamics, and value estimation. No observation-reconstruction loss requires the latent to preserve every visual detail. This resembles the decision-relevant representation discussed in MuZero and Value-Equivalent Models, now with continuous actions and sampling-based rather than tree search.
- The model should be trusted for a measured horizon. The reported TD-MPC experiments use a maximum planning horizon of five latent steps across tasks, scheduled upward from one early in training. A bootstrapped action-value estimate supplies the tail, analogous to the terminal value in an -step return.
These choices shape both the method and its limitations. A latent without a decoder need not preserve enough detail to render an observation. A short plan relies more heavily on the critic, while still depending on the model's predictions along its short rollout. Both dependencies reappear in the training objective and failure modes below.
If you have read The Dreamer Family, you already know a world model that trains its latent by reconstruction and then trains an actor inside imagined rollouts. Sampling-Based Planning and Model Predictive Control introduced the MPPI and cross-entropy shooting lineage. TD-MPC combines sampling-based planning, value bootstrapping, and a decoder-free representation. MPPI searches over actions without differentiating through a learned rollout; temporal-difference learning estimates returns beyond the short plan; the representation is trained for control-related predictions rather than pixel fidelity.
First we build the model and separate its weighted multistep loss from policy optimization. Then we assemble the latent-space planner and work through an exact terminal-value example. Finally we examine TD-MPC2's 80-task generalist, run a small executable diagnostic that may fail to learn under its tiny budget, and trace the remaining failure modes.
Throughout, keep three things separate in your head: what the model predicts, how the planner uses those predictions, and how the value function summarizes what lies beyond. Most confusion about this family of methods comes from blurring those three roles. A model that predicts poorly may still support good decisions if the planner only relies on it briefly and the value function covers the rest. A value function that looks accurate in isolation may still be exploited by a planner that searches the edges of its support. Keeping prediction, planning, and value estimation conceptually distinct is what lets you diagnose which one is failing when the agent underperforms.
Joint Latent Dynamics and Value Learning
A TD-MPC agent has five learned components that share a latent representation. The encoder, dynamics, reward, and value components are optimized with a weighted multistep objective; the policy prior has a separate objective and optimizer. The representation is shaped by control-related losses rather than a pixel-reconstruction loss.
To make this precise, let denote the observation (the raw state vector in a proprioceptive task, an image in a visual one), the action, and the scalar reward. The agent maintains a latent state and five functions of it:
- Encoder : maps an observation to a latent state, .
- Latent dynamics : predicts the next latent from the current latent and action, .
- Reward model : predicts the immediate reward, .
- Value function : estimates the discounted return from a latent state and an action.
- Policy prior : maps a latent state to an action and supplies noisy proposal trajectories to the planner.
Notice that the five functions divide naturally into two groups. The encoder, dynamics, and reward model describe the world: they say where you are, where you go next, and what you earn along the way. The value function and policy prior make decisions: one estimates how good a state-action pair is, and the other proposes actions worth trying. Joint training is what binds the two groups, because the value loss reaches back through the encoder and forces the world-modeling half to preserve whatever information the decision-making half needs.
A latent world model is decoder-free when no part of the training objective asks the latent to reproduce the observation. The latent is instead trained to predict quantities that matter for decisions: rewards (here), values (here and in MuZero), and multi-step latent consistency. TD-MPC's model is decoder-free. Dreamer's RSSM is not. The tradeoff is developed in Predictive and Generative Objectives.
The original method has deterministic latent dynamics rather than a belief distribution or stochastic transition model. Sampling candidate actions during planning does not make those latent transitions probabilistic. This distinction matters when the environment has stochastic outcomes.
The training objective
Training samples sequences of observed transitions from replay, encodes the first observation, and unrolls the latent dynamics along recorded actions. Write and . For an unroll of transitions, the original paper's model objective has three coefficient-weighted losses at each step, discounted in the loss by :
Here balance reward prediction, TD value fitting, and latent consistency; they are not equal by definition. denotes slow-moving target parameters. The consistency target encodes the observed next state, while its prediction is reached by rolling forward from the first encoded state. A collapsed constant latent can make consistency easy, but reward and value prediction generally require distinctions between states. Their losses supply task-relevant pressure; none individually guarantees a useful representation.
For example, a constant encoder and matching dynamics achieve zero consistency error, yet cannot distinguish two observations whose rewards or returns differ under the same action. In tasks where those targets vary, minimizing the other losses disfavors that collapse. This is a task-dependent argument, not a proof of convergence or a guarantee against every degenerate local minimum.
The released TD-MPC implementation computes the one-step TD target using the observed successor encoded by the online encoder. With its twin target critics and a noisy action sampled from the current policy, the target is
The current policy samples around its mean with truncated noise for this target; the paper's single- notation abstracts the released twin critics. The value loss evaluates both on predicted latent states during the multistep unroll. The bootstrap target above instead uses the next real observation. The original paper writes its TD target with in the recurrent loss; reading that notation as a claim that the released code bootstraps from the predicted successor confuses the two paths. A target network slows target drift; it does not guarantee convergence of nonlinear off-policy learning.
The bootstrap action comes from the learned policy rather than an explicit maximization over a continuous action space. This is a tractable actor-critic approximation, and its errors can affect the critic. It does not imply a strict upper bound on critic quality or that an action omitted by the prior is never represented in replay.
The released implementation optimizes the policy separately on the initial and predicted recurrent latents. It detaches these latents, freezes both critics' parameters for the actor step, and maximizes the smaller critic estimate:
This writes the released update at the level of its mean objective and omits the action sampling noise. The detached latents prevent this actor step from updating the encoder or dynamics, while the frozen critics still pass gradients through the policy's action. Neither an entropy bonus nor an action-magnitude penalty appears in the original objective. The small code example below adds an explicit action-magnitude penalty as a toy stabilizer and labels it accordingly.
The toy code below includes an episodic termination mask, , but not a time-limit truncation mask. That is a sensible convention for its finite episodes, not a term shown in the original paper's displayed TD equation. Whether a timeout should bootstrap depends on the environment's observation and episode semantics; confusing timeout with a true terminal state can bias fitted values.
Why no decoder
Dropping the decoder is not an oversight. Three arguments motivate it, and they are worth stating because they apply well beyond this one algorithm.
Sufficiency beats fidelity. If the robot's arm color carries no information about the task, dynamics, or reward, a control policy need not preserve it. It does need quantities that change the return: joint angles, velocities, contacts, and distances. Color could matter, however, if it identifies a task-relevant object or instruction. Reconstruction asks the latent to retain visual detail even when that detail is irrelevant to control. A decoder-free objective does not impose that demand. This is the same argument that motivates the predictive-embedding family covered in Predictive Embeddings, JEPA, and Energy-Based Models.
The gradient budget goes to task-related predictions. The encoder is trained through reward, latent consistency, and value losses rather than a pixel decoder. This avoids allocating a direct reconstruction loss to visual detail irrelevant to the reward. It does not prove that every retained feature matters for control, nor that reconstruction would dominate every alternative model.
It changes what you monitor. Low reconstruction error alone cannot certify good decisions. Removing the decoder also removes that potentially misleading metric, but a falling value loss is no guarantee of good control either: the critic can fit biased targets or be exploited by the planner. Closed-loop evaluation remains necessary.
Reconstruction can nevertheless provide a useful auxiliary signal when reward and value targets are weak or sparse. The original paper evaluates decoder-free learning on continuous-control tasks, including sparse-reward cases; it does not establish a universal dense-versus-sparse boundary. Whether observation reconstruction helps is an empirical question for the task and data regime.
The representation itself
The original TD-MPC model and TD-MPC2 should not be conflated here. The following normalization choices are associated with TD-MPC2, not the original paper's encoder and dynamics.
Layer normalization is used throughout TD-MPC2's MLP components. The released original TD-MPC implementation does not place it in the dynamics and reward networks as previously claimed. Normalization can improve optimization stability, but its exact effect cannot be inferred from one layer in isolation.
Simplicial normalization (SimNorm), introduced in TD-MPC2, splits the latent vector into fixed-size groups and applies a softmax within each group. Each coordinate lies in and each group sums to one. This bounds the coordinates of a normalized latent; it does not guarantee accurate or stable multistep predictions. Our toy below borrows SimNorm from TD-MPC2, so it is an illustrative hybrid rather than a faithful reproduction of the 2022 architecture.
The coordinate bound can be checked on a single deterministic vector. This demonstrates the operation of SimNorm, not a claim about rollout errors.
## SimNorm applied to one latent vector: group-wise softmax over 8 groups of 8 units.
## Seeded so the displayed vector is reproducible.
## The raw logits are given unequal group maxima so the softmax produces visible
## within-group variation rather than a near-uniform (flat) output.
simnorm_groups = 8
import numpy as np
simnorm_rng = np.random.default_rng(0)
latent_raw = simnorm_rng.normal(size=64) * 1.0
## Raise the scale of a couple of groups so some groups are peaked and others flat.
_group_scale = np.repeat(np.array([4.0, 1.5, 3.0, 0.6, 4.5, 1.0, 3.5, 0.8]), 8)
latent_raw = latent_raw * _group_scale
grouped = latent_raw.reshape(simnorm_groups, -1)
grouped = grouped - grouped.max(axis=-1, keepdims=True)
exp_grouped = np.exp(grouped)
simnorm_out = (exp_grouped / exp_grouped.sum(axis=-1, keepdims=True)).reshape(
-1
)
group_sums = simnorm_out.reshape(simnorm_groups, -1).sum(axis=-1)
coord_index = np.arange(1, simnorm_out.size + 1)
These normalization choices affect representation scale in TD-MPC2. The planner does not score candidates by latent distance: it sums predicted rewards and a terminal value. Scale can still affect optimization, but a fixed latent distance need not have a fixed semantic meaning during training.
Finally, note what is not in the original model: no observation decoder, stochastic latent transition, or dynamics ensemble. Those are architectural choices, not proofs that uncertainty is unimportant.
Temporal-Difference Model Predictive Control
We now have a model that predicts rewards and latents, and a value function that estimates returns. The remaining question is how to convert them into an action at every control step. TD-MPC uses a receding-horizon sampling planner, run entirely inside latent space, with the value function supplying the tail of the objective.
Recall from Sampling-Based Planning and Model Predictive Control that the MPPI-style planner samples action sequences, evaluates them under the model, and refits a sampling distribution toward high-scoring candidates. It does not differentiate through the planned rollout. That avoids relying on rollout gradients, but sampling can still exploit errors in the learned model or critic.
The planning objective
Let be the current latent, and let be a candidate action sequence over a horizon of steps. Roll the latent dynamics forward deterministically, . The planner scores each candidate by
The first term estimates rewards along a candidate model rollout; the second evaluates the tail after that rollout. The paper's Equation (3) writes the terminal term as , leaving the terminal action abstract in that notation. The released 2022 implementation instantiates it as
The terminal action is sampled from the learned policy using the configured minimum standard deviation because the candidate sequence ends at step . The two online critics are combined by a minimum; this choice alone is not a guarantee against optimistic scoring. The planner needs only a numeric score, while the policy's separate training objective uses gradients through its proposed action. The score resembles an -step bootstrapped return, but planning optimizes actions using model predictions, whereas TD learning fits a value function to targets. They are related operations, not identical algorithms.
Two error sources enter the score. Errors in reward and latent predictions can change the scores of candidate sequences; for a fixed sampled terminal action, error in the twin-critic terminal estimate contributes , where is the estimation error of at that terminal state and action. Policy sampling adds another source of score variation. These effects depend on which states the optimizer visits. A shorter horizon reduces exposure to model rollouts but puts more weight on the critic. No fixed ranking of the errors follows without measuring the model, critic, and candidate distribution.
The MPPI update
MPPI maintains a time-varying Gaussian over action sequences and improves it over a few iterations. Each iteration does four things.
- Sample. Draw action sequences from and clip each action to the environment's bounds. Add policy-guided trajectories; the original implementation injects noise into policy actions so these proposals are not duplicates.
- Score. Roll every candidate through and , accumulate the discounted predicted reward, and add the terminal value.
- Fit. Keep the top candidates (the elites) and refit:
where is the largest elite score and controls weight sharpness. Subtracting the maximum is a numerically stable softmax shift; it does not normalize score variance or make the weighting invariant to reward scale. The original planner uses this max shift, not the z-score normalization illustrated separately below. 4. Repeat for a fixed number of iterations.
In the released TD-MPC planner, the first action comes from a trajectory sampled among the final elites according to their weights. During training, exploration noise is then added using the final elite spread; evaluation omits that extra noise. The code does not clamp the action again after adding this final noise. This is distinct from the toy planner below, which samples from its refitted Gaussian and clips the result. Neither action-selection rule guarantees useful coverage when the model and critic are poor.
The next illustration asks a counterfactual question: what would z-score normalization do to weights under a reward-scale change? It is useful for understanding softmax sensitivity, but it is not a diagram of the original TD-MPC planner. The weights use deterministic return vectors, with no sampling.
## Elite weights under raw versus normalized returns at two reward scales (deterministic).
mppi_temperature = 1.5
elite_index = np.arange(1, 9)
elite_returns_small = np.linspace(1.0, 3.0, 8)
elite_returns_large = elite_returns_small * 100.0
def _softmax_weights(values, temperature):
scaled = temperature * (values - values.max())
weights = np.exp(scaled)
return weights / weights.sum()
def _normalized_weights(values, temperature):
normed = (values - values.mean()) / (values.std() + 1e-6)
return _softmax_weights(normed, temperature)
weights_raw_small = _softmax_weights(elite_returns_small, mppi_temperature)
weights_raw_large = _softmax_weights(elite_returns_large, mppi_temperature)
weights_normalized = _normalized_weights(elite_returns_large, mppi_temperature)
Policy-guided proposals can put candidates near actions the learned critic rates highly, especially in larger action spaces. The original implementation adds noise to these trajectories so repeated proposals can explore different continuations. This guidance may help but can also inherit critic errors; it is not a guarantee that prior and planner improve together.
What the model is and is not used for
This is the point where TD-MPC most sharply diverges from Dreamer, so it is worth being precise.
The model is used in training and at decision time. Training samples real trajectories from replay, then repeatedly applies learned latent dynamics along their recorded actions. Reward and value predictions on these predicted latents receive gradients through the unroll. At decision time, the planner searches over new candidate action sequences in the same learned model.
The TD bootstrap in released code uses observed successors. This is narrower than saying model predictions never enter learning: the value prediction is fitted on model-unrolled latent states, and its gradient shapes the model. Unlike Dreamer's policy-learning rollouts, these unrolls follow replayed actions rather than a newly imagined policy trajectory.
This distinction matters for diagnosing failure. A model error can affect both planning scores and value predictions during recurrent training, even when the TD target begins from an encoded observed successor. Off-policy replay permits reuse of past data but does not make arbitrary stale or poorly covered trajectories harmless.
The planner and critic remain coupled through the terminal score and policy proposals. If the critic is optimistic for candidate actions or terminal states outside replay coverage, optimization can exploit that error. The method limits rollout length but offers no general divergence guarantee.
Putting the pieces in order
The complete online loop is short enough to state in full.
- Reset the environment, encode the observation to .
- At each step, run MPPI on the current latent with the model and value function, take the sampled action, and step the environment.
- Store in the replay buffer.
- Sample replay sequences, optimize the weighted model/critic losses and the separate policy objective, then update target parameters by exponential moving average.
- Repeat for a fixed number of episodes or environment steps, using a short random-action warm-up at the start to seed the buffer.
Warm-up supplies initial replay data before relying on an untrained planner. The released loop normally runs one optimization update per collected decision transition, not per underlying simulator frame when actions repeat; our short toy groups updates differently. Neither setting establishes that the toy has learned a useful controller.
The original implementation plans at each control step, uses a short horizon and sampled action sequences, and normally takes one training update per collected decision transition. The horizon is scheduled upward during early training. TD-MPC2 uses a short horizon too, but differences in architecture, task mix, and evaluation prevent attributing its gains solely to a stronger critic or a shorter plan.
Short-Horizon Planning with Terminal Values
When can a short plan plus a terminal value outperform a longer model rollout? A good terminal value can replace many predicted steps, but its own errors then matter more. The choice is a measured tradeoff rather than a theorem that three steps beats thirty.
Over steps, an imperfect dynamics model feeds its predicted state into the next prediction. Errors may accumulate, cancel, or vary with the visited states; neither mean-squared prediction error nor return variance is guaranteed to increase at every horizon. Longer rollouts can also search more reward opportunities. Their net effect has to be evaluated on held-out trajectories and in closed-loop control.
If the critic's error at the reached terminal state-action pair is , its direct contribution to the planning score is . The discount attenuates a fixed error, but the terminal state itself changes with and may enter a region where critic error is larger. One critic evaluation is computationally cheap; that does not make its error constant in horizon.
The direct critic term shrinks with only under the fixed-error comparison; model error may grow with . Those tendencies can create an interior optimum, but they can also produce a boundary optimum or a nonmonotone curve. If the critic is weak, planning longer helps only if the model remains accurate enough and the added search cost is worthwhile. Horizon selection belongs to validation, not a rule of thumb based on critic quality alone.
The next plot deliberately assumes one increasing model-error component and one fixed critic-error magnitude discounted by . Under these chosen functions, their sum has an interior minimum. It is a hypothetical error budget, not a measured TD-MPC curve or a general upper bound on planning error.
## Illustrative error decomposition for the horizon tradeoff (deterministic, no sampling).
## Assumed model-error term rises exponentially over this finite range.
## The terminal value's bias is fixed while its weight gamma**H decays.
h_grid = np.arange(1, 16)
gamma_tradeoff = 0.9
model_scale, model_rate = 0.2, 0.2
model_error = model_scale * (np.exp(model_rate * h_grid) - 1.0)
value_bias = 1.2
value_error = (gamma_tradeoff**h_grid) * value_bias
total_error = model_error + value_error
best_h = int(h_grid[int(np.argmin(total_error))])
There is also a classical control-theoretic reading. With exact dynamics and the exact optimal cost-to-go as terminal cost, optimizing the first action in a one-step plan recovers the optimal action by the Bellman equation. This requires the optimal terminal state value , equivalently exact evaluated at an optimal terminal action; exact dynamics and with a suboptimal terminal policy do not suffice. Stability claims for approximate MPC require additional assumptions; an arbitrary terminal cost that upper-bounds cost-to-go is not by itself a sufficient blanket condition. The worked example below isolates the oracle-value case before the neural approximation complicates it.
Worked example: one-step planning with an exact terminal value
Make the argument concrete with a problem we can solve by hand. Take a one-dimensional linear system with quadratic cost,
with discount factor . For this linear-quadratic problem the optimal value function is exactly quadratic, , where solves the discounted Riccati equation
The Riccati equation is the fixed point of the Bellman equation for this problem: substituting a quadratic value function into the Bellman optimality condition and matching coefficients yields exactly this algebraic equation for . Substituting the numbers gives . Multiplying both sides by and collecting terms yields
We take the positive root because must be positive for the value function to be a cost (negative reward) that increases in magnitude with the state. The optimal controller is linear, , and the optimal return from is .
Now run a one-step planner from with . Here we replace TD-MPC's learned-policy, twin-critic terminal score with the oracle ; this is the exact optimal continuation, not a claim that any terminal policy action would work:
Expanding the square and collecting terms gives . Setting the derivative to zero yields , and substituting back gives . A one-step plan with the exact terminal value reproduces the infinite-horizon optimum exactly, both the action and the value.
Remove the terminal value and the one-step planner maximizes at : do nothing now. If it follows that single myopic action and then switches to the optimal controller, its total return is , a cost magnitude about 31.6% above the optimum . If it repeatedly chooses the myopic action, the state stays at and the return is . These are different continuation assumptions; the 31.6% number belongs only to the first.
The same argument becomes a picture if we plot the planning objective over actions. The terminal value moves the maximum away from the tempting "do nothing" action and onto the optimum. All quantities are computed from the closed-form worked example, so the curves are analytic and deterministic.
## Analytic curves for the one-dimensional LQR worked example.
## V(s) = -P s**2 with P the positive Riccati root; s_0 = 1, A = 1, B = 0.5, gamma = 0.9.
lqr_gamma = 0.9
P_pos = (0.125 + np.sqrt(0.125**2 + 0.9)) / 0.45
s0 = 1.0
a_grid = np.linspace(-1.5, 1.5, 401)
def _j_with_terminal(a):
s_next = 1.0 * s0 + 0.5 * a
return -(s0**2 + a**2) + lqr_gamma * (-P_pos * s_next**2)
def _j_myopic(a):
return -(s0**2 + a**2)
j_terminal = _j_with_terminal(a_grid)
j_myopic = _j_myopic(a_grid)
a_star_terminal = float(
-lqr_gamma * 0.5 * P_pos / (1.0 + lqr_gamma * 0.5**2 * P_pos)
)
a_star_myopic = float(a_grid[int(np.argmax(j_myopic))])
Three lessons fall out of this exercise, and all three survive into the neural setting:
- With exact dynamics and the exact optimal terminal value (equivalently, exact at an optimal terminal action), one-step maximization gives the optimal action. A suboptimal terminal action, approximate dynamics, or approximate values can break that equivalence.
- Myopic planning is not merely suboptimal in magnitude; it systematically prefers actions that look good now because it cannot see the consequence. The wrongness is structured, not random.
- The terminal value can change the maximizing action. For an approximate critic, its score error enters with a factor and can alter action ranking; the size of that effect depends on the competing score gaps.
The exact example explains why terminal estimates can matter, not how accurately a learned critic will perform. A short horizon shifts responsibility to the critic; whether that helps depends on the quality of both components near the candidate trajectories.
Choosing the horizon in practice
The horizon is a hyperparameter with a real tradeoff. Use open-loop prediction checks to identify model weaknesses, then test candidate horizons in closed-loop control. A terminal-value ablation answers a different question.
Open-loop error by horizon. Roll the learned model forward with recorded actions from held-out trajectories and compare predictions with observed rewards and encoded successor states. Flat average error alone does not certify that a longer planner is safe: the planner may select different actions or exploit rare optimistic errors.
Controlled horizon sweep. With trained weights, terminal-value setting, initial states, seeds, action-repeat setting, and evaluation budget fixed, compare closed-loop returns across several horizons. The best horizon may differ from the one with lowest recorded-action prediction error because the planner selects its own actions. Include inference latency when the controller has a deadline.
Terminal-value ablation. At a fixed horizon, remove the terminal value while holding weights and evaluation conditions fixed. This isolates its contribution for the trained agent; it does not select a horizon. A small gap can mean the value estimate is unnecessary or that it is not yet useful.
Open-loop prediction quality, horizon-dependent closed-loop return, and terminal-value contribution are distinct measurements. The toy below prints the first and two small closed-loop comparisons, but its three evaluation starts are far too few for a reliable horizon choice.
One more practical consideration: horizon and action repeat interact. If the environment repeats each action for several simulator steps, a horizon of covers times as much real time, and can correspond to a substantial fraction of a second. Report horizons in units of real time when comparing across papers. A horizon of three in a setup with action repeat of four is really a twelve-step horizon in simulator time, and comparing it to a paper with no action repeat would be misleading.
TD-MPC2 and Multi-Task Control
Most original TD-MPC benchmarks train separate agents per task, but the 2022 paper also reports a single policy trained jointly on ten Meta-World tasks (MT10). TD-MPC2 (Hansen et al., ICLR 2024) retains decoder-free latent planning with a terminal value while changing several architectural and optimization details. Its single-task online evaluations and its large offline multitask training are separate experiments; they should not be reported as one result.
Scaling
The paper reports online single-task evaluations over 104 continuous-control tasks across DMControl, Meta-World, MyoSuite, and ManiSkill2 using a common hyperparameter recipe. Separately, it trains offline multitask agents on 545 million transitions from 80 tasks (50 Meta-World and 30 DMControl); the largest has 317 million parameters. On this 80-task dataset, reported normalized score increases across the tested sizes from 1 million to 317 million parameters. That observed scaling trend does not guarantee smooth gains beyond the measured models, tasks, or data regime.
Three implementation changes accompany the scale-up.
Discrete regression for reward and value. TD-MPC2 represents scalar targets on a fixed 101-bin support in a signed-log-transformed space, with soft two-hot target weights and cross-entropy loss. The transform compresses large magnitudes before binning; a categorical head is not intrinsically immune to outliers or support mismatch. The paper motivates this change with reward-scale variation across tasks and reports ablations, rather than proving universal stability.
Task conditioning. The multitask agent learns a bounded-norm embedding for each task. That embedding conditions all five components: encoder, dynamics, reward, value, and policy. Padding and action masks reconcile differing observation and action dimensions. Sharing parameters can transfer useful structure, but interference remains possible; the paper's results do not isolate a universal best embedding placement.
Visual observations. In its image-based experiments TD-MPC2 replaces the state MLP encoder with a shallow four-layer convolutional encoder on 64-by-64 images and applies random-shift augmentation. It does not use a frozen DINOv2 backbone in those reported experiments. Pretrained visual features are a different research lineage, covered later in Pretrained Visual Models and DINO-WM.
What multi-task actually buys, and what it costs
The multi-task setting is not just "more data." TD-MPC2 shares model parameters across tasks while conditioning the encoder, dynamics, reward, value, and policy on a learned task embedding. This lets the model use similarities across tasks without requiring a task-agnostic latent or a clean decomposition into shared dynamics and task-specific rewards. The reported results do not establish that such a factorization emerged.
That has three consequences worth carrying forward.
- The value function becomes task-conditional. Across tasks, reward and value targets can differ substantially in magnitude. This helps explain why TD-MPC2 uses signed-log categorical regression rather than treating the change as merely cosmetic.
- Representation sharing is conditional. The encoder shares parameters but also receives the task embedding, so two tasks need not map similar observations to identical latents. The paper's multitask and few-shot results show that this design can reuse structure in its tested settings; they do not prove a task-agnostic representation or rule out interference on different task mixes.
- Reward scale becomes a first-class engineering problem. Reward and value targets from different tasks can have different magnitudes; TD-MPC2 uses a signed-log transform and discrete targets to improve robustness. Each plan is scored within one task, so a larger reward scale on another task does not by itself make that other task's candidates enter the current elite set.
The decoder-free latent model, short-horizon planning, terminal value, and off-policy replay remain. The critic's distributional head and default ensemble size (five heads rather than the original two), policy distribution, normalization, and task conditioning change substantially. TD-MPC2 is a continuation of the design, not the identical 2022 algorithm with more parameters.
The clearest invariant is the division of labor: learn task-relevant latent predictions, search a few steps, and use a critic for the tail. Training still unrolls learned latent dynamics along replay sequences, and neither the single-task nor multitask experiments establish that every retained component is necessary in every domain.
Code Implementation: A Minimal TD-MPC Agent
We will now build a compact TD-MPC-style cart-pole agent from scratch so the environment and controller are visible. It uses one-step replay samples rather than either paper's recurrent multistep sequence training, and borrows SimNorm from TD-MPC2. Its eight episodes are a diagnostic of code behavior, not a reproduction of either paper or a guaranteed learning demonstration.
Three diagnostics come out of this deliberately short run:
- A learning curve, to check whether this tiny training budget improves behavior at all.
- An open-loop error measurement on held-out trajectories, without assuming error must rise with horizon.
- Two small planning comparisons: terminal value on versus off at fixed horizon, and horizon three versus six with terminal value off. The analytic LQR example above, not this short run, establishes the terminal-value mechanism.
Start with imports and the environment. The cart-pole uses an Euler integrator and a near-upright initial state, so the task is balance rather than swing-up. The dense reward is ; episodes end when the pole falls past twelve degrees or the cart leaves the track. Unlike a constant survival reward, this signal gives the reward head a state-dependent target to learn. The experiment can therefore test its predictions separately from the terminal value estimate.
import copy
import math
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
torch.manual_seed(0)
np.random.seed(0)
GAMMA = 0.97 # effective horizon 1/(1-GAMMA) ~ 33 steps
LATENT_DIM = 64 # divisible by the SimNorm group count
HIDDEN = 128
MAX_EPISODE_STEPS = 200class CartPole:
"""Near-upright-start cart-pole with a dense, state-dependent reward."""
dt = 0.02
gravity = 9.8
masscart = 1.0
masspole = 0.1
length = 0.5 # half the pole length
force_mag = 10.0
theta_limit = 12 * math.pi / 180
x_limit = 2.4
def observe(self):
return np.asarray(self.state, dtype=np.float32)
def reset(self, rng=None, state=None):
if state is None:
state = rng.uniform(-0.05, 0.05, size=4)
self.state = np.asarray(state, dtype=np.float64).copy()
self.steps = 0
return self.observe()
def step(self, action):
u = float(np.asarray(action).reshape(-1)[0])
force = self.force_mag * max(-1.0, min(1.0, u))
x, xdot, theta, thetadot = self.state
cos, sin = math.cos(theta), math.sin(theta)
total_mass = self.masscart + self.masspole
polemass_length = self.masspole * self.length
temp = (force + polemass_length * thetadot**2 * sin) / total_mass
thetaacc = (self.gravity * sin - cos * temp) / (
self.length * (4.0 / 3.0 - self.masspole * cos**2 / total_mass)
)
xacc = temp - polemass_length * thetaacc * cos / total_mass
self.state = np.asarray(
[
x + self.dt * xdot,
xdot + self.dt * xacc,
theta + self.dt * thetadot,
thetadot + self.dt * thetaacc,
]
)
self.steps += 1
x, xdot, theta, thetadot = self.state
reward = math.cos(theta) - 0.05 * x**2
terminated = bool(
abs(x) > self.x_limit or abs(theta) > self.theta_limit
)
truncated = self.steps >= MAX_EPISODE_STEPS
return self.observe(), reward, terminated, truncatedA random policy gives us a reference line and confirms the environment terminates sensibly.
_rand_rng = np.random.default_rng(0)
_env = CartPole()
random_returns, random_lengths = [], []
for _ in range(20):
_env.reset(_rand_rng)
total = 0.0
for _ in range(MAX_EPISODE_STEPS):
_, reward, terminated, truncated = _env.step(
_rand_rng.uniform(-1, 1, size=1)
)
total += reward
if terminated or truncated:
break
random_returns.append(total)
random_lengths.append(_env.steps)
random_return_baseline = float(np.mean(random_returns))
random_length_baseline = float(np.mean(random_lengths))Random policy: mean return 24.2, mean episode length 24.3 steps
Now the five components. This compact implementation combines dynamics and reward prediction in one module and borrows SimNorm from TD-MPC2. It is an executable teaching example, not a reproduction of either paper's full architecture. A combined module shares computation; no benchmark here establishes a specific speedup.
class SimNorm(nn.Module):
"""Softmax within groups of latent units (Simplex Normalization)."""
def __init__(self, dim, groups=8):
super().__init__()
self.groups = groups
self.group_size = dim // groups
def forward(self, x):
shape = x.shape
x = x.view(*shape[:-1], self.groups, self.group_size)
x = F.softmax(x, dim=-1)
return x.view(*shape)
class Encoder(nn.Module):
def __init__(self, obs_dim, latent_dim=LATENT_DIM, hidden=HIDDEN):
super().__init__()
self.net = nn.Sequential(
nn.Linear(obs_dim, hidden),
nn.Mish(),
nn.Linear(hidden, hidden),
nn.Mish(),
nn.Linear(hidden, latent_dim),
SimNorm(latent_dim),
)
def forward(self, obs):
return self.net(obs)
class LatentDynamics(nn.Module):
"""Predicts the next latent and the reward for one transition."""
def __init__(self, latent_dim=LATENT_DIM, action_dim=1, hidden=HIDDEN):
super().__init__()
self.trunk = nn.Sequential(
nn.Linear(latent_dim + action_dim, hidden),
nn.Mish(),
nn.Linear(hidden, hidden),
nn.Mish(),
)
self.next_latent = nn.Sequential(
nn.Linear(hidden, latent_dim), SimNorm(latent_dim)
)
self.reward_head = nn.Linear(hidden, 1)
def forward(self, z, a):
h = self.trunk(torch.cat([z, a], dim=-1))
return self.next_latent(h), self.reward_head(h).squeeze(-1)
class QFunction(nn.Module):
def __init__(self, latent_dim=LATENT_DIM, action_dim=1, hidden=HIDDEN):
super().__init__()
self.net = nn.Sequential(
nn.Linear(latent_dim + action_dim, hidden),
nn.Mish(),
nn.Linear(hidden, hidden),
nn.Mish(),
nn.Linear(hidden, 1),
)
def forward(self, z, a):
return self.net(torch.cat([z, a], dim=-1)).squeeze(-1)
class PolicyPrior(nn.Module):
def __init__(self, latent_dim=LATENT_DIM, action_dim=1, hidden=HIDDEN):
super().__init__()
self.net = nn.Sequential(
nn.Linear(latent_dim, hidden),
nn.Mish(),
nn.Linear(hidden, hidden),
nn.Mish(),
nn.Linear(hidden, action_dim),
nn.Tanh(),
)
def forward(self, z):
return self.net(z)
class TDMPCAgent(nn.Module):
def __init__(
self, obs_dim=4, action_dim=1, latent_dim=LATENT_DIM, hidden=HIDDEN
):
super().__init__()
self.action_dim = action_dim
self.encoder = Encoder(obs_dim, latent_dim, hidden)
self.dynamics = LatentDynamics(latent_dim, action_dim, hidden)
self.q = QFunction(latent_dim, action_dim, hidden)
self.prior = PolicyPrior(latent_dim, action_dim, hidden)
def encode(self, obs):
return self.encoder(obs)Here is a small MPPI-style planner. The policy contributes noisy proposal sequences; Gaussian samples fill out the population; elites are weighted by their raw return scores after the numerically stable softmax shift. The executed action is sampled from the refitted Gaussian and clipped. Unlike the released planner, this toy starts each decision with a zero-mean, fixed-width Gaussian and has no previous-plan warm start or scheduled proposal variance.
def plan(
agent,
z,
horizon=3,
num_samples=128,
num_elites=16,
iterations=3,
num_prior_trajs=8,
sigma=0.5,
temperature=0.5,
use_terminal_value=True,
generator=None,
action_low=-1.0,
action_high=1.0,
):
"""MPPI in latent space with an optional terminal value."""
with torch.no_grad():
action_dim = agent.action_dim
mean = torch.zeros(horizon, action_dim)
std = torch.full((horizon, action_dim), sigma)
# Distinct noisy proposal trajectories from the learned policy prior.
z_prior = z.expand(num_prior_trajs, -1)
prior_steps = []
for _ in range(horizon):
a = (
agent.prior(z_prior)
+ 0.1
* torch.randn(num_prior_trajs, action_dim, generator=generator)
).clamp(action_low, action_high)
prior_steps.append(a)
z_prior, _ = agent.dynamics(z_prior, a)
prior_seq = torch.stack(prior_steps, dim=1) # (P, H, A)
for _ in range(iterations):
eps = torch.randn(
num_samples, horizon, action_dim, generator=generator
)
sampled = (mean + std * eps).clamp(action_low, action_high)
candidates = torch.cat([prior_seq, sampled], dim=0)
zs = z.expand(candidates.shape[0], -1)
returns = torch.zeros(candidates.shape[0])
for t in range(horizon):
zs, r = agent.dynamics(zs, candidates[:, t])
returns = returns + (GAMMA**t) * r
if use_terminal_value:
returns = returns + (GAMMA**horizon) * agent.q(
zs, agent.prior(zs)
)
# Softmax subtracts the elite maximum for numerical stability.
elite_idx = torch.topk(returns, num_elites).indices
elites = candidates[elite_idx]
weights = torch.softmax(temperature * returns[elite_idx], dim=0)
mean = (weights[:, None, None] * elites).sum(dim=0)
std = (
torch.sqrt(
(weights[:, None, None] * (elites - mean) ** 2).sum(dim=0)
)
+ 1e-6
)
# Sample the executed action for exploration.
action = mean[0] + temperature * std[0] * torch.randn(
action_dim, generator=generator
)
action = action.clamp(action_low, action_high)
return action, mean, stdThe toy optimizes the model/critic losses and then the policy loss in separate steps. It uses one-step replay samples rather than the papers' recurrent multi-step training and adds a small action-magnitude penalty to the toy policy objective. The penalty is not entropy and is not part of the original TD-MPC objective.
class ReplayBuffer:
def __init__(self, capacity=50_000):
self.capacity = capacity
self.storage = []
self.next_index = 0
def add(self, obs, action, reward, next_obs, terminated):
item = (
np.asarray(obs, dtype=np.float32),
np.asarray(action, dtype=np.float32).reshape(-1),
np.float32(reward),
np.asarray(next_obs, dtype=np.float32),
np.float32(terminated),
)
if len(self.storage) < self.capacity:
self.storage.append(item)
else:
self.storage[self.next_index] = item
self.next_index = (self.next_index + 1) % self.capacity
def __len__(self):
return len(self.storage)
def sample(self, batch_size, generator):
idx = torch.randint(
len(self.storage), (batch_size,), generator=generator
).tolist()
obs, act, rew, nobs, term = zip(*(self.storage[i] for i in idx))
return (
torch.as_tensor(np.stack(obs)),
torch.as_tensor(np.stack(act)),
torch.as_tensor(np.stack(rew)),
torch.as_tensor(np.stack(nobs)),
torch.as_tensor(np.stack(term)),
)
def td_update(agent, target, model_optimizer, prior_optimizer, batch):
obs, action, reward, next_obs, terminated = batch
z = agent.encode(obs)
with torch.no_grad():
z_next = agent.encode(next_obs)
z_next_consistency_target = target.encode(next_obs)
q_target = reward + GAMMA * (1.0 - terminated) * target.q(
z_next, agent.prior(z_next)
)
pred_next_z, pred_reward = agent.dynamics(z, action)
consistency_loss = F.mse_loss(pred_next_z, z_next_consistency_target)
reward_loss = F.mse_loss(pred_reward, reward)
q_loss = F.mse_loss(agent.q(z, action), q_target)
model_loss = consistency_loss + reward_loss + q_loss
model_optimizer.zero_grad()
model_loss.backward()
# Clip only parameters updated by this optimizer. Prior gradients from the
# previous actor step may still exist and must not scale the model update.
model_parameters = [
p for group in model_optimizer.param_groups for p in group["params"]
]
nn.utils.clip_grad_norm_(model_parameters, 20.0)
model_optimizer.step()
# Policy: maximize Q through the action, not through critic parameters.
z_prior = z.detach()
q_flags = [p.requires_grad for p in agent.q.parameters()]
for p in agent.q.parameters():
p.requires_grad_(False)
prior_optimizer.zero_grad()
prior_action = agent.prior(z_prior)
prior_loss = -agent.q(z_prior, prior_action).mean()
action_magnitude_penalty = 0.1 * (prior_action**2).mean()
(prior_loss + action_magnitude_penalty).backward()
nn.utils.clip_grad_norm_(agent.prior.parameters(), 20.0)
prior_optimizer.step()
for p, was_trainable in zip(agent.q.parameters(), q_flags):
p.requires_grad_(was_trainable)
return q_loss.detach().item()The optimizer boundary is worth testing: gradients left on the policy prior after an actor step must not change the next model update. This assertion runs on two identical agents and the same batch; only one starts with large stale prior gradients. SGD makes any unintended gradient rescaling visible in the model parameters.
def assert_model_update_ignores_stale_prior_grads():
with torch.random.fork_rng():
torch.manual_seed(47)
clean = TDMPCAgent()
stale = copy.deepcopy(clean)
clean_target, stale_target = copy.deepcopy(clean), copy.deepcopy(stale)
for target_model in (clean_target, stale_target):
for parameter in target_model.parameters():
parameter.requires_grad_(False)
batch_rng = torch.Generator().manual_seed(47)
batch = (
torch.randn(8, 4, generator=batch_rng),
torch.tanh(torch.randn(8, 1, generator=batch_rng)),
torch.randn(8, generator=batch_rng),
torch.randn(8, 4, generator=batch_rng),
torch.zeros(8),
)
assert batch[2].shape == batch[4].shape == (8,)
def optimizers(agent):
model_parameters = (
list(agent.encoder.parameters())
+ list(agent.dynamics.parameters())
+ list(agent.q.parameters())
)
return (
torch.optim.SGD(model_parameters, lr=0.01),
torch.optim.SGD(agent.prior.parameters(), lr=0.01),
model_parameters,
)
clean_model_opt, clean_prior_opt, clean_model_params = optimizers(clean)
stale_model_opt, stale_prior_opt, stale_model_params = optimizers(stale)
for parameter in stale.prior.parameters():
parameter.grad = torch.full_like(parameter, 1e5)
td_update(clean, clean_target, clean_model_opt, clean_prior_opt, batch)
td_update(stale, stale_target, stale_model_opt, stale_prior_opt, batch)
for clean_parameter, stale_parameter in zip(
clean_model_params, stale_model_params
):
torch.testing.assert_close(
clean_parameter, stale_parameter, rtol=0, atol=1e-7
)
assert_model_update_ignores_stale_prior_grads()The training loop alternates collection and learning. The buffer stores true termination separately from time-limit truncation, so the toy bootstraps at time limits. Two random-action episodes seed replay. Four updates every fourth environment step give roughly one update per post-warm-up step once the buffer gate opens; this is not evidence of improvement by itself.
def collect_and_train(num_episodes=8, warmup_episodes=2, seed=1):
torch.manual_seed(seed)
rng = np.random.default_rng(seed)
gen = torch.Generator().manual_seed(seed)
agent = TDMPCAgent()
target = copy.deepcopy(agent)
for p in target.parameters():
p.requires_grad_(False)
model_params = (
list(agent.encoder.parameters())
+ list(agent.dynamics.parameters())
+ list(agent.q.parameters())
)
model_optimizer = torch.optim.Adam(model_params, lr=3e-4)
prior_optimizer = torch.optim.Adam(agent.prior.parameters(), lr=3e-4)
env, buffer = CartPole(), ReplayBuffer()
log = {
"episode_return": [],
"episode_length": [],
"env_steps": [],
"updates": [],
}
total_steps = 0
total_updates = 0
for episode in range(num_episodes):
obs = env.reset(rng)
ep_return = 0.0
for t in range(MAX_EPISODE_STEPS):
if episode < warmup_episodes:
action = rng.uniform(-1.0, 1.0, size=1).astype(np.float32)
else:
with torch.no_grad():
z = agent.encode(torch.as_tensor(obs).unsqueeze(0))
action, _, _ = plan(agent, z, generator=gen)
action = action.numpy()
next_obs, reward, terminated, truncated = env.step(action)
buffer.add(obs, action, reward, next_obs, float(terminated))
obs = next_obs
ep_return += reward
total_steps += 1
if episode >= warmup_episodes and len(buffer) > 64 and t % 4 == 0:
for _ in range(4):
td_update(
agent,
target,
model_optimizer,
prior_optimizer,
buffer.sample(128, gen),
)
total_updates += 1
with torch.no_grad():
for p, tp in zip(agent.parameters(), target.parameters()):
tp.mul_(0.99).add_(p, alpha=0.01)
if terminated or truncated:
break
log["episode_return"].append(ep_return)
log["episode_length"].append(env.steps)
log["env_steps"].append(total_steps)
log["updates"].append(total_updates)
return agent, log
agent, log = collect_and_train()Environment steps collected: 127 Model/critic and policy update pairs: 72 Episode returns: 39.8 30.9 8.9 8.9 8.9 8.9 9.9 9.9 Episode lengths: 40 31 9 9 9 9 10 10 Random-policy reference return: 24.2
The executed notebook collected 127 environment steps and performed 72 model/critic-plus-policy update pairs. Its two random-action warm-up returns were 39.8 and 30.9; the six planner-driven returns were 8.9, 8.9, 8.9, 8.9, 9.9, and 9.9, below the independently measured random-policy mean of 24.2. Updates happened, but this run did not demonstrate improved control. The sharp drop is precisely why an executable learning curve matters more than a plausible training loop.
Now a small controlled diagnostic with the same trained weights and initial evaluation states. Comparing the first two settings isolates terminal-value use at horizon three. Comparing the second and third holds terminal-value use off and changes only the horizon from three to six. Three starts are too few to infer a reliable terminal-value benefit or choose a generally useful horizon.
def evaluate_planning(
agent, planning_kwargs, eval_states, max_steps=120, seed=0
):
gen = torch.Generator().manual_seed(seed)
env = CartPole()
returns, lengths = [], []
for state in eval_states:
obs = env.reset(state=state)
total = 0.0
for t in range(max_steps):
with torch.no_grad():
z = agent.encode(torch.as_tensor(obs).unsqueeze(0))
action, _, _ = plan(agent, z, generator=gen, **planning_kwargs)
obs, reward, terminated, truncated = env.step(action.numpy())
total += reward
if terminated or truncated:
break
returns.append(total)
lengths.append(t + 1)
return {
"return": float(np.mean(returns)),
"length": float(np.mean(lengths)),
}
_eval_rng = np.random.default_rng(7)
eval_states = [_eval_rng.uniform(-0.05, 0.05, size=4) for _ in range(3)]
planning_configs = {
"3 steps + terminal value": dict(horizon=3, use_terminal_value=True),
"3 steps, no terminal value": dict(horizon=3, use_terminal_value=False),
"6 steps, no terminal value": dict(horizon=6, use_terminal_value=False),
}
eval_results = {
name: evaluate_planning(agent, cfg, eval_states)
for name, cfg in planning_configs.items()
}planning configuration mean return mean length 3 steps + terminal value 9.9158 10.0 3 steps, no terminal value 9.9156 10.0 6 steps, no terminal value 9.9337 10.0
For open-loop diagnostics, collect trajectories with the learned prior plus noise, fit a linear state probe on half the segments, and evaluate both probe and rollout errors on the other half. These trajectories are a diagnostic distribution; we have not shown that they match MPPI deployment. Latent MSE, reward MSE, and state-probe L2 error have different units and answer different questions. The probe's held-out no-rollout error is a reference for interpreting its rolled-out error, not a measure of distance to the latent manifold.
def collect_validation_segments(
agent, num_segments=8, max_seg_steps=40, noise=0.3, seed=11
):
rng = np.random.default_rng(seed)
env = CartPole()
segments = []
for _ in range(num_segments):
obs = env.reset(rng)
obs_list, act_list, rew_list = [obs.copy()], [], []
for _ in range(max_seg_steps):
with torch.no_grad():
a = (
agent.prior(agent.encode(torch.as_tensor(obs).unsqueeze(0)))
.squeeze(0)
.numpy()
)
a = np.clip(a + noise * rng.normal(size=1), -1.0, 1.0).astype(
np.float32
)
next_obs, reward, terminated, truncated = env.step(a)
obs_list.append(next_obs.copy())
act_list.append(a)
rew_list.append(np.float32(reward))
obs = next_obs
if terminated or truncated:
break
if len(act_list) >= 2:
segments.append(
(np.stack(obs_list), np.stack(act_list), np.stack(rew_list))
)
return segments
def open_loop_diagnostics(agent, segments, max_horizon=8):
# Fit a linear probe on one split; evaluate all metrics on held-out segments.
assert len(segments) >= 2
split = len(segments) // 2
fit_segments, heldout_segments = segments[:split], segments[split:]
latents_all = []
states_all = []
for obs_seq, _, _ in fit_segments:
with torch.no_grad():
latents_all.append(agent.encode(torch.as_tensor(obs_seq)).numpy())
states_all.append(obs_seq)
X = np.concatenate(latents_all, axis=0)
Y = np.concatenate(states_all, axis=0)
design = np.hstack([X, np.ones((len(X), 1))])
probe, *_ = np.linalg.lstsq(design, Y, rcond=None)
heldout_obs = np.concatenate(
[obs_seq for obs_seq, _, _ in heldout_segments]
)
with torch.no_grad():
heldout_latents = agent.encode(torch.as_tensor(heldout_obs)).numpy()
heldout_design = np.hstack(
[heldout_latents, np.ones((len(heldout_latents), 1))]
)
probe_baseline = float(
np.mean(np.linalg.norm(heldout_design @ probe - heldout_obs, axis=1))
)
latent_err = {h: [] for h in range(1, max_horizon + 1)}
reward_err = {h: [] for h in range(1, max_horizon + 1)}
state_err = {h: [] for h in range(1, max_horizon + 1)}
for obs_seq, act_seq, rew_seq in heldout_segments:
with torch.no_grad():
latents = agent.encode(torch.as_tensor(obs_seq))
T = len(act_seq)
for i in range(T):
z = latents[i : i + 1]
for h in range(1, max_horizon + 1):
j = i + h - 1
if j >= T:
break
with torch.no_grad():
z, r_hat = agent.dynamics(
z, torch.as_tensor(act_seq[j : j + 1])
)
latent_err[h].append(
float(F.mse_loss(z, latents[j + 1 : j + 2]))
)
reward_err[h].append(float((r_hat - float(rew_seq[j])) ** 2))
decoded = np.hstack([z.numpy(), np.ones((1, 1))]) @ probe
state_err[h].append(
float(np.linalg.norm(decoded - obs_seq[j + 1]))
)
return latent_err, reward_err, state_err, probe_baseline
segments = collect_validation_segments(agent)
latent_err, reward_err, state_err, probe_baseline = open_loop_diagnostics(
agent, segments
)
horizons = np.arange(1, 9)
mean_latent_err = np.array([np.mean(latent_err[int(h)]) for h in horizons])
mean_reward_err = np.array([np.mean(reward_err[int(h)]) for h in horizons])
mean_state_err = np.array([np.mean(state_err[int(h)]) for h in horizons])Held-out diagnostic transitions: 44
horizon starts latent MSE reward MSE state L2
1 44 0.00085 0.00129 2.3866
2 40 0.00079 0.00447 2.4079
3 36 0.00072 0.00417 2.4470
4 32 0.00065 0.00404 2.4912
5 28 0.00058 0.00384 2.5446
6 24 0.00052 0.00390 2.5968
7 20 0.00046 0.00382 2.6571
8 16 0.00042 0.00387 2.7131
Held-out no-rollout probe error (state L2): 0.0001With all quantities computed, we can plot. The first figure makes the failed improvement test visible instead of presenting the training loop as proof of learning.

Next, the open-loop diagnostics compare predictions against held-out recorded trajectories. In the executed run, reward MSE is 0.00129 at horizon one, 0.00447 at two, and 0.00387 at eight: it rises, then falls, rather than increasing monotonically. Latent MSE falls from 0.00085 to 0.00042, while state-probe L2 error rises from 2.3866 to 2.7131. The held-out no-rollout probe reference is about 0.0001 L2. These metrics have different units. Their eligible start-state cohorts also shrink with horizon, so these descriptive means do not isolate how error changes along a fixed cohort.


Finally, the planner comparison. Rows one and two isolate terminal-value use at horizon three; rows two and three compare horizons three and six with terminal value off. Unlike the exact LQR example, a poorly trained toy critic is not expected to supply a reliable terminal estimate.

Read the printed tables before interpreting the shapes. Three questions matter:
First, is there a meaningful gap between the horizon-three rows? Both round to about 9.9, so this run does not demonstrate a terminal-value benefit. It does not refute the exact LQR mechanism. These three returns cannot tell us whether the learned critic's ranking is uninformative or whether the planner, model errors, or small evaluation sample prevent a useful ranking from changing the observed return.
Second, what happens when the no-terminal plan grows from horizon three to six? Both rows round to about 9.9. That pair changes horizon alone for this agent, but three starts and a single undertrained model cannot support a general conclusion about planning depth.
Third, the independent random-policy baseline averaged 24.2, well above these roughly 9.9 planner returns, even though 72 update pairs executed. The baseline uses different sampled starts, so it is a warning signal rather than a paired performance estimate. More episodes, seeds, and tuning would be needed before an empirical learning or ablation claim; the exact worked example above supplies the clean positive demonstration.
The open-loop table reports errors on held-out prior-plus-noise trajectories. In this run the reward MSE is nonmonotone, the latent MSE falls, and the state-probe error rises mildly from an already large one-step value. The H1 and H8 means use 44 and 16 start states respectively, so changing cohort composition may also move them. A low average error on this distribution would not certify accuracy on actions selected by MPPI, and a state-probe error includes both probe approximation and latent rollout error. Use these measurements alongside closed-loop evaluations, not as a formula for choosing the horizon.
Limitations and Impact
One limitation is optimization against learned scores. The planner searches for sequences with high predicted reward and terminal ; it may therefore favor errors in either the dynamics/reward model or critic, especially outside replay coverage. A short horizon restricts the number of model steps, but gives the terminal estimate more weight through than a longer horizon would at fixed critic error. Policy-guided proposals and action sampling can affect exploration, not guarantee that elites remain in-distribution or that optimism disappears. Offline and Conservative Model-Based RL develops the related coverage problem.
Another limitation is architectural. Deterministic latent dynamics do not explicitly represent a distribution over next states. The value target uses a policy proposal instead of solving a continuous-action maximization, so proposal quality matters. The learned reward head is fitted to the task's scalar reward; changing the task objective generally requires adapting the learned components, not simply changing a prompt. These constraints distinguish this controller from the foundation-model lineage in Part IX.
For a risk-neutral objective with a well-defined, finite expected return, optimizing that expectation remains appropriate even when outcome tails are heavy. A deterministic model can still mis-rank actions if its single predicted path misses decision-relevant stochastic effects, and it cannot directly evaluate risk-sensitive criteria such as tail loss. A critic trained on real outcomes may capture some expected effects; explicit stochastic or ensemble models offer different tools when uncertainty itself matters.
Planning runs at each control step, so inference cost grows with candidate count, iterations, and horizon. A deployment with a tight control deadline must measure that cost. Reproducing either paper also requires its version-specific details: policy-proposal noise, target computation, update schedule, normalization layers, and, in TD-MPC2, distributional heads. The original planner does not z-score returns; our earlier illustration explicitly labels that operation as hypothetical.
The contribution is a practical division of labor: a decoder-free latent model supports short local search while a learned value estimates the tail. The original paper reports strong continuous-control results; TD-MPC2 later reports both 104 online single-task evaluations and a separately trained 80-task generalist with scaling across tested model sizes. These results support the approach on those benchmarks, while the earlier toy illustrates how easily a small implementation can fail to demonstrate the same outcome.
Summary
- TD-MPC blends three traditions. Sampling-based model predictive control searches action sequences, temporal-difference learning trains the bootstrapped tail value, and control-centric representation learning trains TD-MPC's decoder-free TOLD latent state.
- Five learned components, two objectives. A temporally weighted model/critic objective combines reward, value, and latent-consistency losses; the policy has a separate action-value objective. Model-unrolled latent states receive training gradients.
- No decoder, deliberately. Reconstruction fidelity is not control sufficiency. Removing its loss can free optimization pressure for task-related predictions, but does not guarantee better features and may remove a helpful auxiliary signal; usefulness depends on the task and data.
- Planning is short but not free. The original MPPI-style controller scores sampled sequences, weights elites with a max-shifted softmax, and refits a Gaussian. It does not z-score returns. Planning cost scales with samples, iterations, and horizon.
- The terminal value supplies the tail. The released planner appends for a sampled terminal policy action . The analytic LQR example instead uses the oracle ; a learned critic and policy can help or hurt depending on their errors at candidate terminal states.
- Training and planning both use the model. Training unrolls predicted latents along replayed actions, while the released TD target bootstraps from encoded observed successors. Planning searches fresh candidate actions in the learned model.
- TD-MPC2 changes the implementation materially. It adds SimNorm and widespread LayerNorm, distributional reward/value heads, and all-component task conditioning. Reported image runs use a shallow CNN, not DINOv2; the 104-task online suite and 80-task offline generalist are distinct.
- Failure modes remain. Search can exploit learned-score errors, and deterministic dynamics cannot explicitly represent outcome distributions. The short toy run is a diagnostic rather than evidence that these issues are solved.
Quiz
Ready to test your understanding? Take this quick quiz to reinforce what you've learned about TD-MPC and control-centric representations.
TD-MPC and Control-Centric Representations
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!