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
246 changes: 246 additions & 0 deletions docs/source/reference/data_replaybuffers.rst
Original file line number Diff line number Diff line change
Expand Up @@ -506,6 +506,9 @@ trajectory-edge padding, while ``rebin``'s stack keeps a separate frame axis::
... ),
... ) # doctest: +SKIP

The same ``dim=-3`` convention, and the collector / ``extend`` / ``sample``
wiring, is spelled out in :ref:`catframes-collector-replay`.

**Multiple files.** A clip is often split across many small files (one per episode)
rather than one large mp4. :meth:`VideoClipRef.from_files` addresses a list of files
as a single logical sequence, so slicing, :meth:`rebin` and decoding work across
Expand All @@ -529,3 +532,246 @@ When camera and control loops run at different rates, prefer
VideoClipRef
clear_video_decoder_cache
set_video_decoder_cache_size

.. _catframes-collector-replay:

Frame stacking images with CatFrames
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~

Visual off-policy training almost always stacks the last ``N`` frames so a
CNN can see motion. :class:`~torchrl.envs.transforms.CatFrames` can do that
in two places; pick **one** of them for the tensor that is stored.

- **Env-side** (policy / inference): a stateful rolling buffer on the
:class:`~torchrl.envs.TransformedEnv`. Documented in
:ref:`CatFrames for visual RL <catframes-images>`.
- **Buffer-side** (this section): store **unstacked** raw pixels and
rebuild the stack in ``rb.sample()``. A stack of ``N`` float32 frames
occupies ``N`` times the RAM of one uint8 image; putting
:class:`~torchrl.envs.transforms.CatFrames` on the sample path is how
you avoid paying that cost in the storage.

The two placements share ``N`` and ``dim`` so the policy and the loss see
the same layout. They must not both write the stacked key: env stacking
**and** buffer stacking of the already-stacked tensor produces
``[N * N * C, H, W]``.

For CHW images the stack dimension is ``dim=-3`` (channel), not the
vector default ``dim=-1``. After :class:`~torchrl.envs.transforms.ToTensorImage`
(and optional :class:`~torchrl.envs.transforms.GrayScale`) a frame is
``[C, H, W]``; concatenating along ``-3`` yields ``[N * C, H, W]``.

Why the raw frame is what you store
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^

Keep the env transform and the stored tensor on **different keys**.
:class:`~torchrl.envs.transforms.ToTensorImage` should write a processed
copy (``out_keys=["pixels_trsf"]``) and leave the env's ``"pixels"``
untouched. :class:`~torchrl.envs.transforms.CatFrames` then stacks
``pixels_trsf`` for the policy. On ``rb.extend(...)`` you drop
``pixels_trsf`` (and its ``("next", ...)`` counterpart) so the storage
holds only the uint8 frame. The buffer transform recreates
``pixels_trsf`` from ``"pixels"`` at sample time.

:meth:`~torchrl.envs.transforms.CatFrames.make_rb_transform_and_sampler`
builds the sample-time half of that pipeline: a
:class:`~torchrl.data.replay_buffers.SliceSampler` with ``slice_len=N``,
and a transform that reshapes each sampled slice to ``[B, N]``, unfolds
:class:`~torchrl.envs.transforms.CatFrames` along time, keeps the last
step of every window, and -- on the inverse / write path -- excludes the
stacked ``out_keys``. Offline :class:`~torchrl.envs.transforms.CatFrames`
(``forward`` / ``unfolding``) needs a time dimension and
``("next", "done")`` to stop stacks at episode boundaries; the helper
sampler is what provides that time axis. Sampling independent transitions
and then unfolding treats the batch index as time and mixes unrelated
frames.

Collector, ``extend`` and ``sample``
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^

A :class:`~torchrl.collectors.Collector` writes ``("collector", "traj_ids")``
on every batch. Pass that key to the helper so the
:class:`~torchrl.data.replay_buffers.SliceSampler` agrees with the
collector on trajectory boundaries. Then ``extend`` each collected batch
and ``sample`` as usual:

.. code-block:: python

from torchrl.collectors import Collector, RandomPolicy
from torchrl.data import LazyTensorStorage, ReplayBuffer
from torchrl.envs import (
CatFrames,
Compose,
GrayScale,
GymEnv,
InitTracker,
Resize,
StepCounter,
ToTensorImage,
TransformedEnv,
)

frame_stack = 4
batch_size = 32

