# Implementing Categorical CEM for Discrete Action Spaces in Stable-WorldModel

> Learn how to implement categorical CEM for discrete action spaces in Stable-WorldModel. Optimize action sequences efficiently without gradients using Gumbel-max sampling and elite refitting.

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

---

**Categorical CEM in Stable-WorldModel represents discrete actions as per-timestep categorical distributions, using Gumbel-max sampling and elite refitting to optimize action sequences without gradient computation.**

Stable-WorldModel is a model-based reinforcement learning library that provides planners for both continuous and discrete control. For discrete action environments, the repository implements a dedicated **categorical Cross-Entropy Method (CEM)** solver that maintains probability distributions over action symbols rather than Gaussian distributions. This implementation resides in [`stable_worldmodel/solver/categorical_cem.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/categorical_cem.py) and conforms to the generic `Solver` protocol defined in [`stable_worldmodel/solver/solver.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/solver.py).

## How Categorical CEM Works in Stable-WorldModel

The categorical CEM solver treats discrete action selection as a distribution estimation problem, iteratively refining categorical probabilities based on trajectory costs.

### Action Representation and Probability Tensors

For an environment with `gym.spaces.Discrete(K)` actions, the solver represents the action space as a **categorical distribution** over `K` possible symbols. When the planner repeats actions for `action_block` frames (frame-skipping), the distribution is flattened such that `action_simplex_dim = K × action_block`.

The `init_probs()` method initializes a uniform probability tensor of shape `(n_envs, horizon, action_block, K)`, creating one categorical distribution for every timestep and block across all parallel environments.

### Gumbel-Max Sampling and Elite Selection

The `_sample_indices()` method draws candidate action sequences using **Gumbel-max** sampling (`log_probs + Gumbel`). This produces integer indices of shape `(B, N, H, action_block)`, where `B` is the batch size, `N` is the number of samples, and `H` is the horizon. The first sample is forced to the current mean (`argmax`) to preserve the elite sample from the previous iteration.

Sampled indices are converted to one-hot vectors via `torch.nn.functional.one_hot` and reshaped to `(B, N, H, action_block·K)` to match the input format expected by the world-model's `get_cost` method.

After evaluating costs through the **Costable** protocol (`model.get_cost(info, candidates)`), the solver selects the `topk` lowest-cost trajectories using `torch.topk` and computes empirical category frequencies from their one-hot representations to form the new distribution.

### Distribution Updates: Smoothing and EMA

The solver prevents premature distribution collapse through two stabilization mechanisms:

- **Laplace smoothing**: Adding a small constant (`smoothing > 0`) before renormalizing ensures no category probability collapses to zero.
- **Exponential moving average**: When `alpha > 0`, the solver mixes old and new probabilities as `batch_probs = α·old + (1‑α)·new`, stabilizing convergence across CEM iterations.

## Source Code Architecture

### Core Implementation File

