Skip to content

[BUG] TruncatedNormal mean lands outside [low, high], variance goes negative and log_prob is off by tens of nats once loc is a few scales past a bound #4492

Description

@Nicholas022400701

Describe the bug

TruncatedStandardNormal (torchrl/modules/distributions/truncated_normal.py) clamps Phi(a), Phi(b) and Z = Phi(b) - Phi(a) to [1e-6, 1 - 1e-6] and then builds log_Z, the mean, the variance and the entropy from the clamped values. As soon as loc sits a few scales past one of the bounds, both Phi values are below 1e-6 (or above 1 - 1e-6), Z is clamped to 1e-6 instead of its real value and everything built on it is wrong:

  • mean lands outside [low, high]: with low=-1, high=1, loc=2.0, scale=0.1 it is exactly 2.0, with loc=1.5, scale=0.1 it is 1.35. deterministic_sample is mean, so the deterministic action of a TruncatedNormal policy is out of bounds there.
  • variance goes negative (loc=3.0, scale=0.5 gives -0.145).
  • entropy() is 8 to 11 nats too low in three of the four cases below, which is what a PPO or A2C entropy bonus sees.
  • log_prob is off by up to 40 nats (loc=2.0, scale=0.1 at x=0.95: -39.93 against scipy's -0.51) and exp(log_prob) integrates to 7.5e-18 over [low, high] instead of 1.

tanh_loc defaults to False, so an unconstrained policy head can drift into this regime. Only 5 scales past the bound are needed (Phi(-5) = 2.9e-7).

To Reproduce

import torch
from scipy.stats import truncnorm

from torchrl.modules.distributions import TruncatedNormal

low, high = -1.0, 1.0
x = torch.tensor([[0.95]])
for loc, scale in [(1.5, 0.1), (2.0, 0.1), (3.0, 0.5), (-2.5, 0.3)]:
    d = TruncatedNormal(torch.tensor([loc]), torch.tensor([scale]), low=low, high=high)
    ref = truncnorm((low - loc) / scale, (high - loc) / scale, loc=loc, scale=scale)
    print(
        f"loc={loc:5} scale={scale}  mean {d.mean.item():.4f} (scipy {ref.mean():.4f})"
        f"  var {d.variance.item():.2e} (scipy {ref.var():.2e})"
        f"  entropy {d.entropy().item():.2f} (scipy {ref.entropy():.2f})"
        f"  log_prob(0.95) {d.log_prob(x).item():.2f} (scipy {ref.logpdf(0.95):.2f})"
    )

On main (b8b6f0c):

loc=  1.5 scale=0.1  mean 1.3513 (scipy 0.9813)  var 6.22e-02 (scipy 3.27e-04)  entropy -10.98 (scipy -2.98)  log_prob(0.95) 0.07 (scipy 1.32)
loc=  2.0 scale=0.1  mean 2.0000 (scipy 0.9902)  var 1.00e-02 (scipy 9.45e-05)  entropy -14.70 (scipy -3.62)  log_prob(0.95) -39.93 (scipy -0.51)
loc=  3.0 scale=0.5  mean 0.8189 (scipy 0.8872)  var -1.45e-01 (scipy 1.17e-02)  entropy -0.94 (scipy -1.18)  log_prob(0.95) 1.76 (scipy 1.73)
loc= -2.5 scale=0.3  mean -2.0540 (scipy -0.9440)  var 5.60e-01 (scipy 2.94e-03)  entropy -9.88 (scipy -1.88)  log_prob(0.95) -52.02 (scipy -50.77)

Expected behavior

The values scipy's truncnorm gives, which is what the patch below produces:

loc=  1.5 scale=0.1  mean 0.9813 (scipy 0.9813)  var 3.27e-04 (scipy 3.27e-04)  entropy -2.98 (scipy -2.98)  log_prob(0.95) 1.32 (scipy 1.32)
loc=  2.0 scale=0.1  mean 0.9902 (scipy 0.9902)  var 9.36e-05 (scipy 9.45e-05)  entropy -3.62 (scipy -3.62)  log_prob(0.95) -0.51 (scipy -0.51)
loc=  3.0 scale=0.5  mean 0.8872 (scipy 0.8872)  var 1.17e-02 (scipy 1.17e-02)  entropy -1.18 (scipy -1.18)  log_prob(0.95) 1.73 (scipy 1.73)
loc= -2.5 scale=0.3  mean -0.9440 (scipy -0.9440)  var 2.94e-03 (scipy 2.94e-03)  entropy -1.88 (scipy -1.88)  log_prob(0.95) -50.77 (scipy -50.77)

Screenshots

None.

System info

torchrl main (b8b6f0c) from source, torch 2.14.0+cpu, scipy 1.18.1, Python 3.13.14, Linux.

Additional context

The existing tests do not reach this regime: test_truncnormal uses loc=scale for its far locations (a and b stay within about one scale of zero) and test_truncnormal_against_scipy uses loc=0, scale=1.

rsample and icdf are left as they are in the patch below. They still go through the clamped Phi, which in these far cases returns the nearby bound, close to the real distribution (its std is 0.018 for loc=1.5, scale=0.1).

Reason and Possible fixes

Compute log Z with torch.special.log_ndtr instead of clamping. For an interval in the left tail log(Phi(b) - Phi(a)) = log_ndtr(b) + log1p(-exp(log_ndtr(a) - log_ndtr(b))), intervals right of zero are mirrored into the left tail and intervals containing zero use log1p(-ndtr(a) - ndtr(-b)), the same case split as scipy's _log_gauss_mass. phi(a) / Z and phi(b) / Z are then exp(log phi - log Z), and the mean, variance and entropy are built from those two ratios, so nothing tiny is ever divided by something clamped. log_ndtr has no half precision kernel on CPU, so the helper computes in float32 for float16 and bfloat16 inputs and casts back.

With the patch:

  • float64 matches scipy to about 1e-14 on mean, variance, entropy and log_prob for all four cases above and for a grid of 99 points in [low, high].
  • float32 mean, entropy and log_prob stay within 3e-4 of scipy up to 40 scales past the bound. The float32 variance carries a relative error of 1.4e-4 at 5 scales, 0.9% at 10 and 35% at 20, which is the float32 precision of log Z (about -200 at 20 scales) entering 1 - b phi(b)/Z - (phi(b)/Z)**2; on main it is 190 times too large already at 5 scales.
  • Autograd of log_prob + entropy + mean with respect to loc and scale matches central finite differences to 1e-6 in float64 at these locations, and rsample gradients stay finite.
  • pytest test/test_distributions.py passes 995 tests including the float16 rsample ones. The new test_truncnormal_loc_beyond_bounds_against_scipy (5 loc, scale pairs x float32/float64) fails 10 of 10 on main, on the low <= mean <= high or variance > 0 assert, and passes with the patch.
  • ufmt and flake8 are clean.

Branch with the patch and the test: https://github.com/Nicholas022400701/rl/tree/fix/truncnorm-log-mass

Patch
diff --git a/test/test_distributions.py b/test/test_distributions.py
index e2e7850..5794e07 100644
--- a/test/test_distributions.py
+++ b/test/test_distributions.py
@@ -678,6 +678,51 @@ class TestTruncatedNormal:
             pi_x, torch.as_tensor(pdf_scypi_truncnorm, dtype=torch.float32)
         )
 
