MCPcopy Create free account
hub / github.com/SpatialVLA/SpatialVLA / apply_trajectory_transforms

Function apply_trajectory_transforms

data/rlds.py:246–348  ·  view source on GitHub ↗

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,
)

Source from the content-addressed store, hash-verified

244 return dataset, dataset_statistics
245
246def 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"]:

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected