Skip to content

[BugFix] Track VecNormV2 variance with a Welford-style update - #4516

Draft
theap06 wants to merge 2 commits into
pytorch:mainfrom
theap06:vecnorm-welford
Draft

theap06 wants to merge 2 commits into
pytorch:mainfrom
theap06:vecnorm-welford

Conversation

@theap06

@theap06 theap06 commented Oct 2, 2026

Copy link
Copy Markdown
Collaborator

Stacked on #4514 (its commit is the first one here); I'll rebase once it lands.

Problem

VecNormV2 stores moving averages of x and x**2 and computes the variance as E[x**2] - E[x]**2. In float32 this cancels out when the data offset is large relative to its spread. With observations at 1e4 + N(0, 1) and decay=0.99:

offset main #4514 this PR
1e2 std 0.97 std 0.99 matches float64 reference
1e3 std 0.75 std 0.91 matches float64 reference
1e4 NaN finite values up to ~3e4 matches float64 reference

Change

  • _loc / _var now hold the bias-corrected, exponentially weighted mean and variance (they used to hold biased moving averages of x and x**2). The update uses weight w = (1 - decay**n) / (1 - decay**count) (n / count for decay=1), so the stored stats are already bias-corrected and _norm / loc / scale just read them.
  • Variance update: var <- (1 - w) * var + w * (batch_var + (1 - w) * (batch_mean - loc)**2). Only the deviation from the current mean is squared, so the offset never enters a subtraction. For a single sample batch_var = 0; with reduce_batch_dims=True the batch variance is computed two-pass around the batch mean.
  • This is exact algebra for the same estimator [BugFix] Apply VecNormV2 bias correction before subtracting loc**2 #4514 computes; only the floating-point behaviour changes.
  • Checkpoints: get_extra_state now writes a version entry. A state without it is treated as the old format and converted in place on load (keeps shared-memory stats shared), using the instance's decay.
  • Stateless mode carries the same _vecnorm_loc / _vecnorm_var keys, which now hold mean/variance. They are re-initialised at reset, so only tensordicts saved mid-trajectory by an older version would be read differently.

Tests

  • test_vecnorm2_large_offset (stateful, stateless, reduce_batch_dims on a SerialEnv): offset 1e4, compares every normalized output and the final loc / scale to a float64 reference computed from the observations. Fails on main and on [BugFix] Apply VecNormV2 bias correction before subtracting loc**2 #4514, passes here.
  • test_vecnorm2_load_legacy_state_dict: loads an old-format state built from known samples and checks loc / scale, then round-trips the new format.
  • test/transforms/test_normalization.py and the VecNorm / extra_state tests in collectors, evaluator, configs, render, rb and auto-reset pass locally.

Cost

CPU microbenchmark of _step (17-dim obs): single-sample path 137 -> 131 us/step, reduce_batch_dims with batch 64: 163 -> 201 us/step (extra pass for the batch variance).

Mohit-Ak and others added 2 commits October 1, 2026 02:13
The exponential moving averages of x and x**2 were bias-corrected after
loc**2 had already been subtracted, leaving a decay**count * mean**2 term
in the variance. Divide first, then subtract, in both _norm and
_get_loc_scale; use expm1 for the correction and clamp the variance at 0.

Fixes pytorch#4491
VecNormV2 stored moving averages of x and x**2 and computed the variance
as E[x**2] - E[x]**2. In float32 this cancels catastrophically when the
data offset is large compared to its spread (NaN on main, eps-scaled
outputs after the bias-correction fix at offset 1e4).

loc and var now hold the bias-corrected weighted mean and variance and
are updated from the deviation to the current mean. Batches merged with
reduce_batch_dims use a two-pass batch variance. Checkpoints saved in the
old format are converted on load.

[skip ci]
@pytorch-bot

pytorch-bot Bot commented Oct 2, 2026 •

Copy link
Copy Markdown

🔗 Helpful Links

🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/rl/4516

Note: Links to docs will display an error until the docs builds have been completed.

✅ No Failures

As of commit 8cf68b4 with merge base 0c6f682 (image):
💚 Looks good so far! There are no failures yet. 💚

This comment was automatically generated by Dr. CI and updates every 15 minutes.

@meta-cla meta-cla Bot added the CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. label Oct 2, 2026

This branch has not been deployed

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

Labels

BugFix CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. Transforms

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants