MCPcopy Create free account
hub / github.com/drinkingcoder/FlowFormer-Official / fetch_dataloader

Function fetch_dataloader

core/utils/datasets.py:424–466  ·  view source on GitHub ↗

Create the data loader for the corresponding trainign set

(args, TRAIN_DS='C+T+K+S+H')

Source from the content-addressed store, hash-verified

422
423
424def 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
468if __name__ == "__main__":
469 aug_params = {'crop_size': [400, 720], 'min_scale': -0.2, 'max_scale': 0, 'do_flip': True}

Callers

nothing calls this directly

Calls 6

AutoFlowClass · 0.85
FlyingChairsClass · 0.70
FlyingThings3DClass · 0.70
MpiSintelClass · 0.70
KITTIClass · 0.70
HD1KClass · 0.70

Tested by

no test coverage detected