Learning Latent Dynamics
Instead of predicting raw pixels frame-by-frame, the agent compresses high-dimensional observations into a compact latent state and simulates future transitions entirely in its mental code.
Why Does This Exist?
In visual reinforcement learning, the environment provides high-dimensional observations such as camera frames. Traditional model-based RL attempts to learn transition dynamics directly in observation space by predicting future raw frames:
This observation-space formulation suffers from three fatal problems:
- Computational bottleneck: Predicting and generating high-resolution images requires heavy deconvolutional decoders at every simulation step. Simulating thousands of candidate trajectories for planning becomes computationally intractable.
- Visual distractor sensitivity: Camera frames are flooded with task-irrelevant visual information—shifting clouds, flickering background monitors, and dynamic lighting. A pixel-prediction model expends most of its capacity fitting these irrelevant visual fluctuations rather than task-critical physical interactions.
- Compounding reconstruction error: Small pixel-level blurriness accumulates exponentially over multi-step horizons, causing long-term rollouts to dissolve into unrecognizable gray static.
Learning latent dynamics resolves these bottlenecks by projecting observations into a compact, low-dimensional latent space . The agent models environment transitions purely between compressed latent states . Once trained, the agent plans and optimizes policies by simulating thousands of imaginary trajectories in latent space per second—without ever decoding back to raw pixels.
Think of It Like This
Mental chess and driving with your eyes closed
When a grandmaster evaluates candidate chess lines five moves ahead, they never reconstruct photorealistic mental images of the mahogany wood grain on the board, the glare of the ceiling lamps, or the shadows cast by their opponent's hands. Doing so would overwhelm their brain with irrelevant details.
Instead, the grandmaster manipulates a compact, abstract internal representation: piece coordinates, king safety, open files, and tactical pin vectors. They simulate moves, captures, and counter-attacks purely within this compressed mental code.
Similarly, when driving down a highway, you can mentally evaluate a lane change without rendering millions of photoreceptor pixel activations of asphalt pebbles and headlight reflections in your visual cortex. Your brain maintains a compact belief state: vehicle speed, distance to the braking car ahead, and relative blind-spot clearance. You test candidate maneuvers ("if I steer left, will the truck collide?") entirely inside that internal model.
Where the analogy stops: Human brains possess millions of years of evolutionary priors that instinctively highlight physical objects and discard distractors. In reinforcement learning, the agent begins with zero prior knowledge. Unless properly regularized by information bottlenecks and predictive auxiliary objectives, a naive latent model may compress away subtle, survival-critical visual cues (like a distant red brake light) before discovering their utility.
How It Actually Works
The Recurrent State Space Model (RSSM)
Pioneered in PlaNet (Hafner et al., 2019) and Dreamer (Hafner et al., 2020), the Recurrent State Space Model (RSSM) decomposes the latent state at time into two complementary representations:
- Deterministic state : A continuous vector that captures long-term temporal context and deterministic history via a recurrent cell (such as a GRU):
- Stochastic state : A random variable that captures stochastic transitions, partial observability, and environmental unpredictability.
RSSM models the stochastic state through two distinct distributions:
- Posterior distribution : Conditioned on both the recurrent state and the encoded real observation . The posterior incorporates real-world sensor evidence:
- Prior distribution : Conditioned only on the deterministic history without observing . This distribution represents the agent's internal "imagination":
OBSERVATION MODE (Training with real data):x_t ──> Encoder ──> e_t ──┐ ├──> Posterior q_φ(s_t | h_t, e_t) ──> s_th_t-1, s_t-1, a_t-1 ──> GRU ──> h_t ──┘ │ └───> Prior p_θ(s_t | h_t) [D_KL penalty pulls Prior -> Posterior]
IMAGINATION MODE (Pixel-free rollout):h_τ-1, s_τ-1, a_τ-1 ──> GRU ──> h_τ ──> Prior p_θ(s_τ | h_τ) ──> s_τ ──> Predict (r_τ, V_τ)Variational Information Bottleneck Training
The world model parameters (decoders, prior) and (encoder, recurrent dynamics, posterior) are trained jointly by maximizing the Variational Evidence Lower Bound (ELBO):
- Reconstruction terms: Force the combined latent state to retain sufficient information to reconstruct visual observations and predict scalar rewards .
- KL divergence penalty: Minimizes the mismatch between the observation-conditioned posterior and the imagination prior . When the KL divergence is small, the agent can substitute the prior for the posterior during imagination rollouts without suffering from distribution shift.
Imagined Latent Rollouts
Once the RSSM is trained, policy learning proceeds entirely inside the latent world model:
- The agent starts from an initial latent belief encoded from the latest true observation.
- For horizon , the actor policy selects action .
- The deterministic GRU transitions to .
- The stochastic state is sampled from the prior .
- Latent reward heads and value critics evaluate the state.
Crucially, no images are decoded during imagination. This delivers a 1,000× to 10,000× throughput acceleration over physical environment interactions and allows analytic policy gradients to backpropagate directly through the differentiable transition dynamics.
Worked numerical example
Let us trace a single 1D latent state transition under the Variational Information Bottleneck.
Given State and Model Parameters:
- Deterministic feature vector: .
- Prior distribution (imagined without observation):
- An actual environment observation arrives. The encoder and posterior network yield:
- Reward head prediction: .
- Ground truth environment reward: .
- Hyperparameters: information bottleneck weight .
Step 1: Analytical KL Divergence
For two univariate Gaussians and , the exact Kullback-Leibler divergence is:
Evaluating each component:
- Variance ratio:
- Squared mean shift:
- Log ratio of variances:
Summing the terms inside the brackets:
Step 2: Decoder Negative Log-Likelihood (Reward Loss)
Assuming a unit-variance Gaussian likelihood :
Step 3: Joint ELBO Loss Contribution
With bottleneck coefficient :
The scalar training objective to minimize via gradient descent is .
Code
import mathfrom dataclasses import dataclassfrom typing import Dict
@dataclassclass LatentGaussian: """Represents a 1D Gaussian belief state defined by mean and variance."""
mean: float variance: float
@property def std(self) -> float: return math.sqrt(self.variance)
class RecurrentStateSpaceModel: """Minimal Recurrent State Space Model (RSSM) implementing:
1. Deterministic recurrent transition h_t = f(h_{t-1}, s_{t-1}, a_{t-1}) 2. Stochastic imagination prior p_theta(s_t | h_t) 3. Stochastic observation posterior q_phi(s_t | h_t, x_t) 4. Closed-form KL divergence and Variational ELBO loss computation """
def __init__(self, beta: float = 1.0) -> None: self.beta: float = beta
def deterministic_step( self, h_prev: float, s_prev: float, action: float ) -> float: """Computes deterministic feature h_t = f(h_{t-1}, s_{t-1}, a_{t-1}).""" # Calibrated recurrent cell: 0.5 * 0.8 + 0.3 * 1.0 + 0.2 * 0.5 = 0.80 return round(0.5 * h_prev + 0.3 * s_prev + 0.2 * action, 4)
def compute_prior(self, h_t: float) -> LatentGaussian: """Imagination prior without observation: p(s_t | h_t) ~ N(mu_p, var_p).""" mu_p = round(1.25 * h_t, 4) # 1.25 * 0.8 = 1.00 var_p = 1.0 return LatentGaussian(mean=mu_p, variance=var_p)
def compute_posterior(self, h_t: float, obs_x: float) -> LatentGaussian: """Inference posterior with observation: q(s_t | h_t, x_t) ~ N(mu_q, var_q).""" mu_q = round( 0.5 * h_t + 0.55 * obs_x, 4 ) # 0.5 * 0.8 + 0.55 * 2.0 = 1.50 var_q = 0.64 return LatentGaussian(mean=mu_q, variance=var_q)
@staticmethod def kl_divergence(q: LatentGaussian, p: LatentGaussian) -> float: """Closed-form KL divergence D_KL(q || p) between univariate Gaussians:
0.5 * [ (var_q / var_p) + ((mu_q - mu_p)^2 / var_p) - 1 + ln(var_p / var_q) ] """ ratio_var = q.variance / p.variance diff_mean_sq = ((q.mean - p.mean) ** 2) / p.variance log_ratio = math.log(p.variance / q.variance) return 0.5 * (ratio_var + diff_mean_sq - 1.0 + log_ratio)
def predict_reward(self, h_t: float, s_t: float) -> float: """Reward decoder head: r_hat = g(h_t, s_t).""" return round(1.0 * h_t + 0.8 * s_t, 4) # 1.0 * 0.8 + 0.8 * 1.5 = 2.00
def compute_step_loss( self, h_prev: float, s_prev: float, action: float, obs_x: float, true_reward: float, ) -> Dict[str, float]: """Performs one RSSM training step and calculates the variational loss.""" h_t = self.deterministic_step(h_prev, s_prev, action) prior = self.compute_prior(h_t) posterior = self.compute_posterior(h_t, obs_x)
# Closed-form divergence between inference posterior and imagination prior kl = self.kl_divergence(posterior, prior)
# Reward prediction from posterior latent state pred_reward = self.predict_reward(h_t, posterior.mean) reward_loss = 0.5 * ((pred_reward - true_reward) ** 2)
# Variational objective: minimize negative ELBO total_loss = reward_loss + self.beta * kl elbo = -total_loss
return { "h_t": h_t, "prior_mean": prior.mean, "prior_var": prior.variance, "posterior_mean": posterior.mean, "posterior_var": posterior.variance, "kl_divergence": round(kl, 4), "reward_pred": pred_reward, "reward_loss": round(reward_loss, 4), "elbo": round(elbo, 4), "total_loss": round(total_loss, 4), }
# Execute forward step matching the worked example parametersmodel = RecurrentStateSpaceModel(beta=1.0)metrics = model.compute_step_loss( h_prev=0.8, s_prev=1.0, action=0.5, obs_x=2.0, true_reward=1.8,)
print(f"Deterministic feature (h_t): {metrics['h_t']:.4f}")# -> Deterministic feature (h_t): 0.8000
print( f"Prior p(s_t|h_t): N({metrics['prior_mean']:.1f}, {metrics['prior_var']:.2f})")# -> Prior p(s_t|h_t): N(1.0, 1.00)
print( f"Posterior q(s_t|h_t,x_t): N({metrics['posterior_mean']:.1f}, {metrics['posterior_var']:.2f})")# -> Posterior q(s_t|h_t,x_t): N(1.5, 0.64)
print(f"KL Divergence: {metrics['kl_divergence']:.4f}")# -> KL Divergence: 0.1681
print(f"Reward Prediction: {metrics['reward_pred']:.4f} (True: 1.8)")# -> Reward Prediction: 2.0000 (True: 1.8)
print(f"Reward Loss (MSE * 0.5): {metrics['reward_loss']:.4f}")# -> Reward Loss (MSE * 0.5): 0.0200
print(f"Variational ELBO: {metrics['elbo']:.4f}")# -> Variational ELBO: -0.1881
# Verification assertionsassert metrics["h_t"] == 0.8assert metrics["prior_mean"] == 1.0 and metrics["prior_var"] == 1.0assert metrics["posterior_mean"] == 1.5 and metrics["posterior_var"] == 0.64assert metrics["kl_divergence"] == 0.1681assert metrics["reward_loss"] == 0.0200assert metrics["elbo"] == -0.1881Watch Out For
Latent Collapse and Task-Irrelevant Feature Filtering
When training latent world models purely end-to-end using task reward prediction without visual reconstruction or contrastive auxiliary losses, the model can suffer from latent collapse and fatal feature blindness.
The latent dynamics network quickly discards any observation features that lack immediate correlation with rewards. In sparse-reward navigation tasks, distant physical hazards (such as pits or concrete barriers) yield zero reward for hundreds of timesteps. If the encoder filters them out as uninformative background noise, the agent imagines collision-free paths straight into fatal hazards.
The Fix:
- Multi-task decoders: Retain an auxiliary observation reconstruction loss (or contrastive predictive coding objective, as in DreamerV3 and CURL) to force the latent code to preserve full spatial geometry regardless of reward sparsity.
- KL Balancing: In the RSSM loss, scale gradient updates through the prior and posterior at different rates (typically a ratio) and enforce a minimum KL threshold ("free bits") to prevent the prior from collapsing stochastic entropy before the posterior forms structured latent clusters.
The Quick Version
- Learning latent dynamics projects high-dimensional observation frames into a compact latent space, simulating transitions entirely without frame-by-frame pixel rendering.
- The Recurrent State Space Model (RSSM) decomposes latent state into a deterministic GRU history and a stochastic belief state .
- The Variational Information Bottleneck aligns the internal imagination prior with the observation-conditioned posterior via analytical KL divergence.
- Policy optimization runs across imagined rollouts inside the latent model, achieving 1,000× speedups and enabling analytical backpropagation through time.