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

> Learn to implement custom data normalization and obs/rew encoding for world models in galilai-group/stable-worldmodel. Extend the Transformable protocol and register your scaler for seamless integration.

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

---

**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`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/data/normalization.py), and apply it via `column_normalizer`; use the gym-style wrappers in [`stable_worldmodel/wrapper/default.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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:

```python

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

```python

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

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

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

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

- **Extend `Transformable`**: Implement the protocol from [`stable_worldmodel/protocols.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/protocols.py) to create reversible, picklable scalers.
- **Register and Apply**: Add your scaler to `_SCALERS` in [`stable_worldmodel/data/normalization.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/data/normalization.py) and instantiate it via `column_normalizer` in [`stable_worldmodel/data/utils.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/data/utils.py).
- **Use Wrappers for Encoding**: Employ `MegaWrapper` and related classes from [`stable_worldmodel/wrapper/default.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wrapper/default.py) to standardize observations and rewards in the `info` dictionary.
- **Separate Concerns**: Perform normalization on dataset columns outside the environment, and let wrappers handle the encoding of raw environment outputs.

## Frequently Asked Questions

### Where is the `Transformable` protocol defined?

The `Transformable` protocol is defined in [`stable_worldmodel/protocols.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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`](https://github.com/galilai-group/stable-worldmodel/blob/main/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.