diff --git a/test/transforms/test_normalization.py b/test/transforms/test_normalization.py index 330799537c1..f4dcca7c506 100644 --- a/test/transforms/test_normalization.py +++ b/test/transforms/test_normalization.py @@ -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, @@ -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, @@ -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): diff --git a/torchrl/envs/transforms/vecnorm.py b/torchrl/envs/transforms/vecnorm.py index 0c9f4c1b69d..79120cdca5f 100644 --- a/torchrl/envs/transforms/vecnorm.py +++ b/torchrl/envs/transforms/vecnorm.py @@ -10,7 +10,6 @@ from collections import OrderedDict from collections.abc import Sequence from copy import copy - from typing import Any import torch @@ -24,10 +23,8 @@ from torch._dynamo import is_compiling from torchrl.data.tensor_specs import Bounded, Composite, Unbounded - from torchrl.envs.common import EnvBase from torchrl.envs.transforms.transforms import Compose, ObservationNorm, Transform - from torchrl.envs.transforms.utils import _set_missing_tolerance @@ -46,7 +43,10 @@ class VecNormV2(Transform): Stateful vs. Stateless: Stateful Mode (`stateful=True`): - - Maintains internal statistics (`loc`, `var`, `count`) for normalization. + - Maintains internal statistics (`loc`, `var`, `count`) for normalization. `loc` and `var` are the + bias-corrected, exponentially weighted mean and variance of the data. They are updated + incrementally from the deviation to the current mean (as in Welford's algorithm), which keeps the + variance accurate for data with a large offset compared to its spread. - Updates statistics with each call unless frozen. - `state_dict` returns the current statistics. - `load_state_dict` updates the internal statistics with the provided state. @@ -445,7 +445,7 @@ def _step( next_tensordict_select, loc, var, count ) next_tensordict_norm = self._stateless_norm( - next_tensordict_select, loc, var, count + next_tensordict_select, loc, var ) # updates have been done in-place, we're good next_tensordict_norm.set(f"{self.prefix}_loc", loc) @@ -568,24 +568,13 @@ def _in_keys_safe(self): return self.in_keys[:-3] return self.in_keys - def _norm(self, data, loc, var, count): + def _norm(self, data, loc, var): if self.missing_tolerance: loc = loc.select(*data.keys(True, True)) var = var.select(*data.keys(True, True)) - count = count.select(*data.keys(True, True)) if loc.is_empty(): return data - if self.decay < 1.0: - bias_correction = 1 - (count * math.log(self.decay)).exp() - bias_correction = bias_correction.apply(lambda x, y: x.to(y.dtype), data) - else: - bias_correction = 1 - - var = var - loc.pow(2) - loc = loc / bias_correction - var = var / bias_correction - scale = var.sqrt().clamp_min(self.eps) data_update = (data - loc) / scale @@ -601,15 +590,35 @@ def _norm(self, data, loc, var, count): def _stateful_norm(self, data): loc = self._loc var = self._var - count = self._count data_device = data.device if data_device is not None: loc = loc.to(data_device, non_blocking=True) var = var.to(data_device, non_blocking=True) - count = count.to(data_device, non_blocking=True) if data_device.type == "cuda": torch.cuda.current_stream(data_device).synchronize() - return self._norm(data, loc, var, count) + return self._norm(data, loc, var) + + def _update_weight(self, count, n, ref): + """Weight of the latest ``n`` samples in the bias-corrected moving average after ``count`` samples.""" + if self.decay != 1.0: + log_decay = math.log(self.decay) + weight = math.expm1(n * log_decay) / (count * log_decay).expm1() + else: + weight = n / count + return weight.apply(lambda w, r: w.to(r.dtype), ref) + + @staticmethod + def _var_target(loc, batch_mean, batch_var, weight): + """Value ``var`` is interpolated towards so that it tracks the weighted variance. + + Merging a batch of mean ``batch_mean`` and variance ``batch_var`` with weight ``w`` into stats + ``(loc, var)`` gives ``(1 - w) * var + w * (batch_var + (1 - w) * (batch_mean - loc) ** 2)``. Only the + deviation from the current ``loc`` is squared, so large offsets do not cancel out in floating point. + """ + target = (batch_mean - loc).pow(2) * (1 - weight) + if batch_var is not None: + target = target + batch_var + return target def _stateful_update(self, data): if self.frozen: @@ -629,34 +638,29 @@ def _stateful_update(self, data): # The second approach would be to average the data, but that would mean that having one vecnorm per batched # env or one per sub-env will lead to different results as a batch of N elements will actually be # considered as a single one. - # What we go for instead is to average the data (and its squared value) then do the moving average with - # adapted decay. + # What we go for instead is to compute the mean and variance of the batch, then merge them in the + # moving average with adapted decay. n = data.numel() - count += n - data2 = data.pow(2).mean(dim=tuple(range(data.ndim))) - data_mean = data.mean(dim=tuple(range(data.ndim))) - if self.decay != 1.0: - weight = 1 - self.decay**n - else: - weight = n / count + dims = tuple(range(data.ndim)) + data_mean = data.mean(dim=dims) + data_var = (data - data_mean.expand(data.batch_size)).pow(2).mean(dim=dims) + data_var = data_var.named_apply( + lambda name, dv, ref: dv.to(ref.device), + var, + ) else: - count += 1 - data2 = data.pow(2) + n = 1 data_mean = data - if self.decay != 1.0: - weight = 1 - self.decay - else: - weight = 1 / count + data_var = None + count += n + weight = self._update_weight(count, n, loc) data_mean = data_mean.named_apply( lambda name, dm, ref: dm.to(ref.device), loc, ) - data2 = data2.named_apply( - lambda name, d2, ref: d2.to(ref.device), - var, - ) + var_target = self._var_target(loc, data_mean, data_var, weight) loc.lerp_(end=data_mean, weight=weight) - var.lerp_(end=data2, weight=weight) + var.lerp_(end=var_target, weight=weight) def _maybe_stateless_init(self, data): if not self.initialized or f"{self.prefix}_loc" not in data.keys(): @@ -683,8 +687,8 @@ def _maybe_stateless_init(self, data): data[f"{self.prefix}_loc"] = loc data[f"{self.prefix}_var"] = var - def _stateless_norm(self, data, loc, var, count): - data = self._norm(data, loc, var, count) + def _stateless_norm(self, data, loc, var): + data = self._norm(data, loc, var) return data def _stateless_update(self, data, loc, var, count): @@ -692,12 +696,9 @@ def _stateless_update(self, data, loc, var, count): return loc, var, count count = count + 1 data = self._maybe_cast_to_float(data) - if self.decay != 1.0: - weight = 1 - self.decay - else: - weight = 1 / count + weight = self._update_weight(count, 1, loc) + var = var.lerp(end=self._var_target(loc, data, None, weight), weight=weight) loc = loc.lerp(end=data, weight=weight) - var = var.lerp(end=data.pow(2), weight=weight) return loc, var, count def transform_observation_spec(self, observation_spec: Composite) -> Composite: @@ -846,20 +847,10 @@ def to_observation_norm(self) -> Compose | ObservationNorm: def _get_loc_scale(self, loc_only: bool = False) -> tuple: if self.stateful: - loc = self._loc - count = self._count - if self.decay != 1.0: - bias_correction = 1 - (count * math.log(self.decay)).exp() - bias_correction = bias_correction.apply(lambda x, y: x.to(y.dtype), loc) - else: - bias_correction = 1 + loc = self._loc.clone() if loc_only: - return loc / bias_correction, None - var = self._var - var = var - loc.pow(2) - loc = loc / bias_correction - var = var / bias_correction - scale = var.sqrt().clamp_min(self.eps) + return loc, None + scale = self._var.sqrt().clamp_min(self.eps) return loc, scale else: raise RuntimeError("_get_loc_scale() called on stateless vecnorm.") @@ -879,6 +870,7 @@ def __setstate__(self, state: dict[str, Any]): super().__setstate__(state) SEP = ".-|-." + _STATS_VERSION = 2 def set_extra_state(self, state: OrderedDict) -> None: if not self.stateful: @@ -891,6 +883,7 @@ def set_extra_state(self, state: OrderedDict) -> None: "set_extra_state() called with a void state-dict while the instance is initialized." ) td = TensorDict(state).unflatten_keys(self.SEP) + version = td.pop("version", None) if self._loc is None and not all(v.is_shared() for v in td.values(True, True)): warnings.warn( "VecNorm wasn't initialized and the tensordict is not shared. In single " @@ -901,10 +894,29 @@ def set_extra_state(self, state: OrderedDict) -> None: "may not have access to the loaded TensorDict." ) td.share_memory_() + if version is None: + self._convert_legacy_stats(td["loc"], td["var"], td["count"]) self._loc = td["loc"] self._var = td["var"] self._count = td["count"] + def _convert_legacy_stats(self, loc, var, count): + """Converts, in place, stats saved before v2 (biased moving averages of ``x`` and ``x ** 2``) to mean and variance. + + The bias correction uses this instance's ``decay``, which is assumed to match the one the stats were built with. + """ + if self.decay != 1.0: + bias_correction = -(count * math.log(self.decay)).expm1() + bias_correction = bias_correction.apply( + lambda x, y: x.to(y.dtype).clamp_min(torch.finfo(y.dtype).tiny), loc + ) + else: + bias_correction = 1 + mean = loc / bias_correction + variance = (var / bias_correction - mean.pow(2)).clamp_min(0) + loc.copy_(mean) + var.copy_(variance) + def get_extra_state(self) -> OrderedDict: if not self.stateful: return {} @@ -920,6 +932,7 @@ def get_extra_state(self) -> OrderedDict: loc=self._loc, var=self._var, count=self._count, + version=torch.tensor(self._STATS_VERSION), ) return td.flatten_keys(self.SEP).to_dict()