# How to Implement Custom Reward Shaping Wrappers for World Model Training in stable-worldmodel

> Learn to implement custom reward shaping wrappers for world model training in stable-worldmodel. Subclass gym Wrapper to modify rewards for improved agent learning.

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

---

**To implement custom reward shaping wrappers in stable-worldmodel, subclass `gym.Wrapper` to intercept the `step()` method, apply your transformation to the raw reward, and store the result in `info['reward']` so the world model trainer consumes the modified signal.**

The `galilai-group/stable-worldmodel` library standardizes all environment interactions through an `info` dictionary, where base wrappers in [`stable_worldmodel/wrapper/default.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wrapper/default.py) explicitly place the reward signal at `info['reward']`. This architecture allows you to inject custom shaping logic—such as dense rewards or curriculum bonuses—without modifying the core training loop in [`stable_worldmodel/world/world.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/world/world.py). By creating a wrapper that modifies the reward before downstream components read it, you can control the signal that drives world model learning.

## Why the Info Dict is the Injection Point

In [`stable_worldmodel/wrapper/default.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wrapper/default.py), the `EverythingToInfoWrapper` class moves the native Gymnasium reward into the `info` dictionary during the `step()` method:

```python
obs, reward, terminated, truncated, info = self.env.step(action)
...
assert 'reward' not in info
info['reward'] = reward                # ← original reward placement

```

[Source](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wrapper/default.py#L93-L104)

Any downstream component—including the world model learner—reads `info['reward']` to compute losses and gradients. Therefore, a reward-shaping wrapper must execute **before** this insertion point and overwrite `info['reward']` with the transformed value.

## Building a Custom Reward Shaping Wrapper

Create a new file (e.g., [`stable_worldmodel/wrapper/reward.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wrapper/reward.py)) that inherits from `gym.Wrapper` and overrides `step()` to apply your shaping function.

### Step 1: Define the Wrapper Class

```python

# stable_worldmodel/wrapper/reward.py

import gym
from typing import Callable, Any

class RewardShapingWrapper(gym.Wrapper):
    """
    Apply a user-provided shaping function to the environment reward.
    
    The shaping function receives the raw reward and current info dict,
    returning the shaped reward value.
    """
    def __init__(self, env: gym.Env, shaping_fn: Callable[[float, dict], float]):
        super().__init__(env)
        self._shaping_fn = shaping_fn

```

### Step 2: Intercept and Transform

Override `step()` to call the inner environment, apply the shaping function, and modify the info dict:

```python
    def step(self, action: Any):
        # Get original step tuple from wrapped env

        obs, raw_reward, terminated, truncated, info = self.env.step(action)
        
        # Compute shaped reward using custom logic

        shaped_reward = self._shaping_fn(raw_reward, info)
        
        # Store values for debugging and downstream consumption

        info["raw_reward"] = raw_reward
        info["reward"] = shaped_reward   # The field consumed by world model trainer

        
        return obs, shaped_reward, terminated, truncated, info

```

### Step 3: Implement Shaping Functions

Define reusable shaping functions that match the `Callable[[float, dict], float]` signature:

```python
def linear_scale(factor: float) -> Callable[[float, dict], float]:
    """Multiply the raw reward by a constant factor."""
    return lambda r, _: r * factor

def sparse_to_dense(threshold: float, bonus: float) -> Callable[[float, dict], float]:
    """
    Turn a sparse reward (e.g. 0/1) into a dense signal:
    - If raw reward >= threshold → keep it
    - Otherwise → give a small constant bonus
    """
    def fn(r, _):
        return r if r >= threshold else bonus
    return fn

```

## Integration Patterns

Your custom wrapper must be applied **before** `EverythingToInfoWrapper` or `MegaWrapper` so the shaped reward propagates through the pipeline correctly.

### Standalone Usage

```python
import stable_worldmodel.envs.gymnasium_robotics.fetch as fetch_env
from stable_worldmodel.wrapper.reward import RewardShapingWrapper, linear_scale

# Build raw environment

env = fetch_env.FetchPickAndPlaceEnv()

# Apply reward shaping wrapper

env = RewardShapingWrapper(env, shaping_fn=linear_scale(2.0))

# Then apply info/pixel wrappers

from stable_worldmodel.wrapper.default import MegaWrapper
env = MegaWrapper(env, add_pixels=False)

```

### Chaining with MegaWrapper

If you use `MegaWrapper` (which internally chains `AddPixelsWrapper`, `EverythingToInfoWrapper`, and others), inject your wrapper at the base:

```python
from stable_worldmodel.wrapper.default import MegaWrapper
from stable_worldmodel.wrapper.reward import RewardShapingWrapper, sparse_to_dense

base_env = MyEnv()

# Insert shaping before the mega pipeline

env = RewardShapingWrapper(base_env, shaping_fn=sparse_to_dense(0.5, 0.1))

# Apply remaining wrappers

env = MegaWrapper(env, image_shape=(84, 84), add_pixels=True)

```

Because `MegaWrapper` stores the final wrapped environment in `self.env`, the world model trainer automatically receives the shaped reward signal without requiring changes to [`stable_worldmodel/world/world.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/world/world.py).

## Key Files Reference

- **[`stable_worldmodel/wrapper/default.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wrapper/default.py)**: Contains `EverythingToInfoWrapper` and `MegaWrapper` that establish the `info['reward']` convention.
- **[`stable_worldmodel/wrapper/visual.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/wrapper/visual.py)**: Exemplifies the wrapper pattern for visual observations that you should mirror for reward shaping.
- **[`stable_worldmodel/world/world.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/world/world.py)**: Consumes `info['reward']` during world model training to compute prediction losses.

## Summary

- **Subclass `gym.Wrapper`** to create a reward-shaping wrapper that intercepts environment steps.
- **Store the shaped reward in `info['reward']`** to ensure the world model trainer uses your modified signal.
- **Keep the raw reward in `info['raw_reward']`** for debugging and comparison against shaped values.
- **Insert the wrapper before `EverythingToInfoWrapper`** or `MegaWrapper` in the wrapper chain to ensure proper signal flow.
- **Use callable shaping functions** to parameterize transformations like linear scaling or sparse-to-dense conversion.

## Frequently Asked Questions

### Where must the shaped reward be stored for the world model to use it?

The world model trainer reads the reward from `info['reward']` as implemented in the training loop. Your wrapper must write the shaped value to this specific key before returning the step tuple. You can preserve the original reward in `info['raw_reward']` for debugging purposes without affecting training.

### Can the shaping function access observations and actions from the info dict?

Yes. The shaping function receives the entire `info` dictionary as its second argument, which may already contain observations, actions, or pixel data if upstream wrappers have populated it. This allows you to implement context-aware shaping that depends on state features rather than just the scalar reward value.

### How do I verify my reward shaping is working correctly?

Check `info['raw_reward']` versus `info['reward']` in your training logs. The `stable-worldmodel` pipeline logs these values during rollouts in [`stable_worldmodel/world/world.py`](https://github.com/galilai-group/stable-worldmodel/blob/main/stable_worldmodel/world/world.py), so you can compare the shaped and unshaped signals to ensure your transformation is being applied as expected.

### Does reward shaping work with pixel-based observations?

Yes. Reward shaping operates independently of observation type. Whether you use `AddPixelsWrapper` to render visual observations or work with state vectors, the reward-shaping wrapper only modifies the scalar value in `info['reward']`. Simply ensure your wrapper is applied before `MegaWrapper` when using pixel observations.