Applies common transforms that happen at a trajectory level. Such transforms are usually some sort of "relabeling" (e.g., filtering, chunking, adding goals, dropping keys). Transforms in this function should have the following properties: - They require access to an entire traj
(
dataset: dl.DLataset,
*,
train: bool,
goal_relabeling_strategy: Optional[str] = None,
goal_relabeling_kwargs: dict = {},
backward_windows_size: int = 0,
backward_delta: int = 1,
forward_window_size: int = 0,
subsample_length: Optional[int] = None,
skip_unlabeled: bool = False,
max_action: Optional[float] = None,
max_proprio: Optional[float] = None,
task_augment_strategy: Optional[str] = None,
task_augment_kwargs: dict = {},
num_parallel_calls: int = tf.data.AUTOTUNE,
)
| 244 | return dataset, dataset_statistics |
| 245 | |
| 246 | def apply_trajectory_transforms( |
| 247 | dataset: dl.DLataset, |
| 248 | *, |
| 249 | train: bool, |
| 250 | goal_relabeling_strategy: Optional[str] = None, |
| 251 | goal_relabeling_kwargs: dict = {}, |
| 252 | backward_windows_size: int = 0, |
| 253 | backward_delta: int = 1, |
| 254 | forward_window_size: int = 0, |
| 255 | subsample_length: Optional[int] = None, |
| 256 | skip_unlabeled: bool = False, |
| 257 | max_action: Optional[float] = None, |
| 258 | max_proprio: Optional[float] = None, |
| 259 | task_augment_strategy: Optional[str] = None, |
| 260 | task_augment_kwargs: dict = {}, |
| 261 | num_parallel_calls: int = tf.data.AUTOTUNE, |
| 262 | ) -> dl.DLataset: |
| 263 | """ |
| 264 | Applies common transforms that happen at a trajectory level. Such transforms are usually some sort of "relabeling" |
| 265 | (e.g., filtering, chunking, adding goals, dropping keys). |
| 266 | |
| 267 | Transforms in this function should have the following properties: |
| 268 | - They require access to an entire trajectory (i.e., they cannot be applied frame-wise). |
| 269 | - They are generally not CPU-intensive, mostly involving moving and copying data. |
| 270 | - They do not require decoded images. |
| 271 | |
| 272 | Args: |
| 273 | dataset (dl.DLataset): The dataset to transform. |
| 274 | train (bool): Whether the dataset is for training (affects subsampling). |
| 275 | goal_relabeling_strategy (str, optional): The goal relabeling strategy to use, or None for |
| 276 | no goal relabeling. See `goal_relabeling.py`. |
| 277 | goal_relabeling_kwargs (dict, optional): Additional keyword arguments to pass to the goal relabeling function. |
| 278 | window_size (int, optional): The length of the snippets that trajectories are chunked into. |
| 279 | future_action_window_size (int, optional): The number of future actions beyond window_size to include |
| 280 | in the chunked actions. |
| 281 | subsample_length (int, optional): If provided, trajectories longer than this will be subsampled to |
| 282 | this length (after goal relabeling and chunking). |
| 283 | skip_unlabeled (bool, optional): Whether to skip trajectories with no language labels. |
| 284 | max_action: (float, optional): If provided, trajectories in which *any* action dimension |
| 285 | of *any* transition has an absolute value larger than this will be skipped. |
| 286 | max_proprio: (float, optional): If provided, trajectories in which *any* proprio dimension |
| 287 | of *any* transition has an absolute value larger than this will be skipped. |
| 288 | task_augment_strategy (str, optional): The task augmentation strategy to use, or None for no task |
| 289 | augmentation. See `task_augmentation.py`. |
| 290 | task_augment_kwargs (dict, optional): Additional keyword arguments to pass to the task augmentation |
| 291 | function. |
| 292 | num_parallel_calls (int, optional): number of parallel calls for map operations. Default to AUTOTUNE. |
| 293 | """ |
| 294 | if skip_unlabeled: |
| 295 | if "language_instruction" not in dataset.element_spec["task"]: |
| 296 | raise ValueError("skip_unlabeled=True but dataset does not have language labels.") |
| 297 | |
| 298 | dataset = dataset.filter(lambda x: tf.math.reduce_any(x["task"]["language_instruction"] != "")) |
| 299 | |
| 300 | if max_action is not None: |
| 301 | dataset = dataset.filter(lambda x: tf.math.reduce_all(tf.math.abs(x["action"]) <= max_action)) |
| 302 | |
| 303 | if max_proprio is not None and "proprio" in dataset.element_spec["observation"]: |
no outgoing calls
no test coverage detected