Skip to content

[Feature] Read storages and replay buffers through torch.utils.data - #4387

Merged
theap06 merged 8 commits into
pytorch:mainfrom
theap06:rb-dataloader
Sep 28, 2026
Merged

theap06 merged 8 commits into
pytorch:mainfrom
theap06:rb-dataloader

Conversation

@theap06

@theap06 theap06 commented Sep 16, 2026 •

Copy link
Copy Markdown
Collaborator

Description

First step of the DataLoader direction discussed for the next release: interoperate with torch.utils.data at the storage / buffer seam instead of adding a loader of our own. Both entry points are opt-in adapters; Storage and ReplayBuffer keep their class hierarchy.

  • Storage.as_dataset() returns a StorageDataset, a map-style torch.utils.data.Dataset with __getitems__, so a DataLoader with any torch sampler fetches each index batch through a single get call. Usage: DataLoader(rb.storage.as_dataset(), batch_size=32, shuffle=True, collate_fn=tensordict_collate). Multi-dimensional storages are read through storage.flatten().as_dataset(), and StorageEnsemble.as_dataset() raises. DataLoader(storage) without the adapter is unchanged: items are read one by one and the list goes to the user's collate function.
  • tensordict_collate returns a fetched batch unchanged (tensor, tensordict or tuple of them) and stacks lists of samples (lazily for ragged tensordicts, element-wise for mappings and tuples), so per-item storages such as ListStorage and compositions such as ConcatDataset work through the same collate. The default torch collation iterates a TensorDict over its batch dimension and crashes.
  • ReplayBuffer.as_dataset(num_batches=None) returns a ReplayBufferDataset (IterableDataset) that iterates the buffer; it requires the buffer batch_size. Each DataLoader worker holds its own copy of the buffer, so the TorchRL sampler and the buffer transforms run in the worker. num_batches is split between workers; prefetched batches are never serialized to workers and the worker copy gets fresh locks with prefetching disabled (the DataLoader prefetches); a buffer built with a generator is reseeded once per worker from the worker seed, so seeded DataLoaders are reproducible and workers draw distinct batches. RayReplayBuffer rejects as_dataset explicitly.
  • Workers read storage content live under every start method: as_dataset() moves a CPU TensorStorage to shared memory (spawn pickling already did this), memory-mapped storages are read through their files. Without this, forked workers shared the storage length but read a copy-on-write snapshot of the rows and returned stale data after parent writes. A worker forked with a private copy allocated after the dataset was created raises instead. Reads are not synchronized with writes, which the docs state. StorageDataset pickles only the storage, not the buffers attached to it, so unpicklable buffer transforms do not block spawn workers.
  • Sampler.requires_shared_state (new attribute, True by default, False for RandomSampler and SliceSampler) marks samplers whose state every consumer must observe. as_dataset rejects them when num_workers > 0 instead of silently duplicating their state per process (at serialization time for spawn, in the worker for fork). A RateLimitedReplayBuffer that is not shared is rejected the same way, since each worker would spend its own copy of the sample budget.
  • BugFix: TensorStorage.__getstate__ called tree_map(storage, fn) with the arguments swapped, so a plain-tensor TensorStorage could not be sent to a spawned process. Regression test added.
  • BugFix: ReplayBuffer.__getstate__ pickled the generator state as a temporary tensor that, under spawn with the file-descriptor sharing strategy, was freed before the child unpickled it. The state is now pickled as bytes. Regression test added.
  • BugFix: a buffer pickled for a spawned process now drops its prefetch queue. The storage pickles its attached buffer, and cloning the queue there produced temporaries whose shared-memory descriptors were recycled for the spawn pipes (bad value(s) in fds_to_keep, the red tests-cpu (3.10, bulk) job). Regression test added.

Docs: new sections in data_replaybuffers.rst and data_storage.rst. Benchmarks in benchmarks/test_replaybuffer_benchmark.py: test_replay_buffer_dataset_workers (sample path with a per-frame resize, 0/2/4 workers) and test_storage_dataset_fetch (batched adapter fetch vs per-item DataLoader(storage), about 18x on a 256 x 72 float batch, 68x at steady state).

Out of scope, planned as follow-ups: DatasetStorage (map-style datasets backing a read-only buffer), Hugging Face Hub export, the CartPole behavior-cloning example, and DataLoader ergonomics (collate_fn / batch_size=None defaults).

Storage is now a torch.utils.data.Dataset with batched fetching, and
ReplayBuffer.as_dataset() wraps a buffer in an IterableDataset so a
DataLoader owns the parallelism of the sample path. Workers hold their own
buffer copy: the sampler and transforms run in the worker, num_batches is
split between workers, generator-seeded buffers are reseeded per worker and
samplers with cross-process state are rejected. tensordict_collate keeps
tensordicts intact through torch collation. TorchRLBufferDataset builds on
the new dataset. Fix TensorStorage pickling of plain-tensor storages for
spawned processes (tree_map argument order).

