# Implementing Online Learning with Iterative World Model Updates in Stable-WorldModel

> Learn to implement online learning with iterative world model updates in Stable-WorldModel. Enhance agent planning by alternating environment interaction and gradient-based TD-MPC2 model updates.

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

---

**The Stable-WorldModel repository enables online learning by continuously alternating between environment interaction and gradient-based updates of a TD-MPC2 latent dynamics model, allowing the agent to improve planning performance while simultaneously collecting new data.**

The **Stable-WorldModel** framework implements temporal difference model predictive control through iterative world model updates that refine a learned latent dynamics representation in real-time. By tightly coupling data collection with model optimization, this approach achieves sample-efficient online reinforcement learning without requiring offline pre-training datasets.

## Core Architecture Components

The online learning system comprises five tightly integrated modules that manage the cycle of experience collection and model refinement.

### Parallel Environment Management

The **`EnvPool`** class in [`stable_worldmodel/world/env_pool.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/world/env_pool.py) manages a fleet of synchronous gymnasium environments. It guarantees fixed episode lengths and provides batched observations to the policy, enabling high-throughput data generation necessary for stable online updates.

### Trajectory Storage

Experience is stored in the **`ReplayBuffer`** defined in [`stable_worldmodel/data/buffer.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/data/buffer.py). This buffer maintains trajectories with configurable horizon lengths and supports random minibatch sampling via the `sample()` method. Episodes are written atomically using `write_episode()`, ensuring complete trajectory integrity during the online collect-update cycle.

### Latent World Model (TD-MPC2)

The core learning substrate resides in [`stable_worldmodel/wm/tdmpc2/tdmpc2.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/tdmpc2/tdmpc2.py). The **`TDMPC2`** class jointly trains four modules in a shared latent space:

- **Encoder** – Projects raw observations into latent representations
- **Dynamics** – Predicts next latent states given actions
- **Reward predictor** – Estimates immediate rewards in latent space  
- **Q-ensemble** – Estimates value functions for planning

The `tdmpc2_forward()` function computes the multi-task loss that drives iterative updates, handling two-hot reward encoding and SimNorm latent normalization internally.

### Planning and Control

The **`WorldModelPolicy`** in [`stable_worldmodel/policy.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/policy.py) transforms the world model into an actionable planner. It uses the **`CEMSolver`** from [`stable_worldmodel/solver/cem.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/cem.py) to optimize action sequences via the Cross-Entropy Method. With `RECEDING_HORIZON = 1`, only the first action of each plan is executed before replanning occurs using updated latent states.

## The Online Learning Loop

The training orchestration in [`scripts/expert/tdmpc2_online.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/scripts/expert/tdmpc2_online.py) implements a continuous collect-update cycle that drives iterative model improvement.

### Warm-up Phase

Training begins with `SEED_STEPS` (default 5,000) of random action sampling. This phase populates the replay buffer with diverse transitions before gradient updates commence, preventing early model collapse from correlated initial experience.

### Collect-Update Cycle

Inside the `train_task` function, the system alternates between two modes every step:

**1. Data Collection**
- Parallel environments step using either random actions or the current planner (`policy.get_action()`)
- Observations, actions, and rewards accumulate in temporary per-environment buffers
- Upon episode termination, complete trajectories commit to the replay buffer via `buffer.write_episode()`

**2. Model Update**
- When buffer size exceeds `BATCH_SIZE`, the system draws a random minibatch using `buffer.sample()`
- The `_ForwardContext` helper supplies a Lightning-compatible interface for loss computation
- `tdmpc2_forward()` calculates losses across encoder, dynamics, reward, and value predictions
- Separate optimizers step for the encoder, world model, and policy components, using distinct learning rate schedules controlled by `enc_lr_scale`

### Evaluation and Checkpointing

Every `EVAL_FREQ` steps, the script freezes the current model weights and evaluates performance on fresh environments. The checkpointing logic preserves the best-performing model state, ensuring that iterative updates do not degrade planning capability over time.

## Stability Mechanisms for Iterative Updates

Several architectural choices in Stable-WorldModel prevent instability during continuous online learning.

**Two-hot reward/value encoding** replaces scalar regression with a binned classification approach, providing scale-invariant loss gradients that eliminate the need for reward normalization.

**SimNorm latent normalization** bounds latent vector magnitudes in the dynamics model, preventing long-horizon planning from diverging due to uncontrolled state space expansion.

**Automatic discount computation** derives the discount factor γ from episode length parameters, maintaining consistent temporal horizons across diverse task domains without manual tuning.

**Independent optimizer groups** isolate representation learning (`encoder`), dynamics modeling, and policy optimization into separate parameter groups. This prevents gradient interference between world model updates and control policy improvements.

## Practical Implementation Examples

### Running the Standard Online Trainer

Execute the end-to-end training script with environment specifications:

```bash
python scripts/expert/tdmpc2_online.py \
    --domain cheetah \
    --task run \
    --steps 2000000 \
    --base_dir ./models/tdmpc2 \
    --wandb

```

This command initializes the environment pool, executes the warm-up phase, and enters the iterative collect-update loop with automatic logging.

### Custom Online Training Loop

Reuse core components for bespoke implementations:

