(opt, logger)
| 27 | |
| 28 | |
| 29 | def create_train_val_dataloader(opt, logger): |
| 30 | # create train and val dataloaders |
| 31 | train_loader, val_loaders = None, [] |
| 32 | for phase, dataset_opt in opt['datasets'].items(): |
| 33 | dataset_opt['gt_size'] = opt['gt_size'] |
| 34 | if phase == 'train': |
| 35 | dataset_enlarge_ratio = dataset_opt.get('dataset_enlarge_ratio', 1) |
| 36 | train_set = build_dataset(dataset_opt) |
| 37 | train_sampler = EnlargedSampler(train_set, opt['world_size'], opt['rank'], dataset_enlarge_ratio) |
| 38 | train_loader = build_dataloader( |
| 39 | train_set, |
| 40 | dataset_opt, |
| 41 | num_gpu=opt['num_gpu'], |
| 42 | dist=opt['dist'], |
| 43 | sampler= train_sampler, |
| 44 | seed=opt['manual_seed']) |
| 45 | |
| 46 | num_iter_per_epoch = math.ceil( |
| 47 | len(train_set) * dataset_enlarge_ratio / (dataset_opt['batch_size_per_gpu'] * opt['world_size'])) |
| 48 | total_iters = int(opt['train']['total_iter']) |
| 49 | total_epochs = math.ceil(total_iters / (num_iter_per_epoch)) |
| 50 | logger.info('Training statistics:' |
| 51 | f'\n\tNumber of train images: {len(train_set)}' |
| 52 | f'\n\tDataset enlarge ratio: {dataset_enlarge_ratio}' |
| 53 | f'\n\tBatch size per gpu: {dataset_opt["batch_size_per_gpu"]}' |
| 54 | f'\n\tWorld size (gpu number): {opt["world_size"]}' |
| 55 | f'\n\tRequire iter number per epoch: {num_iter_per_epoch}' |
| 56 | f'\n\tTotal epochs: {total_epochs}; iters: {total_iters}.') |
| 57 | elif phase.split('_')[0] == 'val': |
| 58 | val_set = build_dataset(dataset_opt) |
| 59 | val_loader = build_dataloader( |
| 60 | val_set, dataset_opt, num_gpu=opt['num_gpu'], dist=opt['dist'], sampler=None, seed=opt['manual_seed']) |
| 61 | logger.info(f'Number of val images/folders in {dataset_opt["name"]}: {len(val_set)}') |
| 62 | val_loaders.append(val_loader) |
| 63 | else: |
| 64 | raise ValueError(f'Dataset phase {phase} is not recognized.') |
| 65 | |
| 66 | return train_loader, train_sampler, val_loaders, total_epochs, total_iters |
| 67 | |
| 68 | |
| 69 | def load_resume_state(opt): |
no test coverage detected