tensordict_collate#
- torchrl.data.tensordict_collate(batch: Any) Any[source]#
Collate function for a
torch.utils.data.DataLoaderreading TorchRL storages or buffers.A batch that is already a tensor, a tensor collection or a tuple of them, as fetched by
StorageDatasetor yielded byReplayBufferDataset, is returned unchanged. A list of samples, as produced by per-item storages or by composing datasets withtorch.utils.data.ConcatDataset, is stacked: tensor collections lazily when their shapes differ, tensors densely, and mappings or tuples element-wise. The default torch collation iterates a tensordict over its batch dimension and cannot be used.- Parameters:
batch (Tensor, TensorDictBase or list) – the fetched batch.
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])) >>> loader = DataLoader( ... rb.storage.as_dataset(), ... batch_size=4, ... shuffle=True, ... collate_fn=tensordict_collate, ... ) >>> next(iter(loader))["obs"].shape torch.Size([4])