+    @pytest.mark.skipif(not _has_scipy, reason="scipy not installed")
+    @pytest.mark.parametrize("dtype", [torch.float32, torch.float64])
+    @pytest.mark.parametrize(
+        "loc,scale", [(1.5, 0.1), (2.0, 0.1), (3.0, 0.5), (-2.5, 0.3), (-4.0, 0.5)]
+    )
+    def test_truncnormal_loc_beyond_bounds_against_scipy(self, loc, scale, dtype):
+        # loc sits several scales outside [low, high], so Phi(a) and Phi(b) are
+        # both tiny and the moments must be formed from ratios, not from Z
+        from scipy.stats import truncnorm as sp_truncnorm
+
+        low, high = -1.0, 1.0
+        d = TruncatedNormal(
+            torch.tensor([loc], dtype=dtype),
+            torch.tensor([scale], dtype=dtype),
+            low=low,
+            high=high,
+            tanh_loc=False,
+        )
+        ref = sp_truncnorm(
+            (low - loc) / scale, (high - loc) / scale, loc=loc, scale=scale
+        )
+        assert low <= d.mean.item() <= high
+        assert d.variance.item() > 0
+        tol = (
+            {"rtol": 5e-2, "atol": 1e-4}
+            if dtype is torch.float32
+            else {"rtol": 1e-5, "atol": 1e-7}
+        )
+        torch.testing.assert_close(
+            d.mean, torch.tensor([ref.mean()], dtype=dtype), **tol
+        )
+        torch.testing.assert_close(
+            d.variance, torch.tensor([ref.var()], dtype=dtype), **tol
+        )
+        torch.testing.assert_close(
+            d.entropy(), torch.tensor(ref.entropy(), dtype=dtype), **tol
+        )
+        # interior points: log_prob nudges values off the bounds by eps first
+        x = torch.linspace(low, high, 101, dtype=dtype)[1:-1].unsqueeze(-1)
+        torch.testing.assert_close(
+            d.log_prob(x),
+            torch.tensor(ref.logpdf(x.squeeze(-1).numpy()), dtype=dtype),
+            **tol,
+        )
+
     @pytest.mark.parametrize(
         "min", [-torch.ones(3), -1, 3 * torch.tensor([-1.0, -2.0, -0.5]), -0.1]
     )
