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

> Learn to implement Lagrangian and PGD solver callbacks for constrained MPC in Stable-WorldModel. Extend solver constructors, invoke callbacks at each step, and aggregate histories for advanced control.

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

---

**To implement Lagrangian and PGD solver callbacks for constrained MPC in Stable‑WorldModel, extend the solver constructors in [`lagrangian.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/lagrangian.py) and [`pgd.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/solver.py):

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

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

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

**Step 1**: Accept callbacks in `__init__`:

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

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

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

```

### Example Usage

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

**Step 1**: Initialize the callback list:

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

```python

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

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

```

### Example Usage

```python
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`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/lagrangian.py) | Augmented‑Lagrangian solver implementation; add callbacks after lines 59‑78 and 84‑88. |
| [`stable_worldmodel/solver/pgd.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/pgd.py) | Projected Gradient Descent solver; integrate callbacks within the optimization loop at lines 99‑124. |
| [`stable_worldmodel/solver/solver.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/solver.py) | Defines the `Costable` protocol required by all solvers. |
| [`stable_worldmodel/solver/callbacks/common.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/callbacks/common.py) | Base `Callback` class and generic recording utilities. |
| [`stable_worldmodel/solver/callbacks/gd.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/callbacks/gd.py) | Gradient‑specific callbacks (`GradNormRecorder`, `ActionNormRecorder`). |
| [`stable_worldmodel/solver/callbacks/cem.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/callbacks/cem.py) | CEM‑specific diagnostics (`EliteCostRecorder`, `VarNormRecorder`). |
| [`tests/solver/test_callbacks.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`.