diff --git a/test/test_distributions.py b/test/test_distributions.py index e2e785083f1..770459031f9 100644 --- a/test/test_distributions.py +++ b/test/test_distributions.py @@ -921,6 +921,55 @@ def test_entropy(self, neg_inf: float, sparse: bool) -> None: entropy = -entropy.sum(dim=-1) torch.testing.assert_close(dist.entropy(), entropy) + @pytest.mark.parametrize( + "distribution_cls", [MaskedCategorical, MaskedOneHotCategorical] + ) + @pytest.mark.parametrize("padding_value", [-1, 0]) + @pytest.mark.parametrize("neg_inf", [float("-inf"), -10.0]) + @pytest.mark.parametrize("dtype", [torch.float32, torch.float64]) + @pytest.mark.parametrize("device", get_default_devices()) + def test_entropy_sparse_padding( + self, distribution_cls, padding_value, neg_inf, dtype, device + ): + logits = torch.tensor( + [[1.0, 2.0, 3.0, 4.0], [0.0, 3.0, 1.0, 2.0]], dtype=dtype, device=device + ) + logits = logits.expand(2, 2, 4).clone().requires_grad_() + indices = torch.tensor( + [[1, 2, padding_value], [1, 2, 3]], device=device + ).expand(2, 2, 3) + distribution = distribution_cls( + logits=logits, indices=indices, padding_value=padding_value, neg_inf=neg_inf + ) + actual = distribution.entropy() + expected = torch.stack( + [ + torch.distributions.Categorical(logits=logits[:, 0, 1:3]).entropy(), + torch.distributions.Categorical(logits=logits[:, 1, 1:]).entropy(), + ], + dim=-1, + ) + torch.testing.assert_close(actual, expected) + actual_grad = torch.autograd.grad(actual.sum(), logits, retain_graph=True)[0] + expected_grad = torch.autograd.grad(expected.sum(), logits)[0] + torch.testing.assert_close(actual_grad, expected_grad) + + @pytest.mark.parametrize( + "distribution_cls", [MaskedCategorical, MaskedOneHotCategorical] + ) + @pytest.mark.parametrize("sparse", [False, True]) + def test_entropy_zero_probability(self, distribution_cls, sparse): + probs = torch.tensor([0.0, 0.2, 0.3, 0.5]) + mask = torch.tensor([True, False, True, True]) + indices = torch.tensor([0, 2, 3]) + distribution = distribution_cls( + probs=probs, + mask=None if sparse else mask, + indices=indices if sparse else None, + ) + expected = torch.distributions.Categorical(probs=probs[indices]).entropy() + torch.testing.assert_close(distribution.entropy(), expected) + @pytest.mark.parametrize("neg_inf", [-1e20, float("-inf")]) def test_sample_sparse(self, neg_inf: float) -> None: torch.manual_seed(0) diff --git a/torchrl/modules/distributions/discrete.py b/torchrl/modules/distributions/discrete.py index a6a9dc56a4b..8ae19c131c2 100644 --- a/torchrl/modules/distributions/discrete.py +++ b/torchrl/modules/distributions/discrete.py @@ -312,13 +312,12 @@ def entropy(self): # Clamp logits to avoid numerical issues logits = self.logits + mask = ~logits.isfinite() if self._mask.dtype is torch.bool: - mask = expand_as_right(self._mask, logits) - mask = (~mask) | (~logits.isfinite()) - logits = torch.masked_fill(logits, mask, min_real) - else: - # logits are already masked - pass + mask = mask | ~expand_as_right(self._mask, logits) + elif self._padding_value is not None: + mask = mask | (self._mask == self._padding_value) + logits = torch.masked_fill(logits, mask, min_real) logits = logits - logits.logsumexp(-1, keepdim=True) # Get probabilities and mask them