RayReplayBuffer#
- class torchrl.data.RayReplayBuffer(*args, use_ray_service=False, service_backend=None, service_backend_options=None, **kwargs)[source]#
A Ray implementation of the Replay Buffer that can be extended and sampled remotely.
- Keyword Arguments:
replay_buffer_cls (type[ReplayBuffer], optional) – the class to use for the replay buffer. Defaults to
ReplayBuffer.ray_init_config (dict[str, Any], optional) – keyword arguments to pass to ray.init().
remote_config (dict[str, Any], optional) – keyword arguments to pass to cls.as_remote(). Defaults to torchrl.collectors.distributed.ray.DEFAULT_REMOTE_CLASS_CONFIG.
**kwargs – keyword arguments to pass to the replay buffer class.
See also
ReplayBufferfor a list of other keyword arguments.The writer, sampler and storage should be passed as constructors to prevent serialization issues. Transforms constructors should be passed through the transform_factory argument.
Example
>>> from functools import partial >>> import torch >>> from tensordict import TensorDict >>> from torchrl.data import LazyTensorStorage, ReplayBuffer >>> buffer = ReplayBuffer( ... storage=partial(LazyTensorStorage, 100), ... batch_size=4, ... service_backend="ray", ... service_backend_options={"remote_config": {"num_cpus": 0}}, ... transport="auto", ... ) >>> buffer.extend(TensorDict({"value": torch.arange(8)}, batch_size=[8])) >>> buffer.sample().shape torch.Size([4]) >>> buffer.shutdown()
- add(*args, **kwargs)[source]#
Add a single element to the replay buffer.
- Parameters:
data (Any) – data to be added to the replay buffer
- Returns:
index where the data lives in the replay buffer.
- append_transform(*args, **kwargs)[source]#
Appends transform at the end.
Transforms are applied in order when sample is called.
- Parameters:
transform (Transform) – The transform to be appended
- Keyword Arguments:
invert (bool, optional) – if
True, the transform will be inverted (forward calls will be called during writing and inverse calls during reading). Defaults toFalse.
Example
>>> rb = ReplayBuffer(storage=LazyMemmapStorage(10), batch_size=4) >>> data = TensorDict({"a": torch.zeros(10)}, [10]) >>> def t(data): ... data += 1 ... return data >>> rb.append_transform(t, invert=True) >>> rb.extend(data) >>> assert (data == 1).all()
- classmethod as_remote(remote_config=None)#
Creates an instance of a remote ray class.
- Parameters:
cls (Python Class) – class to be remotely instantiated.
remote_config (dict) – the quantity of CPU cores to reserve for this class. Defaults to torchrl.collectors.distributed.ray.DEFAULT_REMOTE_CLASS_CONFIG.
- Returns:
A function that creates ray remote class instances.
- property batch_size#
The batch size of the replay buffer.
The batch size can be overridden by setting the batch_size parameter in the
sample()method.It defines both the number of samples returned by
sample()and the number of samples that are yielded by theReplayBufferiterator.
- clients(num_clients: int) list[Any][source]#
Return one independently routed client per concurrent consumer.
- dumps(path)[source]#
Saves the replay buffer on disk at the specified path.
- Parameters:
path (Path or str) – path where to save the replay buffer.
Examples
>>> import tempfile >>> import tqdm >>> from torchrl.data import LazyMemmapStorage, TensorDictReplayBuffer >>> from torchrl.data.replay_buffers.samplers import PrioritizedSampler, RandomSampler >>> import torch >>> from tensordict import TensorDict >>> # Build and populate the replay buffer >>> S = 1_000_000 >>> sampler = PrioritizedSampler(S, 1.1, 1.0) >>> # sampler = RandomSampler() >>> storage = LazyMemmapStorage(S) >>> rb = TensorDictReplayBuffer(storage=storage, sampler=sampler) >>> >>> for _ in tqdm.tqdm(range(100)): ... td = TensorDict({"obs": torch.randn(100, 3, 4), "next": {"obs": torch.randn(100, 3, 4)}, "td_error": torch.rand(100)}, [100]) ... rb.extend(td) ... sample = rb.sample(32) ... rb.update_tensordict_priority(sample) >>> # save and load the buffer >>> with tempfile.TemporaryDirectory() as tmpdir: ... rb.dumps(tmpdir) ... ... sampler = PrioritizedSampler(S, 1.1, 1.0) ... # sampler = RandomSampler() ... storage = LazyMemmapStorage(S) ... rb_load = TensorDictReplayBuffer(storage=storage, sampler=sampler) ... rb_load.loads(tmpdir) ... assert len(rb) == len(rb_load)
- empty(empty_write_count: bool = True)[source]#
Empties the replay buffer and reset cursor to 0.
- Parameters:
empty_write_count (bool, optional) – Whether to empty the write_count attribute. Defaults to True.
- extend(*args, **kwargs)[source]#
Extends the replay buffer with one or more elements contained in an iterable.
If present, the inverse transforms will be called.`
- Parameters:
data (iterable) – collection of data to be added to the replay buffer.
- Keyword Arguments:
update_priority (bool, optional) – Whether to update the priority of the data. Defaults to True. Without effect in this class. See
extend()for more details.- Returns:
Indices of the data added to the replay buffer.
Warning
extend()can have an ambiguous signature when dealing with lists of values, which should be interpreted either as PyTree (in which case all elements in the list will be put in a slice in the stored PyTree in the storage) or a list of values to add one at a time. To solve this, TorchRL makes the clear-cut distinction between list and tuple: a tuple will be viewed as a PyTree, a list (at the root level) will be interpreted as a stack of values to add one at a time to the buffer. ForListStorageinstances, only unbound elements can be provided (no PyTrees).
- property initialized: bool#
Whether the replay buffer has been initialized.
- insert_transform(index: int, transform: Transform, *, invert: bool = False) ReplayBuffer[source]#
Inserts transform.
Transforms are executed in order when sample is called.
- Parameters:
index (int) – Position to insert the transform.
transform (Transform) – The transform to be appended
- Keyword Arguments:
invert (bool, optional) – if
True, the transform will be inverted (forward calls will be called during writing and inverse calls during reading). Defaults toFalse.
- property is_alive: bool#
Whether the owned Ray replay-buffer actor is available.
- loads(path)[source]#
Loads a replay buffer state at the given path.
The buffer should have matching components and be saved using
dumps().- Parameters:
path (Path or str) – path where the replay buffer was saved.
See
dumps()for more info.
- next()[source]#
Returns the next item in the replay buffer.
This method is used to iterate over the replay buffer in contexts where __iter__ is not available, such as
RayReplayBuffer.
- query(predicate: Callable[[Trajectory], bool] | None = None, *, trajectory_key: NestedKey | None = None) list[Trajectory]#
Filters the stored trajectories with a query predicate.
Splits the buffer content into trajectories (see
iter_trajectories()) and returns those matching the predicate asTrajectoryviews.- Parameters:
predicate (Callable[[Trajectory], bool], optional) – a
TrajectoryPredicatebuilt fromtraj, or any callable mapping a trajectory to a boolean. Defaults to None (return all trajectories).- Keyword Arguments:
trajectory_key (NestedKey, optional) – entry holding per-transition trajectory ids. Defaults to None (auto-detection from
("collector", "traj_ids"),"traj_ids","episode"or the done/terminated/truncated flags).- Returns:
A list of matching trajectory views, ordered chronologically (oldest trajectory first; for multi-dimensional storages, grouped by batch coordinate).
The trajectory boundaries are computed from the stored (untransformed) data with the same machinery
SliceSampleruses, so samplers and queries always agree on where trajectories start and stop. This includes storages withndim > 1(e.g.LazyTensorStorage(..., ndim=2)holding[B, T]batches), whose trajectories are recovered per batch coordinate.Predicates built from
trajreport the keys they read viarequired_keys(); evaluation then only fetches those entries from the storage and only runs the transforms that can affect them. Matching trajectories are extracted in full with the complete transform chain applied, so predicates and results see the same values a sampler would produce. Opaque callables are evaluated against the fully transformed content.Note
Once the buffer has wrapped around (it is at capacity and older entries have been overwritten), the oldest trajectory may have lost its first transitions to overwriting and will appear truncated at the front. A trajectory written across the wrap point is followed through it and returned whole, in time order.
Examples
>>> from torchrl.data import traj >>> good_trajs = rb.query((traj.reward.sum() > 100) & (traj.length >= 50)) >>> observations = good_trajs[0].observation
- read_all_in_order(end: int | None = None) Any#
Read storage contents in physical order.
This is equivalent to
rb[:]whenendisNone.- Parameters:
end (int, optional) – Number of leading storage entries to read. Defaults to the entire storage slice.
- Returns:
A storage slice containing entries
[:end].
- register_load_hook(hook: Callable[[Any], Any])[source]#
Registers a load hook for the storage.
Note
Hooks are currently not serialized when saving a replay buffer: they must be manually re-initialized every time the buffer is created.
- register_save_hook(hook: Callable[[Any], Any])[source]#
Registers a save hook for the storage.
Note
Hooks are currently not serialized when saving a replay buffer: they must be manually re-initialized every time the buffer is created.
- sample(*args, **kwargs)[source]#
Samples a batch of data from the replay buffer.
Uses Sampler to sample indices, and retrieves them from Storage.
- Parameters:
batch_size (int, optional) – size of data to be collected. If none is provided, this method will sample a batch-size as indicated by the sampler.
return_info (bool) – whether to return info. If True, the result is a tuple (data, info). If False, the result is the data.
- Returns:
A batch of data selected in the replay buffer. A tuple containing this batch and info if return_info flag is set to True.
- property sampler: Sampler#
The sampler of the replay buffer.
The sampler must be an instance of
Sampler.
- property service_backend: str#
The canonical deployment backend for this replay buffer.
- set_(key, value)#
Sets the value of a key across the entire replay buffer in-place.
- Parameters:
key (NestedKey) – the key to set.
value (torch.Tensor) – the value to write.
- Returns:
self
- set_at_(key, value, index)#
Sets the value of a key at specified indices in the replay buffer.
- Parameters:
key (NestedKey) – the key to set.
value (torch.Tensor) – the value to write.
index – the indices where to write the value.
- Returns:
self
- set_sampler(sampler)[source]#
Sets a new sampler in the replay buffer and returns the previous sampler.
- set_storage(storage)[source]#
Sets a new storage in the replay buffer and returns the previous storage.
- Parameters:
storage (Storage) – the new storage for the buffer.
collate_fn (callable, optional) – if provided, the collate_fn is set to this value. Otherwise it is reset to a default value.
- start() RayReplayBuffer[source]#
Return this already-started Ray replay-buffer owner.
- stats() dict[str, int | float | bool][source]#
Returns the buffer stats snapshot through a single actor round-trip.
See
stats().
- property storage: Storage#
The storage of the replay buffer.
The storage must be an instance of
Storage.
- property transform: Transform#
The transform of the replay buffer.
The transform must be an instance of
Transform.
- property transport_kind: str#
Physical transport used for replay payloads.
- update_(input_dict_or_td, clone=False, *, keys_to_update=None)#
Updates the replay buffer in-place with the given dict or TensorDict.
- Parameters:
input_dict_or_td (dict or TensorDictBase) – the data to update with.
clone (bool, optional) – whether to clone the values before writing. Defaults to
False.keys_to_update (sequence of NestedKey, optional) – if provided, only these keys will be updated.
- Returns:
self
- update_if_present(*, index, generation, patch, version_key=None, version=None, require_newer=False)[source]#
Conditionally updates live records through a single actor round-trip.
Validation, the generation and version comparisons and the patch write all run inside the replay-buffer actor under its own lock. See
update_if_present().
- write_all(data: Any, end: int | None = None) None#
Write data back to storage in physical order.
This is equivalent to
rb[:end] = data. IfendisNone,enddefaults todata.shape[0]for tensor collections andlen(data)otherwise. Ifdataspans the full storage, this is equivalent torb[:] = data.- Parameters:
data – Data to write to storage.
end (int, optional) – Number of leading storage entries to update. Defaults to
data.shape[0]for tensor collections andlen(data)otherwise.
- property write_count#
The total number of items written so far in the buffer through add and extend.