Troubleshooting Solver Convergence and Planning Horizon Selection in Stable-WorldModel
Solver convergence failures in Stable-WorldModel typically stem from mismatches between the planning horizon and the world model's predictive accuracy, requiring careful tuning of the horizon parameter and solver-specific hyperparameters like n_steps or num_samples.
When deploying latent dynamics models for model-based reinforcement learning, developers frequently encounter optimization stalls or exploding costs during the planning phase. This guide examines the galilai-group/stable-worldmodel codebase to provide concrete debugging strategies for troubleshooting solver convergence and evidence-based methods for selecting an appropriate planning horizon.
Understanding the Solver-World-Model Contract
Stable-WorldModel decouples planning from dynamics through two core protocols defined in stable_worldmodel/solver/solver.py. Understanding this contract is essential before diagnosing convergence issues.
The Costable Protocol
Any world model used for planning must implement the Costable protocol (lines 7-36 of solver.py). This requires exposing a criterion method (or get_cost) that returns a differentiable cost tensor for a batch of candidate actions. If your cost function returns a tensor without requires_grad=True, gradient-based solvers like PGD will fail immediately.
The Solver Protocol
The Solver protocol (lines 39-70 of solver.py) defines the common interface that all planners implement: configure, action_dim, horizon, and solve. When you instantiate a solver such as PGDSolver or CEMSolver, you must first call solver.configure(action_space, n_envs, config) to store the action dimension, number of environments, and critically, the planning horizon (config.horizon).
Warm-Start Handling and Horizon Padding
Before the optimization loop begins, the solver ensures the initial action plan covers the full horizon. The prepare_init_action utility in stable_worldmodel/solver/utils.py (lines 8-67) handles this by:
- Padding shorter plans with zeros or generated actions
- Validating that
init_action.shape[2]matchesaction_dim(see lines 40-44)
If your world model implements the optional Actionable protocol (as PLDM does in stable_worldmodel/wm/pldm/pldm.py), missing horizon steps are generated by rolling the model forward; otherwise, zeros are appended.
Diagnosing Convergence Failures
When solver.solve(info_dict, init_action) returns unstable results, inspect these specific failure modes:
Flat or Exploding Cost Curves
If the cost history (outputs['cost']) remains flat or explodes during optimization, check the planning horizon against model accuracy. According to the source analysis, horizons longer than the model's predictive capability cause error accumulation and noisy gradients. For a model trained at 30 Hz, start with horizon=5–10 steps; for DMControl experiments, successful configurations typically use horizon ∈ [10, 30].
For PGD specifically, examine var_scale and action_noise (lines 24-25 of pgd.py). High variance scales cause oscillation:
# Problematic configuration causing oscillation
solver = PGDSolver(model, n_steps=5, var_scale=1.0, action_noise=0.5)
# Stabilized configuration
solver = PGDSolver(model, n_steps=20, var_scale=0.1, action_noise=0.01)
NaN Gradients and Zero Costs
If you encounter zero or NaN costs during PGD optimization, verify that model.get_cost returns a tensor with requires_grad=True. The PGD implementation includes an explicit assertion at lines 20-23 of pgd.py to catch non-differentiable outputs.
Debug this immediately with:
outputs = solver.solve(info, init_action=None)
cost_history = outputs["cost"]
print("Cost trajectory:", cost_history)
# Visualize to spot NaNs or plateaus
import matplotlib.pyplot as plt
plt.plot(cost_history)
plt.xlabel("Iteration")
plt.ylabel("Cost")
plt.title("PGD Convergence Check")
plt.show()
Action Constraint Violations
When the solver ignores action bounds, check the projection logic in _project_action_simplex (lines 74-99 of pgd.py). This method assumes discrete action spaces. If using continuous actions, switch to a continuous solver like gd.py or verify that action_simplex_dim is correctly computed as self._action_space.n.
Selecting the Optimal Planning Horizon
The planning horizon dramatically impacts memory usage, optimization stability, and final policy performance.
Symptoms of Horizon Mismatch
| Symptom | Root Cause | Resolution |
|---|---|---|
| Cost never decreases | Horizon exceeds model predictive capability | Reduce to 5–10 steps initially |
| GPU out-of-memory | horizon × batch_size × action_dim too large |
Reduce batch_size, use torch.float16, or enable gradient checkpointing |
| Solver stalls early | Horizon shorter than gradient propagation window | Increase n_steps or switch to sampling-based methods (CEM/MPPI) |
| Inconsistent cross-env behavior | Horizon misaligned with episode length | Match config.horizon to eval_every frequency in training scripts |
Practical Horizon Selection
Use a warm-starting strategy and empirical sweep to find the optimal value:
from stable_worldmodel.solver.pgd import PGDSolver
from stable_worldmodel.wm.pldm import PLDM
import gymnasium as gym
env = gym.make("dm_control:cartpole-balance")
model = PLDM(...) # Implements Costable
# Configuration with conservative initial horizon
config = {
"horizon": 15,
"action_block": 1,
"n_steps": 10
}
solver = PGDSolver(model, n_steps=config["n_steps"])
solver.configure(
action_space=env.action_space,
n_envs=4,
config=type("Cfg", (), config)
)
info = env.reset()
plan = solver.solve(info)
# Warm-start the next solve with previous plan for stability
next_plan = solver.solve(next_info, init_action=plan["actions"])
For adaptive selection, implement a horizon sweep that increases length until cost plateaus:
def adaptive_horizon_select(solver, model, env, info, max_horizon=30):
"""Select shortest horizon achieving target cost threshold."""
best_horizon = max_horizon
target_cost = 0.5
for H in range(5, max_horizon + 1, 5):
solver.configure(
action_space=env.action_space,
n_envs=info["obs"].shape[0],
config=type("Cfg", (), {"horizon": H, "action_block": 1})
)
out = solver.solve(info)
if out["cost"][-1] < target_cost:
print(f"Converged at horizon={H}")
return out, H
return out, max_horizon
Switching Solver Strategies
When gradient descent stalls due to long horizons or non-differentiable costs, switch to sampling-based methods implemented in stable_worldmodel/solver/cem.py:
from stable_worldmodel.solver.cem import CEMSolver
# CEM is less sensitive to horizon length than PGD
solver = CEMSolver(
model,
num_samples=256, # Increase for smoother cost estimates
num_elite=16,
n_iters=5,
action_noise=0.2
)
solver.configure(
action_space=env.action_space,
n_envs=4,
config=type("Cfg", (), {"horizon": 12, "action_block": 1})
)
out = solver.solve(info_dict)
print("Best cost after CEM:", out["cost"][-1])
Summary
- Verify the Costable contract: Ensure
model.get_costreturns differentiable tensors withrequires_grad=Trueto prevent NaN errors inpgd.py - Pad horizons correctly: Use
prepare_init_actioninutils.pyto handle warm-starting and validate action dimensions match between environment and model - Start short, increase gradually: Begin with
horizon=5–10for high-frequency models, scaling to15–30for DMControl tasks while monitoring cost trends - Match solver to horizon: Use PGD for short horizons with reliable gradients; switch to CEM or MPPI when gradients become noisy or memory constraints arise
- Monitor
costhistory: Theoutputs['cost']array returned bysolve()is the primary diagnostic tool for convergence failures
Frequently Asked Questions
Why does my PGD solver return NaN costs?
NaN values typically indicate that the world model's cost function returns a tensor without gradient tracking. According to the PGD implementation at lines 20-23 of stable_worldmodel/solver/pgd.py, the solver asserts cost.requires_grad. Ensure your Costable implementation uses differentiable PyTorch operations and explicitly returns tensors with requires_grad=True.
How do I choose between gradient-based and sampling-based solvers?
Use PGD (gradient-based) when your planning horizon is short (< 15 steps) and your world model produces smooth, reliable gradients. Switch to CEM or MPPI (sampling-based) when the horizon exceeds the model's accurate prediction window, when gradients are noisy, or when GPU memory constraints prevent large batch sizes. Sampling-based methods trade computation time for stability on long horizons.
What causes the "action dimension mismatch" error?
This error originates in prepare_init_action at lines 40-44 of utils.py when init_action.shape[2] does not match the solver's action_dim. Verify that your world model's action_dim property matches the environment's action_space.shape[0], and ensure any warm-start tensors have shape (n_envs, horizon, action_dim).
When should I use warm-starting versus zero initialization?
Always use warm-starting (passing init_action=previous_plan["actions"]) when executing sequential planning steps, as implemented in scripts/plan/eval_wm.py. This significantly improves convergence stability by providing the solver with a near-optimal starting point, requiring optimization only on the remaining horizon steps. Use zero initialization only for the first planning step or when the state distribution changes drastically between steps.
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 →