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
96 changes: 87 additions & 9 deletions docs/source/reference/recurrent_state_lifecycle.rst
Original file line number Diff line number Diff line change
Expand Up @@ -16,10 +16,12 @@ primer required by :class:`~torchrl.modules.LSTMModule` or
The main rule to keep in mind is simple: if the loss should replay
sequences, sample sequences. For replay-buffer training, use
:class:`~torchrl.data.replay_buffers.SliceSampler` or another trajectory-aware sampler so
the loss receives contiguous time chunks with ``is_init`` boundaries
preserved. The rest of this page explains what the automated path wires up,
and what to check when building a custom loop, custom replay transform, or
manually constructed training batch.
the loss receives contiguous time chunks. ``is_init`` must then be ``True``
at every slice start — not only at episode starts — or hidden state leaks
across concatenated traj ids (see :ref:`ref_slice_is_init` below). The rest
of this page explains what the automated path wires up, and what to check
when building a custom loop, custom replay transform, or manually
constructed training batch.

Minimal recurrent PPO wiring
----------------------------
Expand Down Expand Up @@ -192,6 +194,12 @@ The path at a glance
Replay buffer (stores (B, T, ...) trajectories with is_init preserved)
│
▼
SliceSampler (default init_key="is_init")
│
├─ samples mid-episode windows (stored is_init often all-false)
└─ ORs True at every slice start so LSTMModule can split
│
▼
Loss / GAE (recurrent mode)
with set_recurrent_mode(True):
value_net(sampled_batch)
Expand All @@ -210,7 +218,8 @@ What ``is_init`` means
set by :class:`~torchrl.envs.transforms.InitTracker` to ``True`` on the *first* step
of every trajectory and ``False`` everywhere else. A trajectory begins
at an explicit :meth:`~torchrl.envs.EnvBase.reset` or right after a
``done`` from the previous step.
``done`` from the previous step. After slice sampling, the first timestep
of each slice must also be ``True`` (see :ref:`ref_slice_is_init`).

