StorageDataset#
- class torchrl.data.replay_buffers.StorageDataset(storage: Storage)[source]#
A map-style
torch.utils.data.Datasetreading a TorchRL storage.The dataset has one item per storage entry and a
torch.utils.data.DataLoaderreads it with any torch sampler and worker processes.__getitems__fetches an index batch with a singleget()call, so the loader receives one batch rather than a list of items: passtensordict_collate()ascollate_fn, which returns such batches unchanged and stacks lists of items. Storages with more than one dimension are read throughflatten(). Reading the storage directly with a DataLoader, without this adapter, fetches items one by one and passes the list to the collate function.DataLoader workers read the storage content live: creating the dataset moves a CPU tensor storage to shared memory and memory-mapped storages are read through their files, so rows written after the workers start are visible to them under every start method. Reads are not synchronized with writes, and a row written while a worker reads it can come back partially updated. A
ListStoragecannot be sent to spawned workers and forked workers read a snapshot of it. Only the storage is sent to the workers, not the buffers attached to it.- Parameters:
storage (Storage) – the storage to read. Must be one-dimensional.
Examples
>>> import torch >>> from tensordict import TensorDict >>> from torch.utils.data import DataLoader >>> from torchrl.data import LazyTensorStorage, ReplayBuffer, tensordict_collate >>> rb = ReplayBuffer(storage=LazyTensorStorage(100)) >>> _ = rb.extend(TensorDict({"obs": torch.arange(100)}, [100])) >>> dataset = rb.storage.as_dataset() >>> len(dataset) 100 >>> loader = DataLoader( ... dataset, batch_size=4, shuffle=True, collate_fn=tensordict_collate ... ) >>> next(iter(loader))["obs"].shape torch.Size([4])