```python
from stable_worldmodel.world.env_pool import EnvPool
from stable_worldmodel.data.buffer import ReplayBuffer
from stable_worldmodel.wm.tdmpc2 import TDMPC2, tdmpc2_forward, load_cfg
from stable_worldmodel.policy import WorldModelPolicy, PlanConfig
from stable_worldmodel.solver.cem import CEMSolver
import torch
import gymnasium as gym

# Initialize environment pool

def make_env():
    return gym.make("swm/CheetahDMControl-v0")

pool = EnvPool([make_env])
obs_dim = pool.envs[0].observation_space.shape[0]
action_dim = pool.envs[0].action_space.shape[0]

# Configure and instantiate model

cfg = load_cfg(obs_dim, action_dim, discount=0.99)
model = TDMPC2(cfg).to("cpu")

# Build planner with CEM

solver = CEMSolver(
    model=model, 
    num_samples=256, 
    n_steps=4, 
    topk=64, 
    var_scale=2.0, 
    device="cpu"
)
plan_cfg = PlanConfig(horizon=cfg.wm.horizon, receding_horizon=1, warm_start=True)
policy = WorldModelPolicy(solver=solver, config=plan_cfg, process={})
policy.set_env(pool)

# Initialize replay buffer

buffer = ReplayBuffer(max_steps=1_000_000, history_len=cfg.wm.horizon + 1)

# Online learning loop

obs = pool.reset()
for step in range(100000):
    # Collect

    actions = policy.get_action({"observation": obs})
    next_obs, reward, terminated, truncated, info = pool.step(actions)
    buffer.write_episode({
        "observation": obs,
        "action": actions,
        "reward": reward
    })
    obs = next_obs
    
    # Update

    if len(buffer) >= 256:
        batch = buffer.sample(256)
        batch = {k: torch.as_tensor(v).to("cpu") for k, v in batch.items()}
        
        ctx = type("Ctx", (), {"model": model, "metrics": {}})()
        tdmpc2_forward(ctx, batch, stage="train", cfg=cfg)
        
        loss = batch["loss"]
        loss.backward()
        torch.nn.utils.clip_grad_norm_(model.parameters(), 20.0)
        # Optimizer steps would follow here

```

### Loading and Inspecting Checkpoints

Analyze trained world models for debugging or transfer learning:

```python
import torch

checkpoint = torch.load("models/tdmpc2/cheetah_run/step_500000_model.pt")
print(checkpoint.cfg)          # Training configuration

print(checkpoint.dynamics)     # Dynamics network weights

print(checkpoint.encoder)      # Observation encoder state

```

## Summary

- **Stable-WorldModel** implements online learning through continuous iteration between data collection and model-based planning using TD-MPC2.
- The **[`tdmpc2_online.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/tdmpc2_online.py)** script orchestrates the full loop: warm-up, environment stepping via **`EnvPool`**, buffer storage through **`write_episode()`**, and gradient updates via **`tdmpc2_forward()`**.
- **Two-hot encoding** and **SimNorm normalization** provide numerical stability during iterative world model updates.
- **Separate optimizers** for encoder, dynamics, and policy components prevent catastrophic forgetting during online training.
- The **receding horizon** approach (executing only the first planned action) ensures the policy adapts immediately to newly updated world models.

## Frequently Asked Questions

### How does Stable-WorldModel prevent overfitting during online updates?

The framework employs **separate optimizers** with distinct learning rate scales (`enc_lr_scale`) to isolate representation learning from policy optimization. Additionally, the **SimNorm** latent normalization in [`stable_worldmodel/wm/tdmpc2/tdmpc2.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/tdmpc2/tdmpc2.py) bounds activation magnitudes, preventing the dynamics model from overfitting to early trajectory distributions. The replay buffer maintains diverse historical data, ensuring minibatch sampling provides stable gradient estimates even as new data arrives continuously.

### What is the purpose of the warm-up phase in online training?

The **warm-up phase** (default `SEED_STEPS = 5,000` in [`scripts/expert/tdmpc2_online.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/scripts/expert/tdmpc2_online.py)) ensures the replay buffer contains sufficiently diverse transitions before gradient updates begin. Acting randomly during this phase prevents the initially untrained world model from generating correlated, low-entropy trajectories that could bias early learning. This diversity is crucial for training stable latent dynamics before the collect-update loop commences.

### Can I modify the planning horizon during online learning?

Yes. The planning horizon is configured through **`PlanConfig`** in [`stable_worldmodel/policy.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/policy.py), specifically the `horizon` parameter. However, the **receding horizon** (`receding_horizon=1`) is fixed to ensure temporal consistency—only the first action of each optimized sequence executes before replanning occurs. Modifying the total horizon length affects computational cost and plan quality but requires maintaining the receding execution pattern for stable online integration with iterative model updates.

### How does the CEM solver interact with the world model during planning?

The **`CEMSolver`** in [`stable_worldmodel/solver/cem.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/cem.py) uses the trained world model to simulate action sequences without environment interaction. During each planning step, the solver samples candidate action sequences, rolls them out through the latent dynamics model to predict rewards and values, then refines the distribution using the top-performing samples. This model-based planning allows the policy to improve immediately after each world model update, creating tight feedback between iterative model learning and control optimization.