Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
130 changes: 125 additions & 5 deletions test/transforms/test_normalization.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,19 +5,15 @@
from __future__ import annotations

import pickle

import sys

import pytest

import torch

from _transforms_common import TIMEOUT, TORCH_VERSION
from packaging import version
from tensordict import assert_close, TensorDict, TensorDictBase
from tensordict.utils import assert_allclose_td
from torch import multiprocessing as mp

from torchrl.data import Composite, Unbounded
from torchrl.envs import (
Compose,
Expand All @@ -32,7 +28,6 @@
from torchrl.envs.libs.gym import _has_gym, GymEnv
from torchrl.envs.transforms import VecNorm
from torchrl.envs.utils import check_env_specs, step_mdp

from torchrl.testing import ( # noqa
BREAKOUT_VERSIONED,
dtype_fixture,
Expand Down Expand Up @@ -121,6 +116,131 @@ def test_vecnorm2_decay1(self, batched):
assert env.transform._loc.ndim == 0
assert env.transform._var.ndim == 0

class OffsetEnv(SimpleEnv):
def __init__(self, offset: float = 100.0, **kwargs):
super().__init__(**kwargs)
self.offset = offset

def _reset(self, tensordict: TensorDictBase, **kwargs) -> TensorDictBase:
tensordict = super()._reset(tensordict, **kwargs)
tensordict["observation"] = tensordict["observation"] + self.offset
return tensordict

def _step(self, tensordict: TensorDictBase) -> TensorDictBase:
tensordict = super()._step(tensordict)
tensordict["observation"] = tensordict["observation"] + self.offset
return tensordict

@staticmethod
def _weighted_stats(samples, decay):
"""Bias-corrected exponentially weighted mean and std of ``samples`` (updates x samples per update)."""
num_updates, n = samples.shape
age = torch.arange(num_updates - 1, -1, -1, dtype=samples.dtype)
weights = (decay ** (n * age)).unsqueeze(-1).expand_as(samples)
mean = (weights * samples).sum() / weights.sum()
std = ((weights * (samples - mean).pow(2)).sum() / weights.sum()).sqrt()
return mean, std

@pytest.mark.parametrize("stateful", [True, False])
def test_vecnorm2_scale_with_offset(self, stateful):
# Observations are OFFSET + N(0, 1): the scale must track the std (1),
# not the mean, as soon as the running stats have warmed up.
torch.manual_seed(0)
env = self.OffsetEnv().append_transform(
VecNormV2(
in_keys=["observation"],
out_keys=["obs_norm"],
decay=0.9999,
stateful=stateful,
)
)
r = env.rollout(1000, break_when_any_done=False)
obs = r["next", "observation"]
obs_norm = r["next", "obs_norm"]
torch.testing.assert_close(
obs_norm[500:].std(unbiased=False), torch.tensor(1.0), atol=0.2, rtol=0
)
if stateful:
torch.testing.assert_close(
env.transform.scale["observation"],
obs.std(unbiased=False),
atol=0.2,
rtol=0,
)
torch.testing.assert_close(
env.transform.loc["observation"], obs.mean(), atol=0.2, rtol=0
)

@pytest.mark.parametrize("mode", ["stateful", "stateless", "reduce_batch_dims"])
def test_vecnorm2_large_offset(self, mode):
torch.manual_seed(0)
decay = 0.99
env = self.OffsetEnv(offset=1e4)
if mode == "reduce_batch_dims":
env = SerialEnv(2, [lambda env=env: env] * 2)
env = env.append_transform(
VecNormV2(
in_keys=["observation"],
out_keys=["obs_norm"],
decay=decay,
stateful=mode != "stateless",
reduce_batch_dims=mode == "reduce_batch_dims",
)
)
r = env.rollout(200, break_when_any_done=False)
if mode == "reduce_batch_dims":
r = r.transpose(0, 1)
obs = torch.cat([r["observation"][:1], r["next", "observation"]]).double()
obs = obs.reshape(obs.shape[0], -1)
obs_norm = torch.cat([r["obs_norm"][:1], r["next", "obs_norm"]]).double()
obs_norm = obs_norm.reshape(obs.shape)
expected = torch.stack(
[
(obs[k] - mean) / std
for k in range(10, obs.shape[0])
for mean, std in [self._weighted_stats(obs[: k + 1], decay)]
]
)
torch.testing.assert_close(obs_norm[10:], expected, atol=1e-2, rtol=0)
if mode != "stateless":
mean, std = self._weighted_stats(obs, decay)
torch.testing.assert_close(
env.transform.loc["observation"].double(), mean, atol=1e-2, rtol=0
)
torch.testing.assert_close(
env.transform.scale["observation"].double(), std, atol=1e-2, rtol=0
)

def test_vecnorm2_load_legacy_state_dict(self):
decay = 0.9
samples = torch.randn(50, dtype=torch.float64) + 3
env = self.SimpleEnv().append_transform(
VecNormV2(in_keys=["observation"], decay=decay)
)
ema_weights = (1 - decay) * decay ** torch.arange(49, -1, -1.0)
sd = env.transform.state_dict()
sep = VecNormV2.SEP
sd["_extra_state"] = {
f"loc{sep}observation": (ema_weights * samples).sum().float(),
f"var{sep}observation": (ema_weights * samples.pow(2)).sum().float(),
f"count{sep}observation": torch.tensor(50),
}
env.transform.load_state_dict(sd)
mean, std = self._weighted_stats(samples.unsqueeze(-1), decay)
torch.testing.assert_close(
env.transform.loc["observation"].double(), mean, rtol=1e-5, atol=1e-5
)
torch.testing.assert_close(
env.transform.scale["observation"].double(), std, rtol=1e-5, atol=1e-5
)

other = self.SimpleEnv().append_transform(
VecNormV2(in_keys=["observation"], decay=decay)
)
other.transform.load_state_dict(env.transform.state_dict())
assert_close(other.transform.loc, env.transform.loc)
assert_close(other.transform.scale, env.transform.scale)

@pytest.mark.skipif(not _has_gym, reason="gym not available")
@pytest.mark.parametrize("stateful", [True, False])
def test_stateful_and_stateless_specs(self, stateful):
Expand Down
Loading
Loading