# How to Implement Custom Cost Functions for World Model Planning in Stable WorldModel

> Learn to implement custom cost functions for world model planning by subclassing Costable and overriding the criterion method in Stable WorldModel. Preserve gradients for optimized planning.

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

---

**Implement custom cost functions by subclassing the `Costable` protocol in [`stable_worldmodel/solver/solver.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/solver.py), specifically overriding the `criterion` method to return a 2-D tensor of shape `(batch, samples)` while preserving gradients for the solver's optimization loop.**

The `stable-worldmodel` repository provides a modular framework for model-based planning where the planning objective is fully decoupled from the solver implementation. All world-model classes expose a standardized `Costable` interface that lets you inject arbitrary differentiable cost functions into MPPI, PGD, or other trajectory optimizers without modifying the solver source code.

## Understanding the Costable Protocol

All planning-capable world models in the repository implement the `Costable` protocol defined in [`stable_worldmodel/solver/solver.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/solver.py). This protocol mandates two methods:

```python
class Costable(Protocol):
    def criterion(self, info_dict: dict, action_candidates: torch.Tensor) -> torch.Tensor: ...
    def get_cost(self, info_dict: dict, action_candidates: torch.Tensor) -> torch.Tensor: ...

```

- **`criterion`** receives the `info_dict` (containing current observations, goal embeddings, and model predictions) and a tensor of candidate actions. It **must return a 2-D tensor of shape `(batch, samples)`** where each element represents the scalar cost for a specific trajectory sample.
- **`get_cost`** acts as a thin wrapper used by solvers. It handles goal embedding caching, runs the model rollout, and delegates to `criterion`. In most custom implementations, you only override `criterion`; `get_cost` can remain unchanged to preserve the default rollout and caching behavior.

## How the Default Cost Works

The reference implementations in [`stable_worldmodel/wm/prejepa/prejepa.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/prejepa/prejepa.py) (PreJEPA) and [`stable_worldmodel/wm/pldm/pldm.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/pldm/pldm.py) (PLDM) provide default MSE-based costs. Their `criterion` computes the mean-squared error between predicted latent states and goal latent states across all modalities:

```python
def criterion(self, info_dict: dict, action_candidates: torch.Tensor):
    emb_keys = [k for k in self.extra_encoders.keys() if k != "action"]
    cost = 0.0
    for key in emb_keys + ["pixels"]:
        preds = info_dict[f"predicted_{key}_emb"]
        goal = info_dict[f"{key}_goal_emb"]
        cost = cost + F.mse_loss(
            preds[:, :, -1:], goal, reduction="none"
        ).mean(dim=tuple(range(2, preds.ndim)))
    return cost

```

When a solver calls `model.get_cost(info, candidates)`, the model executes three internal steps:
1. **Caches goal embeddings** in `_goal_cached_info` to avoid redundant computation.
2. **Runs a rollout** via `self.rollout` to populate `predicted_*_emb` for each modality.
3. **Invokes `criterion`** to aggregate predictions into the final scalar cost.

## Implementing a Custom Cost Function

You have two practical avenues for customization: overriding only `criterion` to modify the loss computation while keeping the standard rollout, or overriding `get_cost` entirely to implement custom rollout logic.

### Option 1: Subclass and Override `criterion`

Override `criterion` when you want to modify how predictions are compared to goals but still rely on the default latent embeddings and rollout mechanism. This is ideal for adding penalties, mixing loss types, or weighting modalities differently.

The following example adds an L1 penalty on a custom "task" embedding to the standard PreJEPA implementation:

```python

# stable_worldmodel/wm/custom_costs.py

import torch.nn.functional as F
from stable_worldmodel.wm.prejepa.prejepa import PreJEPA

class PreJEPAWithL1(PreJEPA):
    """Add an L1 penalty on a user-provided 'task' embedding."""
    def criterion(self, info_dict: dict, action_candidates: torch.Tensor):
        # 1. Default MSE cost

        mse_cost = super().criterion(info_dict, action_candidates)

        # 2. Extra L1 cost on the final timestep of the "task" embedding

        task_pred = info_dict["predicted_task_emb"]          # (B, S, T, D)

        task_goal = info_dict["task_goal_emb"]              # (B, S, T, D)

        l1_cost = F.l1_loss(
            task_pred[:, :, -1, :],
            task_goal[:, :, -1, :].detach(),
            reduction="none",
        ).sum(dim=-1)                                        # (B, S)

        # 3. Combine (weights are arbitrary - tune them for your problem)

        return mse_cost + 0.5 * l1_cost

```

This approach works because the subclass inherits the full encoder and predictor pipeline, ensuring all cached embeddings remain available while your custom logic modifies only the final cost aggregation.

### Option 2: Override `get_cost` Completely

Override `get_cost` when your cost depends on information not captured by the default latent embeddings, such as a learned value network or physics-based constraints that require raw state access.

The following example replaces the MSE goal-matching objective with a learned value function:

```python

# stable_worldmodel/wm/custom_costs.py

import torch
from stable_worldmodel.wm.prejepa.prejepa import PreJEPA

class PreJEPAValueCost(PreJEPA):
    """Cost = -value(predicted final latent state)."""
    def __init__(self, *args, value_net, **kwargs):
        super().__init__(*args, **kwargs)
        self.value_net = value_net          # e.g. nn.Sequential(...)

    def get_cost(self, info_dict: dict, action_candidates: torch.Tensor):
        # Run the standard rollout to obtain predicted_emb

        info_dict = self.rollout(info_dict, action_candidates)

        # Extract the final latent embedding → shape (B·S·T, D)

        final_state = (
            info_dict["predicted_emb"][:, :, -1, :].reshape(-1, info_dict["predicted_emb"].size(-1))
        )

        # Value network returns a scalar per state

        values = self.value_net(final_state).view(action_candidates.shape[:2])   # (B, S)

        # Solvers minimise cost → use negative value

        return -values

```

This implementation preserves the standard rollout mechanics but substitutes the cost computation with a differentiable value network, enabling RL-style planning objectives.

## Integrating with Solvers

Once implemented, plug your custom world model into any solver. The following example integrates `PreJEPAWithL1` with the MPPI solver:

```python
import gymnasium as gym
from types import SimpleNamespace
from stable_worldmodel.solver.mppi import MPPISolver
from stable_worldmodel.wm.custom_costs import PreJEPAWithL1

# 1. Build the world model

model = PreJEPAWithL1(
    encoder=encoder,
    predictor=predictor,
    extra_encoders={"task": task_encoder},
)

# 2. Initialise solver

solver = MPPISolver(
    model=model,
    batch_size=4,
    num_samples=128,
    var_scale=1.0,
    n_steps=20,
    topk=30,
    temperature=0.5,
    device="cuda",
)

# 3. Configure with environment spec

solver.configure(
    action_space=gym.spaces.Box(-1.0, 1.0, shape=(action_dim,)),
    n_envs=8,
    config=SimpleNamespace(
        horizon=15,
        action_block=1,
    ),
)

# 4. Run planning

solution = solver.solve(info_dict)
optimal_actions = solution["actions"]

```

The solver treats your custom model as a black box, requiring only that `get_cost` returns a tensor of shape `(batch, samples)` with valid gradients. Swapping `PreJEPA` for `PreJEPAWithL1` immediately changes the optimization objective without any solver modifications.

## Practical Tips for Debugging and Optimization

- **Maintain tensor shape**: Always return a 2-D tensor `(batch, samples)`. Use `.sum(dim=-1)` or `.mean(dim=...)` to collapse extra dimensions into the final scalar per trajectory.
- **Preserve gradients**: Do not call `.detach()` on the loss before returning. Solvers invoke `cost.backward()` to compute action gradients.
- **Device alignment**: Ensure all tensors (predictions, goals, action candidates) reside on the same device by referencing `action_candidates.device`.
- **Leverage goal caching**: When overriding only `criterion`, reuse the cached goal embeddings that `get_cost` builds. If you override `get_cost`, you must handle goal encoding yourself (see the implementation in [`stable_worldmodel/wm/pldm/pldm.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/pldm/pldm.py) for reference).
- **Weight modalities**: Combine costs inside `criterion` using learned or hand-tuned weights, e.g., `cost = w_pix * pix_cost + w_task * task_cost`.
- **Verify embeddings**: Insert `print(info_dict.keys())` inside `criterion` during development to confirm that expected keys (`predicted_*_emb`, `*_goal_emb`) are present.

## Summary

- The `Costable` protocol in [`stable_worldmodel/solver/solver.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/solver.py) standardizes cost computation through `criterion` and `get_cost` methods.
- **Override `criterion`** to modify loss computation while retaining default rollouts and goal caching from PreJEPA or PLDM.
- **Override `get_cost`** completely when implementing costs that bypass the standard latent-space matching, such as learned value functions.
- Custom cost functions must return a 2-D tensor of shape `(batch, samples)` with gradients enabled for compatibility with MPPI, PGD, and other solvers.
- The repository decouples cost definition from solver implementation, allowing you to swap planning objectives by changing only the world-model class.

## Frequently Asked Questions

### What is the exact signature my custom criterion method must follow?

Your `criterion` method must accept `self`, an `info_dict: dict` containing embeddings and predictions, and `action_candidates: torch.Tensor`, then return a `torch.Tensor` of shape `(batch, samples)`. According to the source code in [`stable_worldmodel/solver/solver.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/solver.py), the returned tensor must contain scalar costs for every trajectory sample to enable gradient-based optimization by the solver.

### Can I use a pretrained neural network as my cost function?

Yes. You can integrate a pretrained value network or discriminator by overriding `get_cost` in your world-model subclass. Instantiate your network in `__init__`, run the standard rollout to obtain latent states, then pass those states through your network and return the negative value (or any differentiable scalar) reshaped to `(batch, samples)`.

### Why does my solver fail with shape mismatch errors?

The most common cause is returning a tensor with incorrect dimensions from `criterion`. The solver expects exactly two dimensions: batch size and number of samples. If your loss computation produces extra dimensions (e.g., from per-timestep or per-modality losses), collapse them using `.mean()` or `.sum()` over the extra axes before returning, ensuring the final shape is `(B, S)`.

### How do I debug which embeddings are available in info_dict?

Insert a print statement or breakpoint inside your `criterion` method to inspect `info_dict.keys()`. The default implementations in [`stable_worldmodel/wm/prejepa/prejepa.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wm/prejepa/prejepa.py) populate keys following the pattern `predicted_{modality}_emb` and `{modality}_goal_emb` for each configured modality (pixels, proprioception, task, etc.). Verify these keys exist before accessing them to avoid KeyError exceptions during the rollout.