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: ...
criterionreceives theinfo_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_costacts as a thin wrapper used by solvers. It handles goal embedding caching, runs the model rollout, and delegates tocriterion. In most custom implementations, you only overridecriterion;get_costcan 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:
- Caches goal embeddings in
_goal_cached_infoto avoid redundant computation. - Runs a rollout via
self.rolloutto populatepredicted_*_embfor each modality. - Invokes
criterionto 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 invokecost.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 thatget_costbuilds. If you overrideget_cost, you must handle goal encoding yourself (see the implementation instable_worldmodel/wm/pldm/pldm.pyfor reference). - Weight modalities: Combine costs inside
criterionusing learned or hand-tuned weights, e.g.,cost = w_pix * pix_cost + w_task * task_cost. - Verify embeddings: Insert
print(info_dict.keys())insidecriterionduring development to confirm that expected keys (predicted_*_emb,*_goal_emb) are present.
Summary
- The
Costableprotocol instable_worldmodel/solver/solver.pystandardizes cost computation throughcriterionandget_costmethods. - Override
criterionto modify loss computation while retaining default rollouts and goal caching from PreJEPA or PLDM. - Override
get_costcompletely 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →