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

Implement custom cost functions by subclassing the Costable protocol in 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. This protocol mandates two methods:

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 (PreJEPA) and 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:

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:


# 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:


# 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:

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 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 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, 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 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.

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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →