Implementing Online Learning with Iterative World Model Updates in Stable-WorldModel
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 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. 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. 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 transforms the world model into an actionable planner. It uses the CEMSolver from 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 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 usingbuffer.sample() - The
_ForwardContexthelper 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:
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:
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:
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.pyscript orchestrates the full loop: warm-up, environment stepping viaEnvPool, buffer storage throughwrite_episode(), and gradient updates viatdmpc2_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 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) 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, 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 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.
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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →