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 intoinfo['observation']and the reward intoinfo['reward'].AddPixelsWrapper: Renders the environment and stores pixel data underinfo['pixels'](orinfo['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
- Extend
Transformable: Implement the protocol fromstable_worldmodel/protocols.pyto create reversible, picklable scalers. - Register and Apply: Add your scaler to
_SCALERSinstable_worldmodel/data/normalization.pyand instantiate it viacolumn_normalizerinstable_worldmodel/data/utils.py. - Use Wrappers for Encoding: Employ
MegaWrapperand related classes fromstable_worldmodel/wrapper/default.pyto standardize observations and rewards in theinfodictionary. - 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. 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →