Create the data loader for the corresponding trainign set
(args, TRAIN_DS='C+T+K+S+H')
| 198 | |
| 199 | |
| 200 | def fetch_dataloader(args, TRAIN_DS='C+T+K+S+H'): |
| 201 | """ Create the data loader for the corresponding trainign set """ |
| 202 | |
| 203 | if args.stage == 'chairs': |
| 204 | aug_params = {'crop_size': args.image_size, 'min_scale': -0.1, 'max_scale': 1.0, 'do_flip': True} |
| 205 | train_dataset = FlyingChairs(aug_params, split='training') |
| 206 | |
| 207 | elif args.stage == 'things': |
| 208 | aug_params = {'crop_size': args.image_size, 'min_scale': -0.4, 'max_scale': 0.8, 'do_flip': True} |
| 209 | clean_dataset = FlyingThings3D(aug_params, dstype='frames_cleanpass') |
| 210 | final_dataset = FlyingThings3D(aug_params, dstype='frames_finalpass') |
| 211 | train_dataset = clean_dataset + final_dataset |
| 212 | |
| 213 | elif args.stage == 'sintel': |
| 214 | aug_params = {'crop_size': args.image_size, 'min_scale': -0.2, 'max_scale': 0.6, 'do_flip': True} |
| 215 | things = FlyingThings3D(aug_params, dstype='frames_cleanpass') |
| 216 | sintel_clean = MpiSintel(aug_params, split='training', dstype='clean') |
| 217 | sintel_final = MpiSintel(aug_params, split='training', dstype='final') |
| 218 | |
| 219 | if TRAIN_DS == 'C+T+K+S+H': |
| 220 | kitti = KITTI({'crop_size': args.image_size, 'min_scale': -0.3, 'max_scale': 0.5, 'do_flip': True}) |
| 221 | hd1k = HD1K({'crop_size': args.image_size, 'min_scale': -0.5, 'max_scale': 0.2, 'do_flip': True}) |
| 222 | train_dataset = 100*sintel_clean + 100*sintel_final + 200*kitti + 5*hd1k + things |
| 223 | |
| 224 | elif TRAIN_DS == 'C+T+K/S': |
| 225 | train_dataset = 100*sintel_clean + 100*sintel_final + things |
| 226 | |
| 227 | elif args.stage == 'kitti': |
| 228 | aug_params = {'crop_size': args.image_size, 'min_scale': -0.2, 'max_scale': 0.4, 'do_flip': False} |
| 229 | train_dataset = KITTI(aug_params, split='training') |
| 230 | |
| 231 | train_loader = data.DataLoader(train_dataset, batch_size=args.batch_size, |
| 232 | pin_memory=False, shuffle=True, num_workers=128, drop_last=True) |
| 233 | |
| 234 | print('Training with %d image pairs' % len(train_dataset)) |
| 235 | return train_loader |
nothing calls this directly
no test coverage detected