Adds a worker benchmark, docs sections, and a CartPole behavior cloning
example whose checkpoint rlrender plays back.
@pytorch-bot

pytorch-bot Bot commented Sep 16, 2026 •

Copy link
Copy Markdown

🔗 Helpful Links

🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/rl/4387

Note: Links to docs will display an error until the docs builds have been completed.

❌ 3 New Failures

As of commit 462bcc0 with merge base b8b6f0c (image):

NEW FAILURES - The following jobs have failed:

This comment was automatically generated by Dr. CI and updates every 15 minutes.

@meta-cla meta-cla Bot added the CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. label Sep 16, 2026
@github-actions github-actions Bot added Documentation Improvements or additions to documentation Benchmarks rl/benchmark changes Examples llm/ LLM-related PR, triggers LLM CI tests ReplayBuffers Modules Integrations/torch_geometric Integrations Feature New feature labels Sep 16, 2026

@vmoens vmoens left a comment

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.

Requesting changes for worker-state correctness and PyTorch Dataset interoperability. The focused interop and storage-spawn tests pass locally, but the current head also has five failing CI jobs (two example jobs and three CPU jobs).

if worker is None:
return None
replay_buffer = self.replay_buffer
self._check_sampler(replay_buffer.sampler)

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.

[P1] Checking only sampler.requires_shared_state misses the existing OpenX streaming path. _StreamingSampler inherits the default False, while _StreamingStorage.get() consumes a dataset_iter created before workers are spawned. Each worker therefore receives the same cursor/state and can silently duplicate stream data. Please either mark this combination as requiring shared state or recreate and shard the stream per worker.

return None
replay_buffer = self.replay_buffer
self._check_sampler(replay_buffer.sampler)
replay_buffer._prefetch = False

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] This disables prefetching only after the worker copy has been created. With spawn, ReplayBuffer.__getstate__ has already captured and serialized the full prefetch queue into every worker; a pickle round-trip with prefetch=3 restores all three queued batches before this line clears them. For large batches this multiplies startup memory and file-descriptor traffic. Please strip prefetch state during dataset serialization instead.

from torchrl.data.replay_buffers.storages.utils import _get_default_collate

storage = self.flatten()
return _get_default_collate(storage)(storage.get(index))

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] PyTorch's Dataset.__getitems__ contract returns a list of samples, but this returns an already-collated batch and relies on an identity collate_fn. Standard composition consequently breaks: wrapping storages in ConcatDataset falls back to scalar __getitem__ calls and tensordict_collate leaves the resulting list uncollated. Third-party Storage subclasses also fail here when _get_default_collate does not recognize them. Please preserve the Dataset batching contract or make the public collation path handle both scalar-sample lists and vectorized batches.

@@ -0,0 +1,168 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.

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.

[P1] This new example is not registered in .github/unittest/examples/scripts/test_examples.py, so both example CI shards fail with unclassified examples. Please add an ExampleSpec with smoke arguments or a documented exclusion.

Strip prefetched batches when a ReplayBufferDataset is serialized, give the
worker copy fresh locks with prefetching disabled, and fail fast in the
parent for shared-state samplers and buffers without a batch size. Make
tensordict_collate stack lists of samples while passing tensor, tensordict
and tuple batches through, so per-item storages and ConcatDataset
compositions work. Read multi-dimensional storages through flatten(),
reject StorageEnsemble and RayReplayBuffer explicitly, mark the OpenX
streaming sampler and the recency prompt-group strategy as requiring shared
state, and warn at construction in TorchRLBufferDataset ahead of v0.17.
Register the behavior cloning example in the examples CI manifest.

Pickle the replay buffer generator state as bytes: the temporary tensor
used before was freed right after pickling, so with the file-descriptor
sharing strategy a spawned child received a recycled descriptor and failed
to unpickle the buffer.
@github-actions github-actions Bot added CI Has to do with CI setup (e.g. wheels & builds, tests...) Data Data-related PR, will launch data-related jobs Data/openx labels Sep 16, 2026

@vmoens vmoens left a comment

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.

Thanks that looks like it's going in the right direction
Can you hold on a bit before merging this I'd like to give it some thought and run it through a couple of former colleagues!

@theap06
theap06 marked this pull request as draft September 16, 2026 20:54
… flag

Move the CartPole behavior cloning example and its CI registration to a
follow-up, leave TorchRLBufferDataset untouched, and default
Sampler.requires_shared_state to True so only samplers with stateless draws
declare themselves. Shorten the storage docs to a pointer.
A worker forked while an asynchronous update is pending inherits the
parent's dependency future and waits on it forever before its first sample.
Mirror the spawn path by clearing the pending futures and executor in the
worker copy. Drop two redundant copies in the collate and RNG restore.

@vmoens vmoens left a comment

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.