catframes = CatFrames(
N=frame_stack,
dim=-3,
in_keys=["pixels_trsf"],
out_keys=["pixels_trsf"],
)
env = TransformedEnv(
GymEnv("CartPole-v1", from_pixels=True, pixels_only=True),
Compose(
InitTracker(),
ToTensorImage(in_keys=["pixels"], out_keys=["pixels_trsf"]),
GrayScale(in_keys=["pixels_trsf"]),
Resize(84, 84, in_keys=["pixels_trsf"]),
catframes,
StepCounter(),
),
) # doctest: +SKIP

rb_catframes, sampler = catframes.make_rb_transform_and_sampler(
batch_size=batch_size,
traj_key=("collector", "traj_ids"),
)
# Sample-path processing must match the env: same ToTensorImage /
# GrayScale / Resize, but applied to both the root and the "next"
# pixels so CatFrames sees the layout it saw during collection.
rb = ReplayBuffer(
storage=LazyTensorStorage(100_000),
sampler=sampler,
batch_size=batch_size,
transform=Compose(
ToTensorImage(
in_keys=["pixels", ("next", "pixels")],
out_keys=["pixels_trsf", ("next", "pixels_trsf")],
),
GrayScale(in_keys=["pixels_trsf", ("next", "pixels_trsf")]),
Resize(84, 84, in_keys=["pixels_trsf", ("next", "pixels_trsf")]),
rb_catframes,
),
) # doctest: +SKIP

collector = Collector(
env,
RandomPolicy(env.action_spec),
frames_per_batch=64,
total_frames=10_000,
) # doctest: +SKIP
for data in collector:
# Inverse ExcludeTransform on rb_catframes already drops the
# stacked keys on write; excluding them here is equivalent and
# makes the stored tensordict obvious.
rb.extend(data.exclude("pixels_trsf", ("next", "pixels_trsf")))
batch = rb.sample()
# batch["pixels_trsf"] has shape [32, 4, 84, 84] (grayscale, dim=-3)

``rb.extend(data)`` is the only write API you need: the collector yields
a tensordict of shape ``[frames_per_batch]`` (or
``[batch, time]`` for a batched env -- then give the storage
``ndim=2``, see :ref:`collectors and replay buffers <ref_collectors>`).
Do not flatten away ``("collector", "traj_ids")`` or ``("next", "done")``
before extending; the sampler uses them to keep each slice inside one
episode.

The same ``extend`` / ``sample`` pairing without the helper looks like
this. Three extra steps are not optional:

- ``SliceSampler(slice_len=N)``: without it,
:meth:`~torchrl.envs.transforms.CatFrames.unfolding` has no time
dimension.
- ``reshape(-1, N)``: ``SliceSampler`` returns the ``B`` windows
flattened as ``[B * N]``. Leaving that rank-1 batch as the time
axis makes later windows inherit the last frames of the previous
window.
- ``[:, -1]``: after unfolding, each window is an ``N``-step
tensordict; the last step is the completed stack.

.. code-block:: python

from torchrl.data import SliceSampler
from torchrl.envs import ExcludeTransform

rb = ReplayBuffer(
storage=LazyTensorStorage(100_000),
sampler=SliceSampler(
slice_len=frame_stack,
traj_key=("collector", "traj_ids"),
),
batch_size=batch_size * frame_stack, # B windows of length N

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

[P2] Reshape sampled windows before CatFrames and retain their last step

SliceSampler returns the B windows flattened as [B*N]; it does not make a [B, N] TensorDict for this transform. Consequently CatFrames treats the entire sample as one time sequence, so the first N-1 rows of later windows incorporate frames from an unrelated preceding window. A synthetic increasing-frame replay reproduces stacks such as [52, 53, 54, 66]; 12 of 16 rows were wrong with N=4. Match the helper by reshaping to [-1, N] before CatFrames and selecting [:, -1] afterward, or remove this purported equivalent recipe.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed in f5ef733. The explicit recipe now matches make_rb_transform_and_sampler: reshape(-1, N) before CatFrames and [:, -1] after it. A synthetic increasing-frame buffer reproduced the bleed ([54, 55, 56, 61] and only 16/64 consecutive rows with N=4); after the reshape/last-step pair every window is consecutive and the sample shape is [B, N] like the helper. The surrounding prose and the pitfalls list call out that SliceSampler returns [B * N], not [B, N].

transform=Compose(
ToTensorImage(
in_keys=["pixels", ("next", "pixels")],
out_keys=["pixels_trsf", ("next", "pixels_trsf")],
),
GrayScale(in_keys=["pixels_trsf", ("next", "pixels_trsf")]),
Resize(84, 84, in_keys=["pixels_trsf", ("next", "pixels_trsf")]),
# SliceSampler returns [B * N]; CatFrames needs a time axis
# per window, then only the completed stack is kept.
lambda td: td.reshape(-1, frame_stack),
CatFrames(
N=frame_stack,
dim=-3,
in_keys=["pixels_trsf", ("next", "pixels_trsf")],
out_keys=["pixels_trsf", ("next", "pixels_trsf")],
),
lambda td: td[:, -1],
ExcludeTransform("pixels_trsf", ("next", "pixels_trsf"), inverse=True),
),
) # doctest: +SKIP
rb.extend(data) # stores raw "pixels" only
batch = rb.sample() # [32, 4, 84, 84] grayscale stacks

The helper is the preferred form: it multiplies the requested
``batch_size`` by ``N`` internally and inserts the reshape / last-step
selection, so ``rb.sample()`` returns ``batch_size`` stacked
transitions rather than a flat ``batch_size * N`` window.

Common pitfalls
^^^^^^^^^^^^^^^

- **Stacking twice.** Env :class:`~torchrl.envs.transforms.CatFrames`
writing ``pixels_trsf`` **and** a buffer
:class:`~torchrl.envs.transforms.CatFrames` that reads that same
already-stacked key. Store raw ``"pixels"`` and rebuild, or store the
stack and do not attach :class:`~torchrl.envs.transforms.CatFrames` to
the buffer -- not both.
- **Wrong ``dim``.** Images after
:class:`~torchrl.envs.transforms.ToTensorImage` are CHW: use
``dim=-3``. ``dim=-1`` is for vector observations. ``dim=-4`` is only
correct after an
:class:`~torchrl.envs.transforms.UnsqueezeTransform` that inserted that
axis (the ``[N, C, H, W]`` variant in
``examples/replay-buffers/catframes-in-buffer.py``).
- **No time axis.** Offline :class:`~torchrl.envs.transforms.CatFrames`
unfolds along time. A uniform random sample of transitions has no
time axis, so the stack is assembled from unrelated steps. Always pair
the buffer transform with a
:class:`~torchrl.data.replay_buffers.SliceSampler` (or
:meth:`~torchrl.envs.transforms.CatFrames.make_rb_transform_and_sampler`).
- **Flattened slices.** ``SliceSampler`` returns ``[B * N]``, not
``[B, N]``. :class:`~torchrl.envs.transforms.CatFrames` then treats
the whole sample as one sequence and the first ``N - 1`` rows of each
later window mix frames from the previous window. Reshape to
``[-1, N]`` before the transform and keep ``[:, -1]`` after it
(the helper does both).
- **Missing ``("next", ...)`` keys.** A buffer transform does not walk
into ``"next"`` on its own. List both ``"pixels"`` and
``("next", "pixels")`` on every sample-path transform that should
apply to both.
- **Same in/out key on**
:class:`~torchrl.envs.transforms.ToTensorImage`. If the processed
tensor overwrites ``"pixels"``, there is no cheap raw frame left to
store and you cannot drop the stack on ``extend``.

.. seealso::

:ref:`CatFrames for visual RL <catframes-images>` for the env-side
placement, reset / :class:`~torchrl.envs.transforms.InitTracker`
behaviour, and the ``dim=-3`` convention.
:ref:`Collectors and replay buffers <ref_collectors>` for
``ndim`` / ``traj_key`` when the collector is batched or
multi-process. A runnable script of the extra-axis (``dim=-4``)
variant is ``examples/replay-buffers/catframes-in-buffer.py``.
102 changes: 101 additions & 1 deletion docs/source/reference/envs_transforms.rst
Original file line number Diff line number Diff line change
Expand Up @@ -192,7 +192,107 @@ transform has an inverse transform that subtracts 1 from the action tensor:
Using a Transform with a Replay Buffer
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~

You can use a transform with a replay buffer by passing it to the ReplayBuffer constructor:
A transform passed to :class:`~torchrl.data.ReplayBuffer` (the ``transform=``
argument, or :meth:`~torchrl.data.ReplayBuffer.append_transform`) runs on the
**sample** path: ``rb.sample()`` applies it after the storage is indexed.
If the transform implements an inverse, that inverse runs when data is
**written** (``rb.add`` / ``rb.extend``). This is how you store cheap raw
observations and reconstruct a processed view only at training time.

The most common instance of this pattern is frame stacking. The env-side
recipe is below; the collector + replay-buffer recipe (including
``extend`` / ``sample``) lives in
:ref:`Frame stacking images with CatFrames <catframes-collector-replay>`.

.. _catframes-images:

CatFrames for visual RL
~~~~~~~~~~~~~~~~~~~~~~~

