Planning with Learned Models
Planning with learned models optimizes action sequences at test time by simulating trajectories through an internal world model. By executing only the first action of an optimized plan and immediately re-planning, Model Predictive Control prevents simulation drift from compounding into real-world failure.
Why Does This Exist?
In classical model-free reinforcement learning, an agent memorizes decisions through a parameterized policy network . Training this network requires millions of trials, and the resulting weights are rigidly coupled to the exact reward function and dynamics encountered during training. If the target goal shifts at test time, or if an obstacle suddenly blocks a corridor, a model-free actor is blind to the change until thousands of new gradient steps update its weights.
Planning with learned models decouples decision-making from fixed policy parameters. Instead of memorizing which action to take in advance, the agent learns a transition dynamics model and reward function . At test time, the agent actively searches for high-return action trajectories inside this internal simulation.
However, naive open-loop trajectory optimization suffers from a fatal pathology: compounding model error. If an agent generates an optimal 30-step action plan using its learned model and executes all 30 steps open-loop, single-step approximation errors compound quadratically or exponentially over the horizon . By step 10, the physical environment has drifted drastically from the imagined state, and the remaining actions produce catastrophic failures.
Model Predictive Control (MPC) and shooting algorithms like the Cross-Entropy Method (CEM) resolve this vulnerability through the Receding Horizon Principle:
- At state , optimize action trajectories across a prospective horizon .
- Execute only the very first action .
- Discard the remainder of the plan.
- Observe the true next environmental state and re-plan from scratch.
By continuously grounding the optimization in true physical feedback at every single timestep, receding-horizon planning eliminates open-loop drift and provides robust, zero-shot adaptation to arbitrary test-time objectives.
Think of It Like This
Navigating a sailboat through shifting wind gusts
Imagine piloting a high-performance sailboat across an unpredictable channel with choppy waves and shifting wind gusts.
A model-free policy approach is like memorizing a rigid rudder and sail timetable three months before the voyage: "At minute 14, turn rudder 12 degrees starboard; at minute 15, ease the mainsail 4 inches." The instant an unexpected wave knocks the hull off course, the pre-memorized timetable drives the boat directly into the reef.
An MPC and Cross-Entropy Method navigator operates with dynamic foresight:
- At every single wave crest, you look ahead across the next 10 seconds (the planning horizon ).
- You mentally simulate 100 hypothetical rudder angles and sail positions (sampling candidate action sequences through your learned model of boat dynamics).
- You identify the top 10% smoothest, fastest courses through the chop (the Cross-Entropy Method elite candidates).
- You refine your planned rudder maneuver by averaging these top candidates (refitting the action distribution ).
- Crucially, you only execute the first second of that maneuver ().
- When the boat lands on the next wave crest, you do not blindly follow the remaining 9 seconds of the old calculation. You observe your true GPS position, wind direction, and boat tilt, and re-simulate 100 new prospective courses from scratch.
Where the analogy stops: human sailors use qualitative visual heuristics and physical intuition. An algorithmic MPC controller computes explicit numerical forward integrations through neural networks, evaluating thousands of candidate floating-point tensors within strict millisecond control deadlines.
How It Actually Works
Model Predictive Control and the Cross-Entropy Method
Planning with learned models replaces policy networks with online trajectory optimization via shooting methods.
┌────────────────────────────────────────────────────────┐ │ MPC Receding Horizon Control Loop │ └───────────────────────────┬────────────────────────────┘ │ ┌────────────────────────────▼────────────────────────────┐ │ Current Environment State s_t │ └────────────────────────────┬────────────────────────────┘ │ ┌────────────────────────────▼────────────────────────────┐ │ Trajectory Optimization (CEM): │ │ Sample N action sequences ~ N(μ, Σ) │ │ Roll out in learned model p̂_θ(s'|s,a) for H steps │ │ Filter elite top-K candidates & refit distribution │ └────────────────────────────┬────────────────────────────┘ │ ┌─────────────────────────┴─────────────────────────┐ │ Optimal Action Sequence: [a_t*, a_{t+1}*, ..., a_{t+H-1}*] │ ▼┌──────────────────┐ ┌─────────────────────────────────┐│ EXECUTE ONLY a_t*│ ─────────────►│ Discard a_{t+1}* ... a_{t+H-1}* │└────────┬─────────┘ └─────────────────────────────────┘ │ ▼┌──────────────────────────┐│ Physical Environment ││ s_{t+1} ~ P_true(·|s, a) │└────────┬─────────────────┘ │ └─────────────────────────► Shift horizon & re-plan from s_{t+1}1. Shooting Methods and the Receding Horizon Principle
In shooting methods, the decision variables are the sequences of prospective actions over a finite planning horizon :
Starting from the current physical state , prospective states are predicted forward sequentially using the learned dynamics model :
The trajectory optimization objective seeks the action sequence maximizing the cumulative predicted return:
where is an optional terminal value function (used in algorithms like TD-MPC) to account for infinite-horizon returns beyond the planning cutoff .
Under Model Predictive Control (MPC):
- The optimization solves for .
- The agent executes only the first action: .
- The remaining actions are discarded.
- The environment steps to the true next state: .
- The optimization problem is re-solved at timestep starting from the ground-truth state .
2. The Cross-Entropy Method (CEM) for Trajectory Optimization
Because learned neural dynamics models create complex, non-convex, and potentially discontinuous reward surfaces, standard gradient ascent on action vectors frequently gets trapped in poor local optima or suffers from exploding/vanishing gradients through time.
The Cross-Entropy Method (CEM) is an evolutionary, derivative-free optimization algorithm that iteratively contracts a sampling distribution around high-performing action sequences.
At each planning timestep , CEM executes optimization iterations:
Step 1: Initialize Sampling Distribution
The action sequence distribution is modeled as a factorized Gaussian across the planning horizon:
where is initialized to zeros (or warm-started from the shifted solution of the previous timestep), and is set to broad variance (e.g., ).
Step 2: Candidate Sampling
At iteration , sample candidate action sequences:
Each candidate is clipped to valid physical control bounds .
Step 3: Trajectory Evaluation
Each candidate sequence is simulated through the learned dynamics model starting from current state , computing its predicted return:
Step 4: Elite Selection
Sort the candidates in descending order of predicted return:
Extract the top candidates (the elite set ), where (typically ):
Step 5: Refit Distribution Parameters
Compute the empirical sample mean and sample variance of the elite set:
Step 6: Polyak Momentum Smoothing
To prevent the sampling distribution from collapsing prematurely into a narrow sub-optimal spike, update parameters using exponential smoothing factor :
After iterations, the mean of the converged distribution at the first timestep is returned as the control action:
Worked numerical example
Let us trace a single CEM refinement step by hand on a 1D control problem () over a planning horizon .
- Population size: candidate action sequences
- Elite fraction: top candidates ()
- Polyak momentum smoothing:
- Initial prior distribution at iteration :
Step 1: Candidate Rollout and Scoring
Four candidate trajectories are evaluated through the learned dynamics model, yielding predicted cumulative returns :
- Candidate 1:
- Candidate 2:
- Candidate 3:
- Candidate 4:
Step 2: Elite Selection
Sorting candidates by return in descending order:
- Rank 1: with
- Rank 2: with
- Rank 3: with
- Rank 4: with
The top elite sequences are and .
Step 3: Elite Parameter Estimation
Compute the empirical mean of the elite subset:
Compute the empirical variance of the elite subset:
- For horizon step 0 ():
- For horizon step 1 ():
Thus:
Step 4: Polyak Smoothing Parameter Update
Apply exponential smoothing with :
Step 5: MPC Action Dispatch
From the updated distribution , the controller extracts the first action:
The motor executes in the physical environment. The planned secondary action () is stored to warm-start the distribution for timestep : .
Code
The following self-contained Python implementation provides a complete CEMTrajectoryOptimizer and ModelPredictiveController, verifies the worked numerical example, and demonstrates receding-horizon control.
from typing import Callable, List, Tupleimport numpy as np
class CEMTrajectoryOptimizer: """Cross-Entropy Method (CEM) for trajectory optimization with learned models."""
def __init__( self, horizon: int, action_dim: int, num_candidates: int = 100, num_elites: int = 10, num_iterations: int = 5, smoothing: float = 0.8, action_low: float = -1.0, action_high: float = 1.0, ) -> None: self.horizon = horizon self.action_dim = action_dim self.num_candidates = num_candidates self.num_elites = num_elites self.num_iterations = num_iterations self.smoothing = smoothing self.action_low = action_low self.action_high = action_high
def plan( self, initial_state: np.ndarray, rollout_fn: Callable[[np.ndarray, np.ndarray], float], warm_start_mean: np.ndarray = None, ) -> Tuple[np.ndarray, np.ndarray]: """Optimizes action sequence over horizon H using CEM.""" if warm_start_mean is not None: mean = warm_start_mean.copy() else: mean = np.zeros((self.horizon, self.action_dim), dtype=np.float64)
var = np.ones((self.horizon, self.action_dim), dtype=np.float64)
for _ in range(self.num_iterations): std = np.sqrt(var) # Sample N candidate sequences: shape (N, horizon, action_dim) candidates = np.random.normal(mean, std, size=(self.num_candidates, self.horizon, self.action_dim)) candidates = np.clip(candidates, self.action_low, self.action_high)
# Evaluate each candidate trajectory return returns = np.array([ rollout_fn(initial_state, candidates[i]) for i in range(self.num_candidates) ], dtype=np.float64)
# Select elite candidate indices elite_indices = np.argsort(returns)[-self.num_elites:] elites = candidates[elite_indices]
# Fit new Gaussian parameters to elite subset elite_mean = np.mean(elites, axis=0) elite_var = np.var(elites, axis=0)
# Polyak momentum parameter smoothing mean = self.smoothing * elite_mean + (1.0 - self.smoothing) * mean var = self.smoothing * elite_var + (1.0 - self.smoothing) * var
return mean, var
def step_numerical_example( self, candidates: np.ndarray, returns: np.ndarray, prior_mean: np.ndarray, prior_var: np.ndarray, num_elites: int = 2, smoothing: float = 0.8, ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: """Executes a single deterministic CEM refitting step matching the worked example.""" elite_indices = np.argsort(returns)[-num_elites:] elites = candidates[elite_indices]
# Empirical elite statistics elite_mean = np.mean(elites, axis=0) elite_var = np.mean((elites - elite_mean) ** 2, axis=0)
# Apply Polyak smoothing update updated_mean = smoothing * elite_mean + (1.0 - smoothing) * prior_mean updated_var = smoothing * elite_var + (1.0 - smoothing) * prior_var
first_action = updated_mean[0] return updated_mean, updated_var, first_action
class ModelPredictiveController: """Receding-horizon Model Predictive Control (MPC) executor."""
def __init__( self, optimizer: CEMTrajectoryOptimizer, rollout_fn: Callable[[np.ndarray, np.ndarray], float], ) -> None: self.optimizer = optimizer self.rollout_fn = rollout_fn self.warm_start: np.ndarray = None
def act(self, current_state: np.ndarray) -> np.ndarray: """Solves finite-horizon problem, returns only the first action, and shifts warm-start.""" planned_seq, _ = self.optimizer.plan( initial_state=current_state, rollout_fn=self.rollout_fn, warm_start_mean=self.warm_start, )
# Receding horizon principle: extract and execute ONLY the first action executed_action = planned_seq[0]
# Warm-start shift for next timestep: shift horizon left and zero-pad tail self.warm_start = np.zeros_like(planned_seq) self.warm_start[:-1] = planned_seq[1:]
return executed_action
# --- Verification & Worked Example Demonstration ---optimizer = CEMTrajectoryOptimizer(horizon=2, action_dim=1)
# Hand-crafted candidate sequences A = [a_0, a_1] matching worked exampletest_candidates = np.array([ [[0.5], [0.8]], # A_1 -> return 3.2 [[-0.5], [0.2]], # A_2 -> return 1.0 [[0.8], [0.6]], # A_3 -> return 4.0 [[-1.0], [-0.5]] # A_4 -> return -0.5], dtype=np.float64)test_returns = np.array([3.2, 1.0, 4.0, -0.5], dtype=np.float64)
mu_0 = np.array([[0.0], [0.0]], dtype=np.float64)var_0 = np.array([[1.0], [1.0]], dtype=np.float64)
new_mu, new_var, a0_star = optimizer.step_numerical_example( candidates=test_candidates, returns=test_returns, prior_mean=mu_0, prior_var=var_0, num_elites=2, smoothing=0.8,)
print(f"Refitted Mean mu^(1): [{new_mu[0, 0]:.2f}, {new_mu[1, 0]:.2f}]")# -> Refitted Mean mu^(1): [0.52, 0.56]
print(f"Refitted Var sigma^(1): [{new_var[0, 0]:.4f}, {new_var[1, 0]:.4f}]")# -> Refitted Var sigma^(1): [0.2180, 0.2080]
print(f"First Executed Action a_0^*: {a0_star[0]:.2f}")# -> First Executed Action a_0^*: 0.52
# Assert exact numerical resultsnp.testing.assert_allclose(new_mu.flatten(), [0.52, 0.56], atol=1e-5)np.testing.assert_allclose(new_var.flatten(), [0.2180, 0.2080], atol=1e-5)np.testing.assert_allclose(a0_star[0], 0.52, atol=1e-5)print("CEM step assertions verified successfully.")# -> CEM step assertions verified successfully.Watch Out For
The Real-Time Latency Trap in High-Frequency Control
In high-frequency robotic systems (such as quadrotors or multi-joint manipulators requiring 50–200 Hz control loops), running online shooting methods creates a severe compute bottleneck.
Evaluating candidate trajectories over horizon across CEM iterations requires forward evaluations of a neural dynamics model per decision step. Even on high-end GPUs with vectorized tensor batching, inference latency frequently exceeds 20–50 milliseconds. When controller latency exceeds the control loop period, the robot acts on stale physical states, inducing phase lag, wild torque oscillations, or mechanical failure.
The Fix:
- Warm-Starting Distributions: Instead of re-initializing to zeros, shift the previous converged distribution forward by one step. This warm-start reduces required CEM iterations from to or , cutting latency by .
- Policy Distillation (Amortization): Concurrently distill the online CEM planner into a fast parametric actor network using behavioral cloning or DAgger (as in Guided Policy Search, POPLIN, or TD-MPC). At runtime, query the distilled actor network in ms, utilizing CEM only in the background or during safety-critical ambiguities.
- Latent Space Planning: Never simulate forward dynamics in high-dimensional observation space (e.g., pixels). Project states into a compact 32-dimensional latent representation (such as Dreamer's RSSM or TD-MPC latent dynamics) where vectorized batch simulations execute with sub-millisecond overhead.
The Quick Version
- Planning with Learned Models searches for optimal action trajectories at test time by rolling forward an internal dynamics model , enabling immediate zero-shot adaptation to new goals without retraining a policy.
- The Receding Horizon Principle (MPC) executes only the very first action of an optimized plan and re-plans at the next timestep, preventing simulation errors from compounding over extended horizons.
- The Cross-Entropy Method (CEM) optimizes action sequences without gradients by sampling candidate trajectories, ranking them by predicted return, and iteratively contracting a Gaussian distribution around the top elite candidates.
- Polyak Smoothing and Warm-Starting prevent premature distribution collapse and accelerate convergence, keeping online planning feasible within real-time robotics deadlines.