Create the data loader for the corresponding trainign set
(args, TRAIN_DS='C+T+K+S+H')
| 422 | |
| 423 | |
| 424 | def fetch_dataloader(args, TRAIN_DS='C+T+K+S+H'): |
| 425 | """ Create the data loader for the corresponding trainign set """ |
| 426 | |
| 427 | if args.stage == 'chairs': |
| 428 | if hasattr(args.percostformer, 'pwc_aug') and args.percostformer.pwc_aug: |
| 429 | aug_params = {'crop_size': args.image_size, 'min_scale': -0.1, 'max_scale': 1.0, 'do_flip': True, 'pwc_aug': True} |
| 430 | else: |
| 431 | aug_params = {'crop_size': args.image_size, 'min_scale': -0.1, 'max_scale': 1.0, 'do_flip': True} |
| 432 | train_dataset = FlyingChairs(aug_params, split='training') |
| 433 | |
| 434 | elif args.stage == 'things': |
| 435 | aug_params = {'crop_size': args.image_size, 'min_scale': -0.4, 'max_scale': 0.8, 'do_flip': True} |
| 436 | clean_dataset = FlyingThings3D(aug_params, dstype='frames_cleanpass') |
| 437 | final_dataset = FlyingThings3D(aug_params, dstype='frames_finalpass') |
| 438 | train_dataset = clean_dataset + final_dataset |
| 439 | |
| 440 | elif args.stage == 'sintel': |
| 441 | aug_params = {'crop_size': args.image_size, 'min_scale': -0.2, 'max_scale': 0.6, 'do_flip': True} |
| 442 | things = FlyingThings3D(aug_params, dstype='frames_cleanpass') |
| 443 | sintel_clean = MpiSintel(aug_params, split='training', dstype='clean') |
| 444 | sintel_final = MpiSintel(aug_params, split='training', dstype='final') |
| 445 | |
| 446 | if TRAIN_DS == 'C+T+K+S+H': |
| 447 | kitti = KITTI({'crop_size': args.image_size, 'min_scale': -0.3, 'max_scale': 0.5, 'do_flip': True}) |
| 448 | hd1k = HD1K({'crop_size': args.image_size, 'min_scale': -0.5, 'max_scale': 0.2, 'do_flip': True}) |
| 449 | train_dataset = 100*sintel_clean + 100*sintel_final + 200*kitti + 5*hd1k + things |
| 450 | |
| 451 | elif TRAIN_DS == 'C+T+K/S': |
| 452 | train_dataset = 100*sintel_clean + 100*sintel_final + things |
| 453 | |
| 454 | elif args.stage == 'kitti': |
| 455 | aug_params = {'crop_size': args.image_size, 'min_scale': -0.2, 'max_scale': 0.4, 'do_flip': False} |
| 456 | train_dataset = KITTI(aug_params, split='training') |
| 457 | |
| 458 | elif args.stage == 'autoflow-pwcaug': |
| 459 | aug_params = {'num_steps': args.trainer.num_steps, 'crop_size': args.image_size, 'log_dir': args.log_dir} |
| 460 | train_dataset = AutoFlow(**aug_params) |
| 461 | |
| 462 | train_loader = data.DataLoader(train_dataset, batch_size=args.batch_size, |
| 463 | pin_memory=False, shuffle=True, num_workers=args.batch_size, drop_last=True) |
| 464 | |
| 465 | print('Training with %d image pairs' % len(train_dataset)) |
| 466 | return train_loader |
| 467 | |
| 468 | if __name__ == "__main__": |
| 469 | aug_params = {'crop_size': [400, 720], 'min_scale': -0.2, 'max_scale': 0, 'do_flip': True} |
nothing calls this directly
no test coverage detected