(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)
| 218 | yield batch |
| 219 | |
| 220 | def 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 | |
| 273 | def calibrate_linear_vel(base_action, c=None): |
| 274 | if c is None: |
no test coverage detected