Implementing Categorical CEM for Discrete Action Spaces in Stable-WorldModel

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 and conforms to the generic Solver protocol defined in 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. 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 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, 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:

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:

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:

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

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 →