If you do not append :class:`~torchrl.envs.transforms.InitTracker` to your env,
``is_init`` will be absent and :class:`~torchrl.modules.LSTMModule` will
Expand All @@ -219,7 +228,9 @@ custom replay buffer / transform drops or rewrites the true boundary
signal), the LSTM has no way to know when a new trajectory has started.
In that case the hidden state will silently carry forward across episode
boundaries — usually the most painful class of recurrent bug to diagnose
because rewards still look plausible. (:ref:`Trajectory boundaries
because rewards still look plausible. The same silent failure happens after
slice sampling when mid-episode windows keep their stored all-false
``is_init`` (see :ref:`ref_slice_is_init`). (:ref:`Trajectory boundaries
<ref_traj_boundaries>` documents which markers samplers use to recover
episode boundaries and the cases where a boundary is unrecoverable.)

Expand Down Expand Up @@ -280,6 +291,66 @@ starts at ``t*+1``. The corresponding ``is_init`` slot is true.
``is_init``, none of these paths sees the boundary and the LSTM treats
the post-done timesteps as a continuation of the pre-done trajectory.

.. _ref_slice_is_init:

Slice sampling and ``is_init``
------------------------------

:class:`~torchrl.envs.transforms.InitTracker` writes ``is_init=True`` only
at *episode* starts (reset, or the step after ``done``). A
:class:`~torchrl.data.replay_buffers.SliceSampler` window can start in the
middle of an episode, in which case every stored ``is_init`` flag in that
window is ``False``. Concatenating such slices — the default flat sampler
output, or a ``[B, T]`` batch whose time dim packs several traj ids —
then looks like one long trajectory to
:class:`~torchrl.modules.LSTMModule`. Hidden state from the last step of
slice A is fed into the first step of unrelated slice B.

The contract is: **after slice sampling, mark the first timestep of each
slice as** ``is_init=True``. :class:`~torchrl.data.replay_buffers.SliceSampler`
does this by default (``init_key="is_init"``): it ORs a ``True`` at every
slice start onto the flags stored in the buffer, so episode-internal
resets that fall inside a slice are preserved. Pass ``init_key=None`` only
when the consumer must not treat slice starts as resets (DreamerV3 RSSM).

If you build the training batch yourself, or you sampled with
``init_key=None``, apply the same OR on the user side. For a
``[num_slices, T]`` batch::

sampled["is_init"][:, 0] = True

For a flat concatenated sample, the sampler also writes
``("next", "truncated")`` at each slice end. Shifting that marker onto
``is_init`` marks the next slice start::

is_init = sampled["is_init"].clone()
truncated = sampled["next", "truncated"]
is_init[0] = True
is_init[1:] |= truncated[:-1]
sampled["is_init"] = is_init

Do not OR ``truncated`` into ``is_init`` in place without the shift:
``truncated`` sits at slice *ends*, ``is_init`` at slice *starts*.

:class:`~torchrl.modules.LSTMModule` then uses ``is_init`` in two ways:

- **Sequential mode** (collection): zeros the incoming hidden so a fresh
trajectory does not inherit the previous one's state::

is_init_expand = expand_as_right(is_init, hidden0)
hidden0 = torch.where(is_init_expand, zeros, hidden0)
hidden1 = torch.where(is_init_expand, zeros, hidden1)

- **Recurrent mode** (loss / GAE): the module splits the time dimension
wherever ``is_init[..., 1:]`` is set and, at each split, restarts from
the *stored* hidden at that position (the value collected with the
slice). Mid-episode slices therefore keep their stored recurrent state
instead of being zeroed.

Without the slice-start OR, a present-but-all-false ``is_init`` is the
silent failure mode: no ``KeyError``, plausible rewards, hidden state
leaking across unrelated traj ids.

Hidden outputs and recurrent backends
-------------------------------------

Expand Down Expand Up @@ -329,7 +400,9 @@ Common debugging symptoms
sequence-aware sampler, verify that the recurrent loss path is wrapped
in ``with set_recurrent_mode(True):``, and check that ``is_init`` is
preserved through your replay buffer (some transforms drop unknown
keys).
keys). If the sampler windows start mid-episode, also check that slice
starts are ``is_init=True`` (SliceSampler's default ``init_key``); an
all-false ``is_init`` on concatenated slices is this same leak.

**Symptom: shapes mismatch in** ``LSTMModule._lstm`` **with cryptic transpose errors.**
The module expects the tensordict-native hidden layout
Expand Down Expand Up @@ -363,8 +436,13 @@ What to check, in order
primer's keys (use :meth:`~torchrl.modules.LSTMModule.make_tensordict_primer`).
5. Replay-buffer training uses :class:`~torchrl.data.replay_buffers.SliceSampler` or
another trajectory-aware sampler when the loss consumes sequences.
6. Loss / advantage code runs under ``with set_recurrent_mode(True):``.
7. The replay buffer preserves ``is_init`` (and any custom recurrent
6. After sampling, slice starts are ``is_init=True`` (SliceSampler's
default ``init_key="is_init"``). A mid-episode slice with all-false
``is_init`` will leak hidden state across concatenated traj ids; OR
the first timestep of each slice, or shift ``("next", "truncated")``
onto ``is_init`` (see :ref:`ref_slice_is_init`).
7. Loss / advantage code runs under ``with set_recurrent_mode(True):``.
8. The replay buffer preserves ``is_init`` (and any custom recurrent
keys) through its transforms.

See also
Expand Down
109 changes: 109 additions & 0 deletions test/modules/test_rnn.py
Original file line number Diff line number Diff line change
Expand Up @@ -2340,6 +2340,115 @@ def test_lstm_collector_replay_mid_batch_done_resets_hidden_state(self):
msg="hidden state leaked across is_init trajectory boundary",
)

def test_lstm_concatenated_slices_all_false_is_init_leaks_hidden(self):
"""Mid-episode concatenated slices leak hidden state unless is_init is set at slice starts.

SliceSampler can return mid-trajectory windows whose stored is_init is
all False (none of the frames were episode starts). LSTMModule in
recurrent mode then treats the concatenated time dim as one trajectory
and carries hidden state from the last step of slice A into the first
step of unrelated slice B.

Marking slice starts with is_init=True (SliceSampler's init_key
default, or shifting ("next", "truncated") onto is_init) restores
independence. Expected features come from the inner LSTM on each
slice alone, not from LSTMModule.
"""
torch.manual_seed(0)
num_rows = 2
slice_len = 3
n_slices = 2
T = slice_len * n_slices
F, H, L = 4, 5, 1
# Mixed traj ids along time, as in https://github.com/pytorch/rl/issues/3147
traj_ids = torch.tensor(
[
[2, 2, 2, 4, 4, 4],
[5, 5, 5, 6, 6, 6],
],
dtype=torch.long,
)
lstm_module = LSTMModule(
input_size=F,
hidden_size=H,
num_layers=L,
in_keys=["obs", "rs_h", "rs_c"],
out_keys=["feat", ("next", "rs_h"), ("next", "rs_c")],
python_based=True,
dropout=0,
)
lstm_module.eval()

obs = torch.randn(num_rows, T, F)
rs_h = torch.randn(num_rows, T, L, H)
rs_c = torch.randn(num_rows, T, L, H)
is_init = torch.zeros(num_rows, T, 1, dtype=torch.bool)
truncated = torch.zeros(num_rows, T, 1, dtype=torch.bool)
truncated[:, slice_len - 1] = True
truncated[:, T - 1] = True
data = TensorDict(
{
"obs": obs,
"rs_h": rs_h,
"rs_c": rs_c,
"is_init": is_init,
("collector", "traj_ids"): traj_ids,
("next", "truncated"): truncated,
},
[num_rows, T],
)

def inner_features(obs_slice, h_slice, c_slice):
h0 = h_slice[:, 0].transpose(0, 1).contiguous()
c0 = c_slice[:, 0].transpose(0, 1).contiguous()
y, _ = lstm_module.lstm(obs_slice, (h0, c0))
return y

expected = torch.empty(num_rows, T, H)
leaked_expected = torch.empty(num_rows, T, H)
for row in range(num_rows):
leaked_expected[row : row + 1] = inner_features(
obs[row : row + 1], rs_h[row : row + 1], rs_c[row : row + 1]
)
for s in range(n_slices):
sl = slice(s * slice_len, (s + 1) * slice_len)
expected[row : row + 1, sl] = inner_features(
obs[row : row + 1, sl],
rs_h[row : row + 1, sl],
rs_c[row : row + 1, sl],
)

with set_recurrent_mode(True), torch.no_grad():
leaked = lstm_module(data.clone())

torch.testing.assert_close(leaked["feat"], leaked_expected)
torch.testing.assert_close(
leaked["feat"][:, :slice_len], expected[:, :slice_len]
)
assert not torch.allclose(
leaked["feat"][:, slice_len:],
expected[:, slice_len:],
atol=1e-5,
rtol=1e-5,
), "second-slice features matched the isolated run; leak was not observed"

marked = data.clone()
marked["is_init"][:, 0] = True
marked["is_init"][:, slice_len] = True

shifted = data.clone()
shifted_is_init = shifted["is_init"].clone()
trunc = shifted["next", "truncated"]
shifted_is_init[:, 0] = True
shifted_is_init[:, 1:] |= trunc[:, :-1]
shifted["is_init"] = shifted_is_init
assert torch.equal(shifted["is_init"], marked["is_init"])

with set_recurrent_mode(True), torch.no_grad():
isolated = lstm_module(marked)

torch.testing.assert_close(isolated["feat"], expected)


class TestGRUModule:
def test_errs(self):
Expand Down
8 changes: 8 additions & 0 deletions torchrl/modules/tensordict_module/rnn.py
Original file line number Diff line number Diff line change
Expand Up @@ -1331,6 +1331,14 @@ def forward(self, tensordict: TensorDictBase):
``is_init`` is sourced from :class:`~torchrl.envs.InitTracker` on the
env side; without that transform there is no signal for boundary
resets and hidden state will silently leak across episodes.
After :class:`~torchrl.data.replay_buffers.SliceSampler`, the first
timestep of each slice must also be ``is_init=True`` (the sampler
does this by default via ``init_key="is_init"``). Mid-episode
slices otherwise keep an all-false ``is_init`` and hidden state
leaks across concatenated traj ids. Sequential mode zeros the
incoming hidden with ``torch.where(is_init, zeros, hidden)``;
recurrent mode splits on ``is_init`` and restarts from the stored
hidden at each split.
"""
# we want to get an error if the value input is missing, but not the hidden states
defaults = [NO_DEFAULT, None, None]
Expand Down
Loading