How to Implement Custom Data Normalization and Observation/Reward Encoding for World Models

To implement custom data normalization and observation/reward encoding in galilai-group/stable-worldmodel, extend the Transformable protocol for your scaler, register it in stable_worldmodel/data/normalization.py, and apply it via column_normalizer; use the gym-style wrappers in stable_worldmodel/wrapper/default.py to lift observations and rewards into the info dictionary.

The stable-worldmodel repository provides a modular framework for world model training that separates data transformation from environment interaction. Understanding how to plug in custom data normalization and observation/reward encoding is essential when integrating non-standard sensor modalities or custom reward functions. This guide walks through the Transformable protocol, the wrapper-based encoding architecture, and the precise file locations where you extend the library.

Understanding the Normalization Architecture

The normalization system relies on a strict protocol-based design to ensure picklable, reversible transformations.

The Transformable protocol in stable_worldmodel/protocols.py defines the required interface: transform, inverse_transform, and a callable __call__ method. All normalizers must implement this protocol to be compatible with the data pipeline.

Built-in scalers are located in stable_worldmodel/data/normalization.py and include IdentityScaler, ZScoreScaler, and PercentileScaler. These classes demonstrate the protocol implementation and are accessible via the _SCALERS mapping defined at lines 48–52 of the file.

The column_normalizer helper in stable_worldmodel/data/utils.py bridges the gap between raw dataset columns and the world model. It instantiates a WrapTorchTransform object that implements Transformable, fits the scaler on the specified column, and returns a picklable transform ready for the training pipeline.

Implementing a Custom Normalizer

Creating a custom normalizer requires three steps: implementing the protocol, registering the class, and applying it to your dataset.

1. Create the Scaler Class

Implement the Transformable protocol in a new class. The following example implements a log-scaling followed by min-max normalization:


# my_scaler.py

import numpy as np
import torch
from stable_worldmodel.protocols import Transformable

class LogMinMaxScaler(Transformable):
    """Log-scale then min-max-normalize to [0, 1]."""

    def __init__(self, eps: float = 1e-8):
        self.eps = eps
        self.min_ = None
        self.max_ = None

    def fit(self, X):
        X = np.log(np.maximum(X, self.eps))
        self.min_ = X.min(axis=0, keepdims=True)
        self.max_ = X.max(axis=0, keepdims=True)
        return self

    def transform(self, X):
        X_alog = np.log(np.maximum(X, self.eps))
        scale = (self.max_ - self.min_).clip(min=self.eps)
        return (X_alog - self.min_) / scale

    def inverse_transform(self, X):
        X_unscaled = X * (self.max_ - self.min_) + self.min_
        return np.exp(X_unscaled)

    def __call__(self, X):
        return torch.tensor(self.transform(X), dtype=X.dtype, device=X.device)

2. Register the Scaler

Add your class to the _SCALERS dictionary in stable_worldmodel/data/normalization.py:


# At the bottom of normalization.py

from my_scaler import LogMinMaxScaler
_SCALERS['logminmax'] = LogMinMaxScaler

This registration allows get_scaler to resolve your scaler by string name.

3. Apply to Dataset Columns

Use column_normalizer to fit and wrap your scaler for specific dataset columns:

from stable_worldmodel.data.utils import column_normalizer

# `ds` is a loaded dataset via load_dataset

temp_normalizer = column_normalizer(
    ds, 
    source='temperature', 
    target='temp_norm',
    method='logminmax'
)

The helper automatically fits the scaler on the source column data and wraps it in a WrapTorchTransform that can be attached to the dataset.

Encoding Observations and Rewards

The world model expects plain observation tensors alongside a metadata dictionary (info). Normalization occurs outside the environment, while encoding is handled by gym-style wrappers that populate the info dict.

The Wrapper Pipeline

Located in stable_worldmodel/wrapper/default.py, the wrapper system includes:

  • EverythingToInfoWrapper: Lifts the observation into info['observation'] and the reward into info['reward'].
  • AddPixelsWrapper: Renders the environment and stores pixel data under info['pixels'] (or info['pixels.<N>'] for multi-view).
  • ResizeGoalWrapper: Uniformly resizes goal images to a specified shape.
  • MegaWrapper: A high-level convenience wrapper that orchestrates the above components.

Typical Configuration

import gymnasium as gym
from stable_worldmodel.wrapper.default import MegaWrapper

base_env = gym.make('walker-walk-v2')

env = MegaWrapper(
    base_env,
    image_shape=(84, 84),
    pixels_transform=None,
    goal_transform=None,
    required_keys=['^pixels(?:\\..*)?$', '^observation$'],
    separate_goal=True,
    image_resample='bilinear',
    add_pixels=True,
)

obs, info = env.reset()

# info now contains:

#   - 'observation' (original observation)

#   - 'pixels' (resized RGB image)

#   - 'reward' (NaN on reset)

During training, the world model encoder consumes this info dictionary, applies the transforms created by column_normalizer, and learns to predict the next state and reward.

End-to-End Integration Example

This complete example demonstrates loading a dataset, creating custom normalizers, and configuring the environment wrapper:

import gymnasium as gym
import stable_worldmodel.data.utils as dutils
import stable_worldmodel.wrapper.default as wwrap

# 1. Load dataset

ds = dutils.load_dataset('lerobot/pusht', cache_dir='~/.stable_worldmodel')

# 2. Build custom normalizers

pos_norm = dutils.column_normalizer(
    ds, source='position', target='pos_norm', method='zscore'
)
img_norm = dutils.column_normalizer(
    ds, source='pixels', target='pixels_norm', 
    method='percentile', low=5, high=95
)
normalizers = [pos_norm, img_norm]

# 3. Construct training environment

raw_env = gym.make('pusht-v0')
env = wwrap.MegaWrapper(
    raw_env,
    image_shape=(84, 84),
    required_keys=['^pixels$', '^observation$'],
)

# 4. Training loop

obs, info = env.reset()
while not info.get('terminated', False):
    # Apply normalizers on-the-fly

    for tr in normalizers:
        info[tr.target] = tr(info[tr.source])
    
    action = policy.act(info)  # Your world-model policy

    obs, reward, terminated, truncated, info = env.step(action)
    # info['reward'] contains the environment reward

Summary

Frequently Asked Questions

Where is the Transformable protocol defined?

The Transformable protocol is defined in stable_worldmodel/protocols.py. According to the source code, it requires implementing transform and inverse_transform methods, plus a __call__ method that returns a torch.Tensor.

How do I handle multi-view camera observations?

Use the AddPixelsWrapper available in stable_worldmodel/wrapper/default.py. When add_pixels=True is passed to MegaWrapper, it automatically stores rendered views under keys like info['pixels.0'], info['pixels.1'], etc., which you can capture using the regex ^pixels(?:\\..*)?$ in your required_keys parameter.

Can I use the normalizer without the column_normalizer utility?

Yes. You can instantiate your scaler directly and wrap it manually in WrapTorchTransform from stable_worldmodel/data/utils.py. However, column_normalizer handles fitting on the column data and ensures the transform is picklable, which is required for distributed training.

What is the performance impact of the wrapper pipeline?

The wrappers in stable_worldmodel/wrapper/default.py operate on CPU and are designed to run once per environment step, before data reaches the GPU. Since normalization is applied via WrapTorchTransform (which converts to tensor on the device), the overhead is minimal compared to the world model's forward pass. For maximum throughput, pre-normalize your dataset offline and use IdentityScaler during training.

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 →