Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
49 changes: 49 additions & 0 deletions test/test_distributions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
11 changes: 5 additions & 6 deletions torchrl/modules/distributions/discrete.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading