Describe the bug
MultiStep replaces ("next", "observation") (and other "next" entries) with the
entry n steps ahead, but explicitly excludes the done keys from that gather:
# torchrl/data/postprocs/postprocs.py, _multi_step_func
tensordict_gather = (
tensordict.get("next")
.exclude(*reward_keys, *done_keys)
.gather(-1, idx_to_gather)
)
As a result, for the n - 1 transitions preceding a termination, ("next", "observation")
points to the terminal observation while ("next", "terminated") is still False
(it is the flag of the original 1-step transition).
Losses that read ("next", "terminated") then bootstrap from the terminal state, e.g. in
DistributionalDQNLoss:
Tz = reward + (1 - terminated.to(reward.dtype)) * discount * support
the target becomes R^(k) + γ^k · Z(s_terminal) instead of R^(k). Terminal observations
are never stored as "observation" in the replay buffer, so the network's value at those
states is never trained and the bias is arbitrary. This affects every episode that ends
with a termination (e.g. every CartPole episode), not just edge cases.
The "nonterminal" key written by MultiStep does not cover this: it is False only at
the step where done itself is set, not at the preceding n - 1 steps.
MultiStepTransform shares _multi_step_func and is affected in the same way.
To Reproduce
import torch
from tensordict import TensorDict
from torchrl.data import MultiStep
T = 6
done = torch.zeros(T, 1, dtype=torch.bool)
done[3] = True # episode terminates after step 3; terminal obs = 4
td = TensorDict({
"observation": torch.arange(T).float().unsqueeze(-1),
"next": TensorDict({
"observation": torch.arange(1, T + 1).float().unsqueeze(-1),
"reward": torch.ones(T, 1),
"done": done.clone(),
"terminated": done.clone(),
}, [T]),
}, [T])
out = MultiStep(gamma=0.99, n_steps=3)(td)
print(out["next", "observation"].squeeze(-1))
print(out["next", "terminated"].squeeze(-1))
print(out["steps_to_next_obs"].shape)
Output (on main at 648f50c):
tensor([3., 4., 4., 4., 6., 6.])
tensor([False, False, False, True, False, False])
torch.Size([6])
| t |
next obs |
next terminated |
expected terminated |
| 0 |
3 |
False |
False |
| 1 |
4 (terminal) |
False |
True |
| 2 |
4 (terminal) |
False |
True |
| 3 |
4 (terminal) |
True |
True |
| 4 |
6 |
False |
False (truncated by batch end, bootstrapping is correct) |
| 5 |
6 |
False |
False |
Expected behavior
("next", "terminated") should describe the state that ("next", "observation") now
points to, i.e. be gathered with the same index (t + steps_to_next_obs - 1).
For the example above: [False, True, True, True, False, False].
System info
- torchrl: 0.14.0+g648f50cf (built from
main)
- tensordict: 2026.10.07 (tensordict-nightly)
- torch: 2.16.0.dev20261007+cpu
- numpy: 2.5.3
- Python: 3.12.13 (MSC v.1944 64 bit, AMD64)
- OS: Windows 11 (10.0.26200)
Reason and Possible fixes
Gather terminated with the same indices as the rest of "next". A user-side workaround
that fixes it:
class FixedMultiStep(MultiStep):
def forward(self, td):
term = td["next", "terminated"].clone() # [..., T, 1]
out = super().forward(td)
steps = out["steps_to_next_obs"] # [..., T]
T = steps.shape[-1]
idx = torch.arange(T, device=steps.device) + steps - 1
out["next", "terminated"] = term.gather(-2, idx.unsqueeze(-1))
out["steps_to_next_obs"] = steps.unsqueeze(-1) # see "Related" below
return out
I only shift terminated and leave done untouched, since I assume done is excluded on
purpose to preserve trajectory boundaries for downstream bookkeeping (episode counting,
trajectory splitting). Whether done / truncated should also be gathered is up to the
maintainers.
The existing tests in test/test_postprocs.py::test_multistep check that the next
observation is replaced (or the step is nonterminal == False), but do not check the
("next", "terminated") entry consumed by the losses, which is why this isn't caught.
I have a fix with regression tests for both this and the issue below ready and will open
a PR referencing this issue.
Related
steps_to_next_obs (and gamma) are written with shape [..., T] (no trailing singleton
dim, see the last line of the output above). DQNLoss's TD0 estimator handles this via
.view_as(reward), but DistributionalDQNLoss uses it as-is:
(1 - terminated) * discount * support broadcasts [B, 1] * [B] to [B, B] and then fails
against support of shape [atoms]:
RuntimeError: The size of tensor a (256) must match the size of tensor b (51) at non-singleton dimension 1
This was previously reported in #2269 and a fix was proposed in #2270, which was closed
without being merged; the line is unchanged on main.
Checklist
Describe the bug
MultiStepreplaces("next", "observation")(and other"next"entries) with theentry
nsteps ahead, but explicitly excludes the done keys from that gather:As a result, for the
n - 1transitions preceding a termination,("next", "observation")points to the terminal observation while
("next", "terminated")is stillFalse(it is the flag of the original 1-step transition).
Losses that read
("next", "terminated")then bootstrap from the terminal state, e.g. inDistributionalDQNLoss:the target becomes
R^(k) + γ^k · Z(s_terminal)instead ofR^(k). Terminal observationsare never stored as
"observation"in the replay buffer, so the network's value at thosestates is never trained and the bias is arbitrary. This affects every episode that ends
with a termination (e.g. every CartPole episode), not just edge cases.
The
"nonterminal"key written byMultiStepdoes not cover this: it isFalseonly atthe step where
doneitself is set, not at the precedingn - 1steps.MultiStepTransformshares_multi_step_funcand is affected in the same way.To Reproduce
Output (on
mainat 648f50c):Expected behavior
("next", "terminated")should describe the state that("next", "observation")nowpoints to, i.e. be gathered with the same index (
t + steps_to_next_obs - 1).For the example above:
[False, True, True, True, False, False].System info
main)Reason and Possible fixes
Gather
terminatedwith the same indices as the rest of"next". A user-side workaroundthat fixes it:
I only shift
terminatedand leavedoneuntouched, since I assumedoneis excluded onpurpose to preserve trajectory boundaries for downstream bookkeeping (episode counting,
trajectory splitting). Whether
done/truncatedshould also be gathered is up to themaintainers.
The existing tests in
test/test_postprocs.py::test_multistepcheck that the nextobservation is replaced (or the step is
nonterminal == False), but do not check the("next", "terminated")entry consumed by the losses, which is why this isn't caught.I have a fix with regression tests for both this and the issue below ready and will open
a PR referencing this issue.
Related
steps_to_next_obs(andgamma) are written with shape[..., T](no trailing singletondim, see the last line of the output above).
DQNLoss's TD0 estimator handles this via.view_as(reward), butDistributionalDQNLossuses it as-is:(1 - terminated) * discount * supportbroadcasts[B, 1] * [B]to[B, B]and then failsagainst
supportof shape[atoms]:This was previously reported in #2269 and a fix was proposed in #2270, which was closed
without being merged; the line is unchanged on
main.Checklist