# Understanding World Model Rollout and Embedding Caching in Stable-WorldModel

> Learn how galilai-group/stable-worldmodel optimizes rollout and caching with vision transformers for efficient model-predictive control. Discover embedding caching techniques.

- Repository: [GalilAI-group/stable-worldmodel](https://github.com/galilai-group/stable-worldmodel)
- Tags: internals
- Published: 2026-05-30

---

**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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/lewm/lewm.py) and [`stable_worldmodel/wm/pldm/pldm.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/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:

```python
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:

```python
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:

```python
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:

```python
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`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/world/world.py) | EnvPool handling, rollout loop, `evaluate` and `collect` APIs |
| **LeWM model** | [`stable_worldmodel/wm/lewm/lewm.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/lewm/lewm.py) | Embedding caching (`emb`, `goal_emb`), autoregressive rollout implementation |
| **PLDM model** | [`stable_worldmodel/wm/pldm/pldm.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/pldm/pldm.py) | Diffusion-based predictor with identical caching strategy |
| **Utility helpers** | [`stable_worldmodel/wm/utils.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/utils.py) | Pretrained encoder loading and dataset utilities |
| **Environment interface** | [`stable_worldmodel/wrapper/mega_wrapper.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/lewm/lewm.py) **and** [`stable_worldmodel/wm/pldm/pldm.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/pldm/pldm.py), checking for key existence before invoking the encoder.
- The `World` class in [`stable_worldmodel/world/world.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/lewm/lewm.py) or [`stable_worldmodel/wm/pldm/pldm.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/pldm/pldm.py) when `get_cost` or `rollout` is invoked, flattening batch dimensions and iterating through the prediction horizon.