ReplayBufferDataset#
- class torchrl.data.ReplayBufferDataset(replay_buffer: ReplayBuffer, *, num_batches: int | None = None)[source]#
A
torch.utils.data.IterableDatasetstreaming batches from a replay buffer.Iterating the dataset iterates the buffer: every item is one batch of
replay_buffer.batch_sizeelements with the buffer transforms applied. Pass the dataset to atorch.utils.data.DataLoaderwithbatch_size=Noneandtensordict_collate()to sample in worker processes. Each worker holds its own copy of the buffer, so the sampler and the transforms run in the worker andnum_batchesis split between workers. The storage content is shared rather than copied and workers observe later writes as described inStorageDataset.Buffer prefetching is disabled in workers and prefetched batches are never serialized to them, the DataLoader prefetches instead. A buffer built with a
torch.Generatoris reseeded once per worker from the worker seed, so sampling in workers is reproducible when the DataLoader is seeded (torch.manual_seedorDataLoader(generator=...)).Samplers whose
requires_shared_stateisTrue, which is every sampler except those that declare their draws stateless such asRandomSamplerandSliceSampler, are rejected when workers are used. So is aRateLimitedReplayBufferthat has not been shared withshare(), since each worker would otherwise spend its own copy of the sample budget.- Parameters:
replay_buffer (ReplayBuffer) – the buffer to sample from. Its
batch_sizemust be set.- Keyword Arguments:
num_batches (int or None, optional) – number of batches yielded by one iterator, shared between DataLoader workers.
Nonestreams batches until the sampler runs out, which samplers with replacement never do. Defaults toNone.
Examples
>>> import torch >>> from tensordict import TensorDict >>> from torch.utils.data import DataLoader >>> from torchrl.data import ( ... LazyMemmapStorage, ... SliceSampler, ... TensorDictReplayBuffer, ... tensordict_collate, ... ) >>> rb = TensorDictReplayBuffer( ... storage=LazyMemmapStorage(1000), ... sampler=SliceSampler(num_slices=4, traj_key="episode", cache_values=True), ... batch_size=32, ... ) >>> _ = rb.extend( ... TensorDict( ... {"obs": torch.randn(1000, 3), "episode": torch.arange(1000) // 50}, ... [1000], ... ) ... ) >>> loader = DataLoader( ... rb.as_dataset(num_batches=8), ... batch_size=None, ... num_workers=4, ... persistent_workers=True, ... collate_fn=tensordict_collate, ... ) >>> batches = list(loader) >>> len(batches), batches[0]["obs"].shape (8, torch.Size([32, 3]))