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

Class OpenXIterableDataset

data/dataset.py:16–172  ·  view source on GitHub ↗

Dataset for supervised fine-tuning.

Source from the content-addressed store, hash-verified

14from .rlds import dataset_statistics, build_interleaved_dataset
15
16class 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.)

Callers 1

build_datasetsFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected