Skip to content

[BUG] VecNormV2 subtracts loc**2 before the decay bias correction, so scale is about |mean| instead of std for tens of thousands of steps #4491

Description

@Nicholas022400701

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

  • I have checked that there is no similar issue in the repo (required)
  • I have read the documentation (required)
  • I have provided a minimal working example to reproduce the bug (required)

No activity

Activity on this issue will appear here.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions