Repository navigation
[Feature Request] Question about LSTMModules #3147
Description
Activity
Wouldn't this mean that during training the recurrent state is being passed between unrelated trajectories?
You should be passing the
"is_init"key to the module: using that, it will know how to reorganize your data in such a way that trajectories are fed to the LSTM independently.rl/torchrl/modules/tensordict_module/rnn.py
Lines 705 to 721 in f7592c6
is_init = tensordict_shaped["is_init"].squeeze(-1) splits = None if self.recurrent_mode and is_init[..., 1:].any(): from torchrl.objectives.value.utils import _get_num_per_traj_init # if we have consecutive trajectories, things get a little more complicated # we have a tensordict of shape [B, T] # we will split / pad things such that we get a tensordict of shape # [N, T'] where T' <= T and N >= B is the new batch size, such that # each index of N is an independent trajectory. We'll need to keep # track of the indices though, as we want to put things back together in the end. splits = _get_num_per_traj_init(is_init) tensordict_shaped_shape = tensordict_shaped.shape tensordict_shaped = _split_and_pad_sequence( tensordict_shaped.select(*self.in_keys, strict=False), splits ) is_init = tensordict_shaped["is_init"].squeeze(-1) - linked a pull request that will close this issue[Tutorials] Upgrade DQN with RNN tutorial #3152
on Sep 8, 2025 But IIUC,
is_initonly tracks whether the a specific frame is the start of an episode when the episode is collected. So the logic withinrnn.pywouldn't be able to appropriately reshape the data to (B, T) when the consecutive trajectories are arbitrary slices.Take this example:
- Rollout trajectories of length 200 from the env.
- Add these to a buffer
- Sample short trajectories from the buffer using a slice sampler. Lets say this results in traj ids [2, 2, 2, 4, 4, 4, 5, 5, 5, 6, 6, 6].
- The
is_initvalues for these slices may all be False since its possible none of the slices start at the beginning of an episode
- The
Reacted by wlruys(cc @vmoens, incase you mised this last comment)
You could be right, I should test this.
In practice slicesampler writes down the truncated key so we should be able to get the same signal from there!@vmoens whats a good way one can go about testing this
/assign
Reacted by github-actions
Motivation
Could I get more details on tensor sizes required when using LSTMModule?
I'm particularly confused about the required input size with recurrent mode... seems like it needs to be (*b, T, Feature), and that the time dimension should not have consecutive trajectories stacked. And since recurrent_mode is enabled for all losses, doesn't this mean every time you call a loss function and happen to have a recurrent model, it needs to be batched in the same way?
I'm a bit confused how this example works at all: https://docs.pytorch.org/rl/main/tutorials/dqn_with_rnn.html. The data passed into the loss function is of size (B1, B2), where both batch dimensions have multiple slices concatentated together. Wouldn't this mean that during training the recurrent state is being passed between unrelated trajectories?
In [2]: s["collector", "traj_ids"]
Out[2]:
tensor([[10, 10, 10, 10, 10, 10, 10, 10, 11, 11, 11, 11, 11, 11, 11, 11, 11, 12,
12, 12, 12, 12, 12, 12, 12, 12, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13,
13, 13, 13, 13, 13, 13, 14, 14, 14, 14, 14, 14, 14, 14],
[30, 30, 30, 30, 30, 30, 30, 30, 30, 30, 30, 31, 31, 31, 31, 31, 31, 31,
31, 32, 32, 32, 32, 32, 32, 32, 32, 33, 33, 33, 33, 33, 33, 33, 33, 33,
33, 33, 33, 34, 34, 34, 34, 34, 34, 34, 34, 34, 34, 34],
[20, 20, 20, 20, 20, 20, 20, 21, 21, 21, 21, 21, 21, 21, 21, 21, 22, 22,
22, 22, 22, 22, 22, 22, 22, 22, 23, 23, 23, 23, 23, 23, 23, 23, 23, 23,
23, 23, 24, 24, 24, 24, 24, 24, 24, 24, 24, 24, 25, 25],
[ 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1,
1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3,
3, 3, 3, 3, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4]])
Solution
Documentation updates.
Checklist