The complete categorical CEM logic resides in **[`stable_worldmodel/solver/categorical_cem.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/categorical_cem.py)**. The `CategoricalCEMSolver` class implements the `Solver` interface with methods for configuration, initialization, and solving.

Key methods include:

- `configure(action_space, n_envs, config)`: Validates that the action space is `Discrete` and extracts planning parameters from `PlanConfig`.
- `solve(info_dict)`: Executes the iterative sampling-evaluation-refitting loop and returns the optimized action sequence.
- `_sample_indices(log_probs, num_samples)`: Internal method handling Gumbel-max sampling with deterministic seeding.

### Protocol Compliance and Integration

The solver integrates seamlessly with the library's planning infrastructure through standard protocols:

- **`Costable` interface**: Any world-model implementing `get_cost(info_dict, action_candidates)` can be plugged in, including `TD-MPC2` implementations or lightweight test stubs.
- **`PlanConfig`**: Configuration dataclass from [`stable_worldmodel/policy.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/policy.py) supplies horizon length, action block size, and receding horizon parameters.
- **Callbacks**: The solver supports diagnostic callbacks such as `BestCostRecorder` and `MeanCostRecorder` from [`stable_worldmodel/solver/callbacks/common.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/callbacks/common.py), storing per-step statistics in the output dictionary.

## Configuration and Setup

Before solving, the solver requires configuration via the `configure()` method, which accepts a `gym.spaces.Discrete` action space and a `PlanConfig` instance. The class supports batched operation across multiple environments (`n_envs`) and maintains an internal `torch.Generator` seeded via the `seed` constructor argument for reproducible sampling.

## Practical Implementation Examples

### Example 1: Basic Setup with Dummy Cost Model

The following demonstrates minimal usage with a synthetic cost function:

```python
import torch
import gymnasium.spaces as gym_spaces
from stable_worldmodel.policy import PlanConfig
from stable_worldmodel.solver.categorical_cem import CategoricalCEMSolver

class DummyCostModel:
    def get_cost(self, info_dict, action_candidates):
        # Returns sum of squared one-hot entries

        return action_candidates.pow(2).sum(dim=(-1, -2))

solver = CategoricalCEMSolver(
    model=DummyCostModel(),
    n_steps=4,
    num_samples=64,
    topk=8,
    smoothing=0.05,
    alpha=0.2,
    batch_size=2,
    seed=123,
)

action_space = gym_spaces.Discrete(5)
config = PlanConfig(horizon=6, receding_horizon=3, action_block=2)
solver.configure(action_space=action_space, n_envs=3, config=config)

info = {"state": torch.zeros(3, 4)}
result = solver.solve(info)

print("Planned actions shape:", result["actions"].shape)  # (3, 6, 2)

```

### Example 2: Monitoring Convergence with Callbacks

Track optimization progress using built-in recorders:

```python
from stable_worldmodel.solver.callbacks import BestCostRecorder, MeanCostRecorder

cbs = [BestCostRecorder(name="best"), MeanCostRecorder(name="mean")]
solver = CategoricalCEMSolver(
    model=DummyCostModel(),
    n_steps=10,
    num_samples=128,
    topk=16,
    callbacks=cbs,
    seed=42,
)

solver.configure(
    action_space=gym_spaces.Discrete(4),
    n_envs=1,
    config=PlanConfig(horizon=5, receding_horizon=2, action_block=1)
)

out = solver.solve({"state": torch.zeros(1, 3)})
best_history = out["callbacks"]["best"]  # List of 10 values

mean_history = out["callbacks"]["mean"]

```

### Example 3: Integration with TD-MPC2 World Model

Plug the solver into a trained world-model pipeline:

```python
from stable_worldmodel.wm.tdmpc2.tdmpc2 import TD_MPC2
from stable_worldmodel.solver.categorical_cem import CategoricalCEMSolver

wm = TD_MPC2(...)  # Pre-trained model

solver = CategoricalCEMSolver(
    model=wm,
    n_steps=15,
    num_samples=256,
    topk=32
)

action_space = env.action_space  # gym.spaces.Discrete

config = PlanConfig(horizon=8, receding_horizon=4, action_block=1)
solver.configure(action_space=action_space, n_envs=1, config=config)

obs = env.reset()
info = {"state": torch.from_numpy(obs["state"]).float().unsqueeze(0)}
plan = solver.solve(info)["actions"]
action = plan[0, 0].item()  # Execute first action

env.step(action)

```

## Summary

- **Categorical CEM** in Stable-WorldModel uses probability distributions over discrete symbols instead of continuous parameters, making it compatible with `gym.spaces.Discrete` environments.
- The implementation in [`stable_worldmodel/solver/categorical_cem.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/solver/categorical_cem.py) leverages **Gumbel-max sampling** for candidate generation and **elite selection** via `torch.topk` for distribution refinement.
- **Laplace smoothing** and exponential moving average (`alpha`) prevent distribution collapse and ensure stable convergence.
- The solver adheres to the **`Costable`** protocol, allowing seamless integration with any world-model implementing `get_cost()`, including TD-MPC2.
- Comprehensive testing in [`tests/solver/test_categorical_cem.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/tests/solver/test_categorical_cem.py) validates initialization, batching, convergence, and deterministic seeding.

## Frequently Asked Questions

### What is the difference between categorical CEM and standard CEM?

**Standard CEM** typically optimizes continuous action sequences using Gaussian distributions with mean and covariance parameters. **Categorical CEM** replaces Gaussian distributions with categorical distributions over discrete action symbols, utilizing Gumbel-max sampling and one-hot encodings. Both methods employ elite selection and iterative refitting, but categorical CEM operates on probability simplexes rather than continuous parameter spaces.

### How does the `action_block` parameter affect planning?

The `action_block` parameter enables **action repetition** (frame-skipping) across multiple time steps. When `action_block > 1`, the categorical distribution is flattened such that each planning timestep represents `action_block` consecutive actions. This reduces the effective planning horizon while allowing the agent to hold actions constant for extended periods, which is common in discrete control tasks with high-frequency observations.

### Can categorical CEM handle large discrete action spaces?

The categorical CEM implementation is optimized for moderate discrete spaces (up to a few dozen actions) where maintaining a full probability vector remains computationally tractable. For very large action spaces (hundreds or thousands of actions), the memory requirements for the probability tensor `(n_envs, horizon, action_block, K)` scale linearly with `K`, and alternative methods such as factorized distributions or gradient-based policies may be more appropriate.

### How is reproducibility ensured in the categorical CEM implementation?

The solver guarantees **deterministic behavior** through an internal `torch.Generator` initialized with the `seed` constructor argument. All random sampling operations, including Gumbel noise generation in `_sample_indices()`, use this dedicated generator rather than the global random state. Consequently, providing the same seed produces identical action sequences across runs, which is critical for debugging and experimental reproducibility.