filter_trajectories#
- torchrl.data.filter_trajectories(data: TensorDictBase, predicate: Callable[[Trajectory], bool] | None = None, *, trajectory_key: NestedKey | None = None) list[Trajectory][source]#
Split
datainto trajectories and keep those matchingpredicate.- Parameters:
data (TensorDictBase) – a tensordict of transitions with a single batch dimension.
predicate (Callable[[Trajectory], bool], optional) – a
TrajectoryPredicatebuilt fromtraj, or any callable mapping aTrajectoryto 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
Trajectoryviews.
Examples
>>> from torchrl.data import filter_trajectories, traj >>> good = filter_trajectories(data, traj.reward.sum() > 100)