Skip to content

[BUG] MultiStep does not shift ("next", "terminated") with the next observation, so losses bootstrap from terminal states #4536

Description

@romanbranovets

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

  • I have checked that there is no similar issue in the repo (required)
  • I have read the documentation (required)
  • I have provided a minimal working example to reproduce the bug (required)

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    bugSomething isn't working

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions