Describe the bug
VecNormV2 keeps loc and var as exponential moving averages of x and x**2 that start from zero, and corrects their bias with 1 - decay**count. _norm and _get_loc_scale subtract loc**2 before dividing by the bias correction:
var = var - loc.pow(2)
loc = loc / bias_correction
var = var / bias_correction
Since loc ~ (1 - d**t) * mean and var ~ (1 - d**t) * E[x**2], this gives
(var - loc**2) / (1 - d**t) = E[x**2] - (1 - d**t) * mean**2 = std**2 + d**t * mean**2
instead of std**2. Any observation or reward with a non-zero mean is scaled by roughly |mean| rather than by its standard deviation, and with the default decay=0.9999 the d**t * mean**2 term needs about 50k updates to fade (0.9999**10000 = 0.37). The old sum-based VecNorm computes ssq / count - loc**2 and does not have this problem. test_vecnorm_rollout does not catch it because Pendulum observations are close to zero mean and it uses decay=0.9, and test_vecnorm2_decay1 uses decay=1, which skips the bias correction.
To Reproduce
Observations are 100 + N(0, 1), VecNormV2 with the default decay=0.9999:
import torch
from tensordict import TensorDict
from torchrl.data import Composite, Unbounded
from torchrl.envs import EnvBase, VecNormV2
class OffsetEnv(EnvBase):
"""Observations are 100 + N(0, 1)."""
def __init__(self):
super().__init__()
self.observation_spec = Composite(observation=Unbounded(()))
self.action_spec = Composite(action=Unbounded(()))
self.reward_spec = Composite(reward=Unbounded((1,)))
def _reset(self, tensordict=None, **kwargs):
return TensorDict({"observation": 100 + torch.randn(()), "done": torch.zeros((), dtype=torch.bool)})
def _step(self, tensordict):
return TensorDict({"observation": 100 + torch.randn(()), "reward": torch.zeros(1), "done": torch.zeros((), dtype=torch.bool)})
def _set_seed(self, seed):
pass
torch.manual_seed(0)
for steps in (10, 100, 1000, 10000):
env = OffsetEnv().append_transform(VecNormV2(in_keys=["observation"], out_keys=["obs_norm"])) # decay=0.9999
r = env.rollout(steps, break_when_any_done=False)
obs, norm = r["next", "observation"], r["next", "obs_norm"]
print(
f"steps={steps:5d} sample mean={obs.mean():7.3f} std={obs.std(unbiased=False):.3f} "
f"VecNormV2 loc={env.transform.loc['observation']:7.3f} scale={env.transform.scale['observation']:7.3f} "
f"std of the normalized second half={norm[steps // 2:].std():.4f}"
)
main (b8b6f0c):
steps= 10 sample mean= 99.281 std=0.702 VecNormV2 loc= 99.485 scale= 99.435 std of the normalized second half=0.0056
steps= 100 sample mean=100.065 std=1.003 VecNormV2 loc=100.079 scale= 99.580 std of the normalized second half=0.0100
steps= 1000 sample mean=100.027 std=1.013 VecNormV2 loc=100.028 scale= 95.150 std of the normalized second half=0.0098
steps=10000 sample mean= 99.991 std=0.996 VecNormV2 loc= 99.992 scale= 60.652 std of the normalized second half=0.0145
scale follows sqrt(std**2 + 0.9999**t * mean**2): 60.7 after 10k steps is sqrt(1 + 0.37 * 100**2). The normalized observations have a spread of 0.01 instead of 1.
With the patch below:
steps= 10 sample mean= 99.281 std=0.702 VecNormV2 loc= 99.486 scale= 0.933 std of the normalized second half=0.5762
steps= 100 sample mean=100.065 std=1.003 VecNormV2 loc=100.079 scale= 1.010 std of the normalized second half=1.0005
steps= 1000 sample mean=100.027 std=1.013 VecNormV2 loc=100.028 scale= 1.015 std of the normalized second half=0.9093
steps=10000 sample mean= 99.991 std=0.996 VecNormV2 loc= 99.992 scale= 0.987 std of the normalized second half=0.9887
Expected behavior
scale should match the standard deviation of the data at every step count, and the normalized data should have unit spread once the running mean has settled, as with VecNorm.
System info
torchrl main (b8b6f0c) from source, torch 2.14.0+cpu, Python 3.13.14, Linux.
Reason and Possible fixes
Divide both moving averages by the bias correction first, then subtract:
loc = loc / bias_correction
var = var / bias_correction
var = (var - loc.pow(2)).clamp_min(0)
Two small companions to that reorder, both found while testing it:
1 - exp(count * log(decay)) in float32 has only about four significant digits for small counts, and with the correct subtraction order that relative error enters the variance multiplied by mean**2. With 10 samples of 100 + N(0, 1) (sample std 1.0707) the reordered code still returned scale=1.1592; -expm1(count * log(decay)) is the same quantity without the cancellation and returns 1.0707.
- The
clamp_min(0): after a single update var/bc - (loc/bc)**2 is exactly zero in math, but rounding can leave a tiny negative value that sqrt turns into nan. test_denorm[True] fails without it.
The stats update in _stateful_update / _stateless_update is unchanged. Patch against main, with a regression test (test_vecnorm2_decay_bias_correction, stateful and stateless) that fails on main with assert 0.5 < tensor(0.0107) and passes with the fix; test/transforms/test_normalization.py is otherwise green (22 passed, 15 skipped for gym) and ufmt / flake8 are clean. Happy to open the PR if the approach looks right.
Patch
diff --git a/test/transforms/test_normalization.py b/test/transforms/test_normalization.py
index 3307995..e1b3a83 100644
--- a/test/transforms/test_normalization.py
+++ b/test/transforms/test_normalization.py
@@ -121,6 +121,54 @@ class TestVecNormV2:
assert env.transform._loc.ndim == 0
assert env.transform._var.ndim == 0
+ @pytest.mark.parametrize("stateful", [True, False])
+ def test_vecnorm2_decay_bias_correction(self, stateful):
+ # The running stats are exponential moving averages of x and x^2 that start from zero. Both have
+ # to be bias-corrected before the variance is formed, otherwise an observation with a large offset
+ # is normalized by sqrt(E[x^2]) ~ |mean| instead of by its standard deviation.
+ offset = 100.0
+
+ class OffsetEnv(self.SimpleEnv):
+ def _reset(self, tensordict, **kwargs):
+ tensordict = super()._reset(tensordict, **kwargs)
+ tensordict["observation"] = tensordict["observation"] + offset
+ return tensordict
+
+ def _step(self, tensordict):
+ tensordict = super()._step(tensordict)
+ tensordict["observation"] = tensordict["observation"] + offset
+ return tensordict
+
+ torch.manual_seed(0)
+ env = OffsetEnv().append_transform(
+ VecNormV2(
+ in_keys=["observation"],
+ out_keys=["obs_norm"],
+ decay=0.9999,
+ stateful=stateful,
+ )
+ )
+ rollout = env.rollout(50, break_when_any_done=False)
+ obs = rollout["next", "observation"]
+ # once the running mean has settled, the normalized observations have unit spread
+ obs_norm = rollout["next", "obs_norm"][-25:]
+ assert 0.5 < obs_norm.std() < 2.0, obs_norm.std()
+ if stateful:
+ torch.testing.assert_close(
+ env.transform.loc["observation"], obs.mean(), atol=0.05, rtol=0
+ )
+ torch.testing.assert_close(
+ env.transform.scale["observation"],
+ obs.std(unbiased=False),
+ atol=0.05,
+ rtol=0,
+ )
+ torch.testing.assert_close(
+ rollout[-1]["next", "obs_norm"],
+ (obs[-1] - env.transform.loc["observation"])
+ / env.transform.scale["observation"],
+ )
+
@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 0c9f4c1..bd58213 100644
--- a/torchrl/envs/transforms/vecnorm.py
+++ b/torchrl/envs/transforms/vecnorm.py
@@ -577,14 +577,18 @@ class VecNormV2(Transform):
return data
if self.decay < 1.0:
- bias_correction = 1 - (count * math.log(self.decay)).exp()
+ # -expm1(count * log(decay)) is 1 - decay**count without the cancellation of 1 - exp(x)
+ bias_correction = -(count * math.log(self.decay)).expm1()
bias_correction = bias_correction.apply(lambda x, y: x.to(y.dtype), data)
else:
bias_correction = 1
- var = var - loc.pow(2)
+ # loc and var are exponential moving averages of x and x^2 that start from zero, so both must be
+ # bias-corrected before the variance is formed: var/bc - (loc/bc)^2, not (var - loc^2)/bc.
loc = loc / bias_correction
var = var / bias_correction
+ # rounding can leave a tiny negative variance, which sqrt would turn into nan
+ var = (var - loc.pow(2)).clamp_min(0)
scale = var.sqrt().clamp_min(self.eps)
@@ -849,16 +853,16 @@ class VecNormV2(Transform):
loc = self._loc
count = self._count
if self.decay != 1.0:
- bias_correction = 1 - (count * math.log(self.decay)).exp()
+ bias_correction = -(count * math.log(self.decay)).expm1()
bias_correction = bias_correction.apply(lambda x, y: x.to(y.dtype), loc)
else:
bias_correction = 1
if loc_only:
return loc / bias_correction, None
var = self._var
- var = var - loc.pow(2)
loc = loc / bias_correction
var = var / bias_correction
+ var = (var - loc.pow(2)).clamp_min(0)
scale = var.sqrt().clamp_min(self.eps)
return loc, scale
else:
Checklist
Describe the bug
VecNormV2keepslocandvaras exponential moving averages ofxandx**2that start from zero, and corrects their bias with1 - decay**count._normand_get_loc_scalesubtractloc**2before dividing by the bias correction:Since
loc ~ (1 - d**t) * meanandvar ~ (1 - d**t) * E[x**2], this givesinstead of
std**2. Any observation or reward with a non-zero mean is scaled by roughly|mean|rather than by its standard deviation, and with the defaultdecay=0.9999thed**t * mean**2term needs about 50k updates to fade (0.9999**10000 = 0.37). The old sum-basedVecNormcomputesssq / count - loc**2and does not have this problem.test_vecnorm_rolloutdoes not catch it because Pendulum observations are close to zero mean and it usesdecay=0.9, andtest_vecnorm2_decay1usesdecay=1, which skips the bias correction.To Reproduce
Observations are
100 + N(0, 1),VecNormV2with the defaultdecay=0.9999:main(b8b6f0c):scalefollowssqrt(std**2 + 0.9999**t * mean**2): 60.7 after 10k steps issqrt(1 + 0.37 * 100**2). The normalized observations have a spread of 0.01 instead of 1.With the patch below:
Expected behavior
scaleshould match the standard deviation of the data at every step count, and the normalized data should have unit spread once the running mean has settled, as withVecNorm.System info
torchrl main (b8b6f0c) from source, torch 2.14.0+cpu, Python 3.13.14, Linux.
Reason and Possible fixes
Divide both moving averages by the bias correction first, then subtract:
Two small companions to that reorder, both found while testing it:
1 - exp(count * log(decay))in float32 has only about four significant digits for small counts, and with the correct subtraction order that relative error enters the variance multiplied bymean**2. With 10 samples of100 + N(0, 1)(sample std 1.0707) the reordered code still returnedscale=1.1592;-expm1(count * log(decay))is the same quantity without the cancellation and returns 1.0707.clamp_min(0): after a single updatevar/bc - (loc/bc)**2is exactly zero in math, but rounding can leave a tiny negative value thatsqrtturns intonan.test_denorm[True]fails without it.The stats update in
_stateful_update/_stateless_updateis unchanged. Patch againstmain, with a regression test (test_vecnorm2_decay_bias_correction, stateful and stateless) that fails onmainwithassert 0.5 < tensor(0.0107)and passes with the fix;test/transforms/test_normalization.pyis otherwise green (22 passed, 15 skipped for gym) andufmt/flake8are clean. Happy to open the PR if the approach looks right.Patch
Checklist