Skip to content

[Feature Request] Question about LSTMModules #3147

Description

@itwasabhi

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

  • I have checked that there is no similar issue in the repo (required)

Activity

  1. vmoens commented on Sep 8, 2025

    @vmoens
    Contributor

    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.

    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)

  2. linked a pull request that will close this issue[Tutorials] Upgrade DQN with RNN tutorial #3152on Sep 8, 2025
  3. itwasabhi commented on Sep 8, 2025

    @itwasabhi
    ContributorAuthor

    But IIUC, is_init only tracks whether the a specific frame is the start of an episode when the episode is collected. So the logic within rnn.py wouldn'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_init values for these slices may all be False since its possible none of the slices start at the beginning of an episode
  4. itwasabhi commented on Sep 8, 2025

    @itwasabhi
    ContributorAuthor

    (cc @vmoens, incase you mised this last comment)

  5. vmoens commented on Sep 8, 2025

    @vmoens
    Contributor

    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!

  6. reopened this on Sep 8, 2025
  7. itwasabhi commented on Sep 29, 2025

    @itwasabhi
    ContributorAuthor

    @vmoens whats a good way one can go about testing this

  8. YeonwooSung commented on Sep 22, 2026

    @YeonwooSung
    Contributor

    /assign

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

Metadata

Metadata

Labels

enhancementNew feature or request

Type

No type

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions