Dataset for supervised fine-tuning.
| 14 | from .rlds import dataset_statistics, build_interleaved_dataset |
| 15 | |
| 16 | class OpenXIterableDataset(IterableDataset): |
| 17 | """Dataset for supervised fine-tuning.""" |
| 18 | |
| 19 | def __init__( |
| 20 | self, |
| 21 | data_root_dir, |
| 22 | output_dir, |
| 23 | data_mix, |
| 24 | image_size=224, |
| 25 | max_length=1024, |
| 26 | is_train=True, |
| 27 | shuffle_buffer_size=1000_000, |
| 28 | tsfm_thread_muti=1, |
| 29 | read_thread_muti=1, |
| 30 | obs_backward_steps=0, |
| 31 | obs_backward_delta=1, |
| 32 | action_forward_steps=0, |
| 33 | use_raw_dataloader=False, |
| 34 | fix_raw_length=None, |
| 35 | vla_processor=None, |
| 36 | ): |
| 37 | super(OpenXIterableDataset, self).__init__() |
| 38 | self.data_root_dir = data_root_dir |
| 39 | self.data_mix = data_mix |
| 40 | self.use_raw_dataloader = use_raw_dataloader |
| 41 | self.vla_processor = vla_processor |
| 42 | self.image_size = image_size |
| 43 | self.max_length = max_length |
| 44 | self.is_train = is_train |
| 45 | |
| 46 | self.total_ranks = torch.distributed.get_world_size() |
| 47 | self.current_rank = torch.distributed.get_rank() |
| 48 | |
| 49 | if self.data_mix in OXE_NAMED_MIXTURES: |
| 50 | mixture_spec = OXE_NAMED_MIXTURES[self.data_mix] |
| 51 | else: |
| 52 | mixture_spec = [(os.path.join(self.data_mix, "1.0.0"), 1.0)] |
| 53 | per_dataset_kwargs, weights = get_oxe_dataset_kwargs_and_weights( |
| 54 | self.data_root_dir, |
| 55 | mixture_spec, |
| 56 | load_camera_views=("primary",), |
| 57 | load_depth=False, |
| 58 | load_proprio=False, |
| 59 | load_language=True, |
| 60 | action_proprio_normalization_type=NormalizationType.BOUNDS_Q99, |
| 61 | ) |
| 62 | self.dataset_num = len(weights) |
| 63 | self.rlds_config = dict( |
| 64 | traj_transform_kwargs=dict( |
| 65 | backward_windows_size=obs_backward_steps, # If we wanted to feed / predict more than one step |
| 66 | backward_delta=obs_backward_delta, |
| 67 | forward_window_size=action_forward_steps, # For action chunking |
| 68 | skip_unlabeled=True, # Skip trajectories without language labels |
| 69 | goal_relabeling_strategy="uniform", # Goals are currently unused |
| 70 | ), |
| 71 | frame_transform_kwargs=dict( |
| 72 | resize_size=(self.image_size, self.image_size), |
| 73 | num_parallel_calls=16, # For CPU-intensive ops (decoding, resizing, etc.) |