Understanding World Model Rollout and Embedding Caching in Stable-WorldModel

The stable-worldmodel library caches initial and goal embeddings in a mutable info dictionary to avoid redundant encoder passes during batch evaluation of action sequences, enabling efficient model-predictive control with vision transformers.

The stable-worldmodel repository separates environment management, representation learning, and temporal prediction into distinct components. During inference, the model must evaluate thousands of candidate action sequences per environment step, making world model rollout and embedding caching essential for computational efficiency. This article examines the caching mechanisms implemented in World, LeWM, and PLDM classes, demonstrating how the codebase avoids wasteful recomputation while maintaining deterministic latent representations.

World-Level Rollout Orchestration

The World class in stable_worldmodel/world/world.py provides the high-level API that drives the rollout loop. It wraps an EnvPool of parallel environments and a MegaWrapper that lifts raw observations into the info dictionary.

After attaching a policy with world.set_policy(policy), calling world.evaluate(...) triggers a rollout loop that repeatedly executes:

  1. Calls policy.get_action(infos) to select actions based on current observations
  2. Steps the environment pool with the selected actions
  3. Handles resets, termination masks, and episode accounting
  4. Forwards the info dict to the world model for cost computation

The info dictionary produced by the wrapper is mutated in-place throughout the loop, serving as the vehicle for cached embeddings that persist across evaluation steps.

Model-Level Rollout and Embedding Caching

Both LeWM and PLDM implement an inference-only rollout method that operates on the info dictionary. These classes are located in stable_worldmodel/wm/lewm/lewm.py and stable_worldmodel/wm/pldm/pldm.py respectively, and share a common caching strategy.

Initial State Embedding Caching

Before predicting future states, the model checks for pre-computed embeddings to avoid re-encoding observations:

if 'emb' not in info:
    _init = {k: v[:, 0] for k, v in info.items() if torch.is_tensor(v)}
    _init = self.encode(_init)
    info['emb'] = _init['emb'].detach().unsqueeze(1).expand(B, S, -1, -1)

If 'emb' is already present in the info dictionary (cached from a previous call), the encoder skip logic eliminates redundant computation. The cached tensor is expanded to shape (B, S, ...) to accommodate batch size B and sample count S for parallel candidate evaluation.

Goal Embedding Caching

When evaluating candidate plans against a target state, the goal encoding is similarly cached:

if 'goal_emb' not in info_dict:
    goal = {k: v[:, 0] for k, v in info_dict.items() if torch.is_tensor(v)}
    goal['pixels'] = goal['goal']
    # Strip the 'goal_' prefix and the action entry...

    goal = self.encode(goal)
    info_dict['goal_emb'] = goal['emb']

Subsequent calls to get_cost reuse info_dict['goal_emb'] instead of re-encoding the goal image, amortizing the vision transformer cost across all action candidates.

Why Caching Matters for World Model Rollout

Embedding caching in stable-worldmodel provides three critical advantages:

  • Computational Efficiency – Encoding a vision transformer backbone dominates the inference budget. Caching reduces encoder calls from thousands (one per candidate sequence) to one per environment step.
  • Determinism – Reusing the same latent vector across all candidates guarantees that identical observations yield identical embeddings, preventing subtle numerical drift during optimization.
  • Scalability – When running model-predictive control with large batch sizes, the cost of a single encoder pass becomes negligible compared to the predictor forward passes across the planning horizon.

Practical Implementation Examples

Batch Evaluation with the World Class

To run evaluation using the high-level orchestration layer:

import stable_worldmodel as swm

# Create a world with 8 parallel environments

world = swm.World('swm/PushT-v1', num_envs=8, image_shape=(64, 64))

# Attach a policy that knows how to query the world model

world.set_policy(my_policy)

# Collect 200 expert episodes

world.collect('data.lance', episodes=200, seed=42)

# Evaluate the policy on 100 fresh episodes

results = world.evaluate(episodes=100, seed=123)
print('Average return:', results['mean_return'])

This workflow handles the info dictionary lifecycle automatically, including embedding caching across rollout steps.

Direct Cost-Based Planning with LeWM

For fine-grained control over the planning process, instantiate the world model directly:

import torch
from stable_worldmodel.wm.lewm import LeWM
from stable_worldmodel.wm.utils import load_encoder, load_predictor, load_action_encoder

# Load pretrained components (paths omitted for brevity)

encoder = load_encoder('path/to/encoder.ckpt')
predictor = load_predictor('path/to/predictor.ckpt')
action_encoder = load_action_encoder('path/to/action_encoder.ckpt')

wm = LeWM(encoder, predictor, action_encoder)

# `info` contains initial observation pixels and (optional) action history

info = {'pixels': initial_pixels}          # shape (B, 1, C, H, W)

action_candidates = torch.randn(B, S, T, 4)  # B envs, S samples, horizon T

# Compute cost for each candidate sequence; caches 'emb' automatically

costs = wm.get_cost(info, action_candidates)   # shape (B, S)

best_idx = costs.argmin(dim=1)
best_action_seq = action_candidates[torch.arange(B), best_idx]

The get_cost method internally calls the rollout logic that populates and reuses cached embeddings.

Core Implementation Files

Component File Path Purpose
World orchestration stable_worldmodel/world/world.py EnvPool handling, rollout loop, evaluate and collect APIs
LeWM model stable_worldmodel/wm/lewm/lewm.py Embedding caching (emb, goal_emb), autoregressive rollout implementation
PLDM model stable_worldmodel/wm/pldm/pldm.py Diffusion-based predictor with identical caching strategy
Utility helpers stable_worldmodel/wm/utils.py Pretrained encoder loading and dataset utilities
Environment interface stable_worldmodel/wrapper/mega_wrapper.py Lifts raw observations into the info dict consumed by world models

Summary

  • The world model rollout and embedding caching mechanism stores initial state embeddings under the 'emb' key and goal embeddings under 'goal_emb' in the mutable info dictionary.
  • Caching occurs in stable_worldmodel/wm/lewm/lewm.py and stable_worldmodel/wm/pldm/pldm.py, checking for key existence before invoking the encoder.
  • The World class in stable_worldmodel/world/world.py orchestrates the rollout loop while the model-level components handle the temporal prediction and cost computation.
  • This design pattern enables evaluation of thousands of action candidates with only a single encoder forward pass, critical for real-time model-predictive control with vision-based world models.

Frequently Asked Questions

What is the purpose of the info dictionary in stable-worldmodel?

The info dictionary serves as a mutable state container that flows between the environment wrapper, the policy, and the world model. It carries observation tensors, cached embeddings (emb, goal_emb), and metadata through the rollout loop. By mutating this dictionary in-place, the system avoids redundant encoder computations while maintaining clean separation between the environment interface and the model implementation.

How do LeWM and PLDM differ in their rollout implementations?

Both LeWM (Latent World Model) and PLDM (Predictive Latent Diffusion Model) implement identical embedding caching logic in their respective files under stable_worldmodel/wm/. The primary difference lies in the predictor architecture: LeWM uses an autoregressive predictor, while PLDM employs a diffusion-based predictor. However, both share the same caching mechanism for initial and goal embeddings to optimize inference performance.

Why is embedding caching critical for model-predictive control?

In model-predictive control (MPC), the system evaluates hundreds or thousands of candidate action sequences at each environment step. Without caching, the vision encoder would process the same initial observation and goal image for every candidate, consuming the majority of the inference budget. By caching these embeddings in the info dict, the computational cost becomes independent of the sample count, enabling real-time planning with large candidate batches.

Where is the rollout loop actually executed in the codebase?

The outer rollout loop executes in stable_worldmodel/world/world.py within the evaluate and collect methods, which iteratively call policy.get_action(infos) and step the environment pool. The inner model rollout (temporal prediction) executes in either stable_worldmodel/wm/lewm/lewm.py or stable_worldmodel/wm/pldm/pldm.py when get_cost or rollout is invoked, flattening batch dimensions and iterating through the prediction horizon.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →