MCPcopy Create free account
hub / github.com/MarkFzp/act-plus-plus / load_data

Function load_data

utils.py:220–271  ·  view source on GitHub ↗
(dataset_dir_l, name_filter, camera_names, batch_size_train, batch_size_val, chunk_size, skip_mirrored_data=False, load_pretrain=False, policy_class=None, stats_dir_l=None, sample_weights=None, train_ratio=0.99)

Source from the content-addressed store, hash-verified

218 yield batch
219
220def load_data(dataset_dir_l, name_filter, camera_names, batch_size_train, batch_size_val, chunk_size, skip_mirrored_data=False, load_pretrain=False, policy_class=None, stats_dir_l=None, sample_weights=None, train_ratio=0.99):
221 if type(dataset_dir_l) == str:
222 dataset_dir_l = [dataset_dir_l]
223 dataset_path_list_list = [find_all_hdf5(dataset_dir, skip_mirrored_data) for dataset_dir in dataset_dir_l]
224 num_episodes_0 = len(dataset_path_list_list[0])
225 dataset_path_list = flatten_list(dataset_path_list_list)
226 dataset_path_list = [n for n in dataset_path_list if name_filter(n)]
227 num_episodes_l = [len(dataset_path_list) for dataset_path_list in dataset_path_list_list]
228 num_episodes_cumsum = np.cumsum(num_episodes_l)
229
230 # obtain train test split on dataset_dir_l[0]
231 shuffled_episode_ids_0 = np.random.permutation(num_episodes_0)
232 train_episode_ids_0 = shuffled_episode_ids_0[:int(train_ratio * num_episodes_0)]
233 val_episode_ids_0 = shuffled_episode_ids_0[int(train_ratio * num_episodes_0):]
234 train_episode_ids_l = [train_episode_ids_0] + [np.arange(num_episodes) + num_episodes_cumsum[idx] for idx, num_episodes in enumerate(num_episodes_l[1:])]
235 val_episode_ids_l = [val_episode_ids_0]
236 train_episode_ids = np.concatenate(train_episode_ids_l)
237 val_episode_ids = np.concatenate(val_episode_ids_l)
238 print(f'\n\nData from: {dataset_dir_l}\n- Train on {[len(x) for x in train_episode_ids_l]} episodes\n- Test on {[len(x) for x in val_episode_ids_l]} episodes\n\n')
239
240 # obtain normalization stats for qpos and action
241 # if load_pretrain:
242 # with open(os.path.join('/home/zfu/interbotix_ws/src/act/ckpts/pretrain_all', 'dataset_stats.pkl'), 'rb') as f:
243 # norm_stats = pickle.load(f)
244 # print('Loaded pretrain dataset stats')
245 _, all_episode_len = get_norm_stats(dataset_path_list)
246 train_episode_len_l = [[all_episode_len[i] for i in train_episode_ids] for train_episode_ids in train_episode_ids_l]
247 val_episode_len_l = [[all_episode_len[i] for i in val_episode_ids] for val_episode_ids in val_episode_ids_l]
248 train_episode_len = flatten_list(train_episode_len_l)
249 val_episode_len = flatten_list(val_episode_len_l)
250 if stats_dir_l is None:
251 stats_dir_l = dataset_dir_l
252 elif type(stats_dir_l) == str:
253 stats_dir_l = [stats_dir_l]
254 norm_stats, _ = get_norm_stats(flatten_list([find_all_hdf5(stats_dir, skip_mirrored_data) for stats_dir in stats_dir_l]))
255 print(f'Norm stats from: {stats_dir_l}')
256
257 batch_sampler_train = BatchSampler(batch_size_train, train_episode_len_l, sample_weights)
258 batch_sampler_val = BatchSampler(batch_size_val, val_episode_len_l, None)
259
260 # print(f'train_episode_len: {train_episode_len}, val_episode_len: {val_episode_len}, train_episode_ids: {train_episode_ids}, val_episode_ids: {val_episode_ids}')
261
262 # construct dataset and dataloader
263 train_dataset = EpisodicDataset(dataset_path_list, camera_names, norm_stats, train_episode_ids, train_episode_len, chunk_size, policy_class)
264 val_dataset = EpisodicDataset(dataset_path_list, camera_names, norm_stats, val_episode_ids, val_episode_len, chunk_size, policy_class)
265 train_num_workers = (8 if os.getlogin() == 'zfu' else 16) if train_dataset.augment_images else 2
266 val_num_workers = 8 if train_dataset.augment_images else 2
267 print(f'Augment images: {train_dataset.augment_images}, train_num_workers: {train_num_workers}, val_num_workers: {val_num_workers}')
268 train_dataloader = DataLoader(train_dataset, batch_sampler=batch_sampler_train, pin_memory=True, num_workers=train_num_workers, prefetch_factor=2)
269 val_dataloader = DataLoader(val_dataset, batch_sampler=batch_sampler_val, pin_memory=True, num_workers=val_num_workers, prefetch_factor=2)
270
271 return train_dataloader, val_dataloader, norm_stats, train_dataset.is_sim
272
273def calibrate_linear_vel(base_action, c=None):
274 if c is None:

Callers 2

mainFunction · 0.90
mainFunction · 0.90

Calls 6

find_all_hdf5Function · 0.85
flatten_listFunction · 0.85
printFunction · 0.85
BatchSamplerFunction · 0.85
get_norm_statsFunction · 0.70
EpisodicDatasetClass · 0.70

Tested by

no test coverage detected