Thanks for addressing the original inline comments. After thinking more about the API direction, I think the storage side of this interop should use an explicit adapter rather than changing the public base class.

For this PR, please keep Storage's inheritance unchanged and expose Storage.as_dataset() returning a small map-style adapter (for example, StorageDataset) that owns the PyTorch Dataset and __getitems__ behavior. DataLoader does not require Dataset inheritance: Storage already has __len__ and __getitem__. The current Storage(Dataset) change has several compatibility costs:

  • it changes the MRO and isinstance/issubclass results for a public extension point, and can break downstream multiple-inheritance classes;
  • it inherits adjacent behavior such as Dataset.__add__, so storage addition gains new semantics unrelated to replay storage;
  • more importantly, adding Storage.__getitems__ silently changes existing DataLoader(storage, ...) calls from N scalar reads followed by a list passed to the user's collator into one batched get(indices) result. Existing custom collators may therefore break even if they worked before this PR.

An adapter gives us a clean opt-in and BC boundary while retaining the optimized batched fetch. ReplayBuffer.as_dataset() / ReplayBufferDataset is likewise the right shape for the buffer side, since that is explicitly an iterable view of replay-buffer sampling rather than a new identity for ReplayBuffer itself.

Two useful follow-ups, explicitly out of scope for this PR:

  1. The inverse adapter, DatasetStorage(dataset), could let a fixed-length map-style PyTorch dataset back a read-only replay buffer together with the existing ImmutableDatasetWriter. The initial contract should probably be deliberately narrow: scalar and batched index forwarding, uniform and without-replacement sampling, explicit collation and checkpoint semantics, and no IterableDataset or trajectory-aware samplers until those contracts are designed. Stochastic __getitem__ also needs documentation because a priority would apply to an index, not necessarily to a stable realized sample.
  2. Utilities to export and upload offline replay data to the Hugging Face Hub would be valuable. I would make an interoperable dataset representation that can be consumed outside TorchRL the primary exchange format (including LeRobot-oriented conversion where the schema fits), while optionally supporting a native TorchRL storage artifact for lossless, fast reload. The portable schema, trajectory metadata, transforms, and native-storage versioning deserve a separate design rather than being coupled to this DataLoader adapter.

Separately, the current head still has a red tests-cpu (3.10, bulk) job in test_as_dataset_workers_drop_prefetched_batches (ValueError: bad value(s) in fds_to_keep), which also needs to be resolved or shown to be unrelated before approval.

@theap06
theap06 marked this pull request as ready for review September 27, 2026 07:36
Keep Storage a plain class and expose Storage.as_dataset(), returning a
map-style StorageDataset that owns the torch Dataset behavior and the
batched __getitems__ fetch. DataLoader(storage) keeps reading items one by
one, so existing collate functions see the same lists as before.
StorageEnsemble.as_dataset() raises.

DataLoader workers read storage content live: as_dataset() moves a CPU
TensorStorage to shared memory, so forked workers no longer pair the shared
length with a copy-on-write snapshot of the rows and return stale data. A
worker forked with a private copy allocated after the dataset was created
raises instead. StorageDataset pickles only the storage, not the buffers
attached to it, so unpicklable buffer transforms do not block spawn workers.

A buffer pickled for a spawned process drops its prefetch queue. The storage
pickles its attached buffer, and cloning the queue there produced
temporaries whose shared-memory descriptors were recycled for the spawn
pipes ("bad value(s) in fds_to_keep").

Tests are consolidated around these behaviors and a benchmark compares the
batched adapter fetch with per-item reads.
# Conflicts:
#	benchmarks/test_replaybuffer_benchmark.py
#	test/rb/test_rb_core.py
#	test/rb/test_storages.py
RateLimitedReplayBuffer, merged from main, keeps its sample budget in the
buffer. Worker copies of an unshared buffer each spend their own budget, so
two workers drew 200 records against a budget of 100. ReplayBufferDataset
now rejects such a buffer when workers are used, like samplers that require
shared state; a shared buffer keeps a single budget across workers.

Forked workers also get a fresh readiness condition when the buffer is not
shared, since the parent may hold it while sampling or writing.
SamplerWithoutReplacement inherits the fail-closed requires_shared_state
default, so it already covers undeclared samplers without a dedicated test
class. The same-shape list collation covered by the ListStorage case is
exercised by the ConcatDataset test.
@theap06

theap06 commented Sep 27, 2026

Copy link
Copy Markdown
Collaborator Author

cc @bsprenger

@theap06
theap06 merged commit 33d8cb8 into pytorch:main Sep 28, 2026
122 of 125 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Benchmarks rl/benchmark changes CI Has to do with CI setup (e.g. wheels & builds, tests...) CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. Data/openx Data Data-related PR, will launch data-related jobs Documentation Improvements or additions to documentation Examples Feature New feature Integrations/torch_geometric Integrations llm/ LLM-related PR, triggers LLM CI tests Modules ReplayBuffers

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants