# Troubleshooting Solver Convergence and Planning Horizon Selection in Stable-WorldModel

> Struggling with Stable-WorldModel solver convergence Stop planning horizon mismatches Learn how to tune horizon and solver hyperparameters for reliable results.

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

---

**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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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]` matches `action_dim` (see lines 40-44)

If your world model implements the optional `Actionable` protocol (as `PLDM` does in [`stable_worldmodel/wm/pldm/pldm.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/pgd.py)). High variance scales cause oscillation:

```python

# 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`](https://github.com/galilai-group/stable-worldmodel/blob/main/pgd.py) to catch non-differentiable outputs.

Debug this immediately with:

```python
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`](https://github.com/galilai-group/stable-worldmodel/blob/main/pgd.py)). This method assumes discrete action spaces. If using continuous actions, switch to a continuous solver like [`gd.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/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:

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

```python
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`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/cem.py):

```python
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_cost` returns differentiable tensors with `requires_grad=True` to prevent NaN errors in [`pgd.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/pgd.py)
- **Pad horizons correctly**: Use `prepare_init_action` in [`utils.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/utils.py) to handle warm-starting and validate action dimensions match between environment and model
- **Start short, increase gradually**: Begin with `horizon=5–10` for high-frequency models, scaling to `15–30` for 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 `cost` history**: The `outputs['cost']` array returned by `solve()` 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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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.