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)
Describe the bug
TruncatedStandardNormal(torchrl/modules/distributions/truncated_normal.py) clampsPhi(a),Phi(b)andZ = Phi(b) - Phi(a)to[1e-6, 1 - 1e-6]and then buildslog_Z, the mean, the variance and the entropy from the clamped values. As soon aslocsits a fewscales past one of the bounds, bothPhivalues are below1e-6(or above1 - 1e-6),Zis clamped to1e-6instead of its real value and everything built on it is wrong:meanlands outside[low, high]: withlow=-1, high=1,loc=2.0, scale=0.1it is exactly2.0, withloc=1.5, scale=0.1it is1.35.deterministic_sampleismean, so the deterministic action of aTruncatedNormalpolicy is out of bounds there.variancegoes negative (loc=3.0, scale=0.5gives-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_probis off by up to 40 nats (loc=2.0, scale=0.1atx=0.95:-39.93against scipy's-0.51) andexp(log_prob)integrates to7.5e-18over[low, high]instead of 1.tanh_locdefaults toFalse, 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
On main (b8b6f0c):
Expected behavior
The values scipy's
truncnormgives, which is what the patch below produces: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_truncnormalusesloc=scalefor its far locations (aandbstay within about one scale of zero) andtest_truncnormal_against_scipyusesloc=0, scale=1.rsampleandicdfare left as they are in the patch below. They still go through the clampedPhi, which in these far cases returns the nearby bound, close to the real distribution (its std is0.018forloc=1.5, scale=0.1).Reason and Possible fixes
Compute
log Zwithtorch.special.log_ndtrinstead of clamping. For an interval in the left taillog(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 uselog1p(-ndtr(a) - ndtr(-b)), the same case split as scipy's_log_gauss_mass.phi(a) / Zandphi(b) / Zare thenexp(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_ndtrhas no half precision kernel on CPU, so the helper computes in float32 for float16 and bfloat16 inputs and casts back.With the patch:
1e-14on mean, variance, entropy andlog_probfor all four cases above and for a grid of 99 points in[low, high].log_probstay within3e-4of scipy up to 40 scales past the bound. The float32 variance carries a relative error of1.4e-4at 5 scales,0.9%at 10 and35%at 20, which is the float32 precision oflog Z(about-200at 20 scales) entering1 - b phi(b)/Z - (phi(b)/Z)**2; on main it is 190 times too large already at 5 scales.log_prob + entropy + meanwith respect tolocandscalematches central finite differences to1e-6in float64 at these locations, andrsamplegradients stay finite.pytest test/test_distributions.pypasses 995 tests including the float16rsampleones. The newtest_truncnormal_loc_beyond_bounds_against_scipy(5loc, scalepairs x float32/float64) fails 10 of 10 on main, on thelow <= mean <= highorvariance > 0assert, and passes with the patch.ufmtandflake8are clean.Branch with the patch and the test: https://github.com/Nicholas022400701/rl/tree/fix/truncnorm-log-mass
Patch
Checklist