Rate this Page

filter_trajectories#

torchrl.data.filter_trajectories(data: TensorDictBase, predicate: Callable[[Trajectory], bool] | None = None, *, trajectory_key: NestedKey | None = None) list[Trajectory][source]#

Split data into trajectories and keep those matching predicate.

Parameters:
  • data (TensorDictBase) – a tensordict of transitions with a single batch dimension.

  • predicate (Callable[[Trajectory], bool], optional) – a TrajectoryPredicate built from traj, or any callable mapping a Trajectory to a boolean. Defaults to None (keep all trajectories).

Keyword Arguments:

trajectory_key (NestedKey, optional) – entry holding per-transition trajectory ids. Defaults to None (auto-detection).

Returns:

A list of matching Trajectory views.

Examples

>>> from torchrl.data import filter_trajectories, traj
>>> good = filter_trajectories(data, traj.reward.sum() > 100)