:class:`CatFrames` concatenates the last ``N`` observations along one
existing dimension so a feed-forward policy can see motion. There are two
legitimate placements; they solve different problems and must not be
stacked on top of each other.

1. **On the env** (this section). The policy sees a stacked observation
at every step. The transform is stateful: it keeps a rolling buffer
that is flushed on reset.
2. **On the replay buffer, at sample time**
(:ref:`catframes-collector-replay`). Raw unstacked frames are stored;
the stack is rebuilt when you call ``rb.sample()``. This is the
memory-efficient path for image off-policy training.

For CHW images (the layout produced by :class:`ToTensorImage`) the stack
dimension is ``dim=-3`` (the channel axis). That is the DQN convention:
a grayscale frame of shape ``[1, H, W]`` becomes ``[N, H, W]``; an RGB
frame of shape ``[3, H, W]`` becomes ``[3 * N, H, W]``. Vector
observations use ``dim=-1`` instead. Using the vector default on pixels,
or stacking along ``-4`` without first inserting that axis, silently
produces the wrong layout.

Env-side stacking (policy input)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^

Put :class:`CatFrames` on the env so the observation spec the policy
reads is already stacked. Write the stack to a **different** key than
the raw pixels: the collector can then persist the cheap uint8 frame
and drop the float stack (see :ref:`catframes-collector-replay`).

:class:`CatFrames` is stateful. ``env.reset()`` (or a ``"_reset"`` flag
in the tensordict) flushes the rolling buffer and pads the new stack.
With the default ``padding="same"`` the missing history is a repeat of
the first post-reset frame; ``padding="constant"`` fills with
``padding_value`` (0 by default). :class:`InitTracker` is not required
for that flush -- :class:`CatFrames` listens to the env ``_reset`` key
-- but it should sit in the same :class:`Compose` so ``"is_init"`` marks
the steps where the stack was re-initialized. Collectors, recurrent
policies and advantage estimators all read that flag.

.. code-block:: python

from torchrl.envs import (
CatFrames,
Compose,
GrayScale,
GymEnv,
InitTracker,
Resize,
StepCounter,
ToTensorImage,
TransformedEnv,
)

env = TransformedEnv(
GymEnv("CartPole-v1", from_pixels=True, pixels_only=True),
Compose(
InitTracker(),
ToTensorImage(in_keys=["pixels"], out_keys=["pixels_trsf"]),
GrayScale(in_keys=["pixels_trsf"]),
Resize(84, 84, in_keys=["pixels_trsf"]),
CatFrames(N=4, dim=-3, in_keys=["pixels_trsf"], out_keys=["pixels_trsf"]),
StepCounter(),
),
) # doctest: +SKIP
# "pixels": raw uint8 frame from the env
# "pixels_trsf": float stack of shape [4, 84, 84] (grayscale, dim=-3)
# "is_init": True on the first step after every reset

A rollout or a :class:`~torchrl.collectors.Collector` then feeds
``pixels_trsf`` to the policy. Do **not** also attach a second
:class:`CatFrames` on the replay buffer that reads the same stacked
key -- that stacks twice (``[N, H, W]`` stored, ``[N * N, H, W]``
sampled). Either:

- store the already-stacked ``pixels_trsf`` and leave the buffer
transform-free (``N`` times more RAM), or
- exclude ``pixels_trsf`` on ``extend`` and rebuild the stack from
raw ``"pixels"`` at sample time (recommended; see
:ref:`catframes-collector-replay`).

If you want a separate stack axis ``[N, C, H, W]`` rather than channel
concatenation, unsqueeze first and pass ``dim=-4`` (that is the variant
in ``examples/replay-buffers/catframes-in-buffer.py``). For a standard
visual-RL CNN the ``dim=-3`` layout above is the one to use.

Cloning transforms
~~~~~~~~~~~~~~~~~~
Expand Down
5 changes: 5 additions & 0 deletions torchrl/envs/transforms/_observation.py
Original file line number Diff line number Diff line change
Expand Up @@ -1089,6 +1089,11 @@ def make_rb_transform_and_sampler(
.. note:: For a more complete example, refer to torchrl's github repo `examples` folder:
https://github.com/pytorch/rl/tree/main/examples/replay-buffers/catframes-in-buffer.py

.. seealso:: :ref:`catframes-images` (env-side stacking, reset /
:class:`~torchrl.envs.transforms.InitTracker`) and
:ref:`catframes-collector-replay` (collector ``extend`` / buffer
``sample``, ``dim=-3`` for CHW images).

"""
from torchrl.data.replay_buffers import SliceSampler

Expand Down
Loading