How to Implement Lagrangian and PGD Solver Callbacks for Constrained MPC in Stable‑WorldModel

To implement Lagrangian and PGD solver callbacks for constrained MPC in Stable‑WorldModel, extend the solver constructors in lagrangian.py and pgd.py to accept a callbacks list, invoke each callback at every optimization step with solver state including costs and constraints, and aggregate the callback histories into the solver's output dictionary.

Stable‑WorldModel is an open‑source framework for model‑predictive control (MPC) that operates on learned world models. While the gradient‑based GradientSolver and CEMSolver already support pluggable diagnostic callbacks, the constrained MPC solvers—LagrangianSolver and PGDSolver—currently lack this instrumentation. This guide shows you how to add callback support to these solvers to enable richer debugging and monitoring of constrained optimization loops.

Core Architecture

The Costable Protocol

All MPC solvers in Stable‑WorldModel interact with the world model through the Costable protocol defined in stable_worldmodel/solver/solver.py:

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

The get_cost method returns a (batch, samples) tensor of scalar costs. For constrained optimization compatible with LagrangianSolver, models may additionally implement get_constraints(info_dict, actions) → (B, S, C), where C is the number of inequality constraints.

LagrangianSolver Workflow

The LagrangianSolver class in stable_worldmodel/solver/lagrangian.py implements augmented‑Lagrangian optimization for inequality constraints. The implementation follows a dual‑ascent pattern:

  1. Initialization (lines 26‑52): Builds a zero‑filled action tensor and adds random perturbations via init_action.
  2. Outer loop: Repeats dual ascent for n_outer_steps. Each iteration:
    • Creates a fresh optimizer (default: Adam).
    • Performs n_steps inner gradient steps on the augmented Lagrangian loss (lines 59‑78).
    • Updates Lagrange multipliers (lines 84‑88) and penalty coefficient rho via _update_multipliers.
  3. Aggregation: Selects the lowest‑cost action per environment and returns the final action tensor plus learned multipliers.

PGDSolver Workflow

The PGDSolver class in stable_worldmodel/solver/pgd.py handles projected gradient descent for discrete action spaces:

  1. Action tensor: Builds one‑hot representations and adds Gaussian noise to all samples except the first.
  2. Optimization loop (lines 99‑124): Runs n_steps of SGD on the raw cost.
  3. Projection (lines 74‑99): Projects the action tensor back onto the probability simplex using _project_action_simplex.
  4. Extraction: Converts the best sample per environment back to discrete indices.

How Callbacks Work in the Solver Ecosystem

The Callback Base Class

All diagnostic callbacks inherit from stable_worldmodel/solver/callbacks/common.py:

class Callback:
    def __init__(self, reduction: Literal['mean','sum','none']='mean'):
        ...
    
    def __call__(self, **state):
        ...
    
    def compute(self, **state):
        raise NotImplementedError
    
    @property
    def output_key(self) -> str:
        ...

Each callback receives solver state (e.g., costs, params, step) and stores history as a list[list[Any]] (batches × steps). The reduction argument determines how per‑environment values are aggregated across the batch dimension.

Existing Diagnostic Callbacks

Stable‑WorldModel ships with several concrete callbacks:

  • Cost recorders: BestCostRecorder, MeanCostRecorder (state key: costs)
  • Gradient/Action monitors: GradNormRecorder, ActionNormRecorder (state key: params)
  • CEM‑specific: EliteCostRecorder, VarNormRecorder, MeanShiftRecorder (state keys: topk_vals, var, mean)

Solvers that support callbacks (GradientSolver, CEMSolver, ICEMSolver) follow this invocation pattern:

self.callbacks = list(callbacks) if callbacks else []

# Inside optimization loop

for cb in self.callbacks:
    cb(step=step_idx, costs=costs, params=action_tensor, **extra_state)

# After solving

if self.callbacks:
    outputs['callbacks'] = {cb.output_key: cb.history for cb in self.callbacks}

Extending LagrangianSolver with Callbacks

To add callback support to LagrangianSolver, modify three key locations in stable_worldmodel/solver/lagrangian.py:

Step 1: Accept callbacks in __init__:

def __init__(self, ..., callbacks: list[Callback] | None = None, ...):
    ...
    self.callbacks = list(callbacks) if callbacks else []

Step 2: Invoke callbacks after computing the augmented Lagrangian loss. Inside the outer loop, after the inner gradient steps:

for cb in self.callbacks:
    cb(
        step=global_step,
        costs=costs,
        constraints=constraints,
        lambdas=self._lambdas[start_idx:end_idx] if self._lambdas is not None else None,
        rho=rho,
        params=batch_init,
    )

Step 3: Export histories in the output dictionary:

if self.callbacks:
    outputs['callbacks'] = {cb.output_key: cb.history for cb in self.callbacks}
return outputs

Example Usage

from stable_worldmodel.solver.callbacks import (
    BestCostRecorder, GradNormRecorder, ActionNormRecorder,
)

solver = LagrangianSolver(
    model=my_world_model,
    n_steps=10,
    n_outer_steps=5,
    callbacks=[BestCostRecorder(), GradNormRecorder(), ActionNormRecorder()],
)

info = {"pixels": torch.randn(4, 1, 3, 64, 64)}
result = solver(info)

# Access diagnostic histories

if "callbacks" in result:
    print(result["callbacks"]["BestCostRecorder"])  # List of per-step best costs

Extending PGDSolver with Callbacks

The PGDSolver follows an identical pattern. In stable_worldmodel/solver/pgd.py:

Step 1: Initialize the callback list:

def __init__(self, ..., callbacks: list[Callback] | None = None, ...):
    ...
    self.callbacks = list(callbacks) if callbacks else []

Step 2: Insert callback invocation after each projection step inside the optimization loop (around lines 99‑124):


# After computing costs and projecting actions

for cb in self.callbacks:
    cb(
        step=step_idx,
        costs=costs,
        params=batch_init,
    )

Step 3: Return callback data:

if self.callbacks:
    outputs['callbacks'] = {cb.output_key: cb.history for cb in self.callbacks}

Example Usage

solver = PGDSolver(
    model=my_discrete_model,
    n_steps=15,
    callbacks=[BestCostRecorder(reduction="none"), ActionNormRecorder()],
)
result = solver(info)

# result['callbacks'] contains per-batch, per-step diagnostics

Benefits for Constrained MPC Diagnostics

Implementing these callbacks provides critical visibility into the constrained optimization process:

  • Constraint tracking: By logging constraints and lambdas, you can verify that the Augmented Lagrangian is driving constraint violations toward zero across outer iterations.
  • Gradient health: GradNormRecorder detects vanishing or exploding gradients, which is essential when the cost surface combines task losses with penalty terms.
  • Action feasibility: ActionNormRecorder reveals whether the optimizer respects action bounds or drifts into unrealistic regions before projection.
  • Offline analysis: Historic logs enable systematic hyper‑parameter tuning for rho_init and rho_scale without re‑running full experiments.

Key Source Files

File Purpose
stable_worldmodel/solver/lagrangian.py Augmented‑Lagrangian solver implementation; add callbacks after lines 59‑78 and 84‑88.
stable_worldmodel/solver/pgd.py Projected Gradient Descent solver; integrate callbacks within the optimization loop at lines 99‑124.
stable_worldmodel/solver/solver.py Defines the Costable protocol required by all solvers.
stable_worldmodel/solver/callbacks/common.py Base Callback class and generic recording utilities.
stable_worldmodel/solver/callbacks/gd.py Gradient‑specific callbacks (GradNormRecorder, ActionNormRecorder).
stable_worldmodel/solver/callbacks/cem.py CEM‑specific diagnostics (EliteCostRecorder, VarNormRecorder).
tests/solver/test_callbacks.py Reference test suite demonstrating callback behavior for existing solvers.

Summary

  • LagrangianSolver and PGDSolver provide powerful constrained optimization but lack the callback instrumentation available to other Stable‑WorldModel solvers.
  • Adding support requires only three changes: accepting a callbacks list in __init__, invoking each callback with solver state (costs, constraints, lambdas, params) during the optimization loop, and exporting the aggregated histories in the output dictionary.
  • You can reuse existing callbacks like BestCostRecorder and GradNormRecorder without modification, or create custom callbacks that monitor constraint‑specific metrics.
  • These diagnostics make constrained MPC pipelines transparent, debuggable, and easier to tune for real‑world robotics tasks.

Frequently Asked Questions

What is the Costable protocol in Stable‑WorldModel?

The Costable protocol is an interface defined in stable_worldmodel/solver/solver.py that requires implementing get_cost(info_dict, action_candidates). It standardizes how MPC solvers query costs from learned world models. For constrained MPC, models can optionally implement get_constraints() to return a tensor of shape (B, S, C) representing C inequality constraints.

How does the LagrangianSolver handle constraints?

LagrangianSolver uses an augmented Lagrangian method. It converts inequality constraints into penalty terms using Lagrange multipliers (lambdas) and a penalty coefficient (rho). During each outer iteration, it performs gradient descent on the combined cost‑plus‑penalty objective, then updates the multipliers via dual ascent (lines 84‑88 in lagrangian.py) to enforce constraint satisfaction.

Can I use existing callbacks like BestCostRecorder with LagrangianSolver?

Yes. Once you add callback support to LagrangianSolver, existing callbacks from stable_worldmodel/solver/callbacks work immediately because they only require standard state keys like costs and params. You can also create custom callbacks that access Lagrangian‑specific keys such as constraints, lambdas, and rho to monitor constraint violations directly.

Where should callbacks be invoked in the PGDSolver optimization loop?

Invoke callbacks after the gradient step and simplex projection (lines 74‑99 in pgd.py) within the main optimization loop. This ensures the params state reflects the projected actions actually used for computing costs. Pass step, costs, and params to match the interface expected by ActionNormRecorder and BestCostRecorder.

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 →