diff --git a/test/objectives/test_values.py b/test/objectives/test_values.py index 581ce05cfae..51f60ca8b48 100644 --- a/test/objectives/test_values.py +++ b/test/objectives/test_values.py @@ -38,6 +38,7 @@ ) from torchrl.objectives.value.functional import ( generalized_advantage_estimate, + reward2go, td0_advantage_estimate, td1_advantage_estimate, td_lambda_advantage_estimate, @@ -1857,6 +1858,65 @@ def test_gae(self, device, gamma, lmbda, N, T, dtype, has_done): torch.testing.assert_close(r1, r2, rtol=1e-4, atol=1e-4) + @pytest.mark.parametrize("device", get_default_devices()) + @pytest.mark.parametrize("lmbda", [0.0, 1e-9]) + def test_vec_estimates_zero_lmbda(self, device, lmbda): + # gamma * lmbda equal to zero or below the truncation threshold of the + # scalar fast paths must match the non-vectorized estimates + torch.manual_seed(0) + gamma = 0.9 + done = torch.zeros(3, 8, 1, device=device, dtype=torch.bool).bernoulli_(0.2) + done[..., -1, :] = True + terminated = done & torch.zeros_like(done).bernoulli_(0.5) + reward = torch.randn(3, 8, 1, device=device) + state_value = torch.randn(3, 8, 1, device=device) + next_state_value = torch.randn(3, 8, 1, device=device) + + r1 = vec_generalized_advantage_estimate( + gamma, + lmbda, + state_value, + next_state_value, + reward, + done=done, + terminated=terminated, + ) + r2 = generalized_advantage_estimate( + gamma, + lmbda, + state_value, + next_state_value, + reward, + done=done, + terminated=terminated, + ) + torch.testing.assert_close(r1, r2, rtol=1e-4, atol=1e-4) + + r1 = vec_td_lambda_advantage_estimate( + gamma, + lmbda, + state_value, + next_state_value, + reward, + done=done, + terminated=terminated, + ) + r2 = td_lambda_advantage_estimate( + gamma, + lmbda, + state_value, + next_state_value, + reward, + done=done, + terminated=terminated, + ) + torch.testing.assert_close(r1, r2, rtol=1e-4, atol=1e-4) + + # with a zero discount the reward-to-go is the reward itself + torch.testing.assert_close( + reward2go(reward, done, gamma=lmbda), reward, rtol=1e-4, atol=1e-4 + ) + @pytest.mark.parametrize("device", get_default_devices()) @pytest.mark.parametrize("N", [(1,), (8,), (7, 3)]) @pytest.mark.parametrize("dtype", [torch.float, torch.double]) diff --git a/torchrl/objectives/value/functional.py b/torchrl/objectives/value/functional.py index b42b3bc7148..5b6f5a691fe 100644 --- a/torchrl/objectives/value/functional.py +++ b/torchrl/objectives/value/functional.py @@ -194,12 +194,13 @@ def _geom_series_like(t, r, thr): if isinstance(r, torch.Tensor): r = r.item() - if r == 0.0: - return torch.zeros_like(t) - elif r >= 1.0: + if r >= 1.0: lim = t.numel() + elif r == 0.0: + lim = 1 # the series is [1, 0, 0, ...] else: - lim = int(math.log(thr) / math.log(r)) + # r ** 0 is 1 whatever r, so the first element is always kept + lim = max(int(math.log(thr) / math.log(r)), 1) rs = torch.full_like(t[:lim], r) rs[0] = 1.0