diff --git a/torchrl/modules/distributions/truncated_normal.py b/torchrl/modules/distributions/truncated_normal.py
index 24a3a55..b8839a5 100644
--- a/torchrl/modules/distributions/truncated_normal.py
+++ b/torchrl/modules/distributions/truncated_normal.py
@@ -20,6 +20,28 @@ CONST_LOG_INV_SQRT_2PI = math.log(CONST_INV_SQRT_2PI)
 CONST_LOG_SQRT_2PI_E = 0.5 * math.log(2 * math.pi * math.e)
 
 
+def _log_gauss_mass(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
+    """log(Phi(b) - Phi(a)) for a < b, evaluated in the left tail so it stays accurate far from zero."""
+    dtype = a.dtype
+    if dtype in (torch.float16, torch.bfloat16):
+        # log_ndtr has no half precision kernel on CPU
+        a, b = a.float(), b.float()
+    # Phi(b) - Phi(a) = Phi(-a) - Phi(-b): mirror intervals right of zero into the left tail
+    right = a > 0
+    a, b = torch.where(right, -b, a), torch.where(right, -a, b)
+    left = b <= 0
+    # both bounds in the left tail: log Phi(b) + log(1 - Phi(a) / Phi(b))
+    log_ndtr_a = torch.special.log_ndtr(torch.where(left, a, -1.0))
+    log_ndtr_b = torch.special.log_ndtr(torch.where(left, b, 0.0))
+    mass_left = log_ndtr_b + torch.log1p(-(log_ndtr_a - log_ndtr_b).exp())
+    # interval containing zero: Phi(b) - Phi(a) = 1 - Phi(a) - Phi(-b)
+    mass_mid = torch.log1p(
+        -torch.special.ndtr(torch.where(left, -1.0, a))
+        - torch.special.ndtr(torch.where(left, -1.0, -b))
+    )
+    return torch.where(left, mass_left, mass_mid).to(dtype)
+
+
 class TruncatedStandardNormal(Distribution):
     """Truncated Standard Normal distribution.
 
@@ -60,20 +82,20 @@ class TruncatedStandardNormal(Distribution):
         self._little_phi_b = self._little_phi(self.b)
         self._big_phi_a = self._big_phi(self.a)
         self._big_phi_b = self._big_phi(self.b)
-        self._Z = (self._big_phi_b - self._big_phi_a).clamp(eps, 1 - eps)
-        self._log_Z = self._Z.log()
+        self._log_Z = _log_gauss_mass(self.a, self.b)
+        self._Z = self._log_Z.exp()
+        # phi(a) / Z and phi(b) / Z as exp(log phi - log Z): both vanish in the
+        # tails, only their ratio to Z is well defined there
+        little_phi_a_d_Z = (self._log_little_phi(self.a) - self._log_Z).exp()
+        little_phi_b_d_Z = (self._log_little_phi(self.b) - self._log_Z).exp()
         little_phi_coeff_a = torch.nan_to_num(self.a, nan=math.nan)
         little_phi_coeff_b = torch.nan_to_num(self.b, nan=math.nan)
         self._lpbb_m_lpaa_d_Z = (
-            self._little_phi_b * little_phi_coeff_b
-            - self._little_phi_a * little_phi_coeff_a
-        ) / self._Z
-        self._mean = -(self._little_phi_b - self._little_phi_a) / self._Z
-        self._variance = (
-            1
-            - self._lpbb_m_lpaa_d_Z
-            - ((self._little_phi_b - self._little_phi_a) / self._Z) ** 2
+            little_phi_b_d_Z * little_phi_coeff_b
+            - little_phi_a_d_Z * little_phi_coeff_a
         )
+        self._mean = little_phi_a_d_Z - little_phi_b_d_Z
+        self._variance = 1 - self._lpbb_m_lpaa_d_Z - self._mean**2
         self._entropy = CONST_LOG_SQRT_2PI_E + self._log_Z - 0.5 * self._lpbb_m_lpaa_d_Z
 
     @constraints.dependent_property
@@ -103,6 +125,10 @@ class TruncatedStandardNormal(Distribution):
     def _little_phi(x):
         return (-(x**2) * 0.5).exp() * CONST_INV_SQRT_2PI
 
+    @staticmethod
+    def _log_little_phi(x):
+        return CONST_LOG_INV_SQRT_2PI - (x**2) * 0.5
+
     def _big_phi(self, x):
         phi = 0.5 * (1 + (x * CONST_INV_SQRT_2).erf())
         return phi.clamp(self.eps, 1 - self.eps)

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)

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