Build dataloader. Args: dataset (torch.utils.data.Dataset): Dataset. dataset_opt (dict): Dataset options. It contains the following keys: phase (str): 'train' or 'val'. num_worker_per_gpu (int): Number of workers for each GPU. batch_size_per_g
(dataset, dataset_opt, num_gpu=1, dist=False, sampler=None, seed=None)
| 38 | |
| 39 | |
| 40 | def build_dataloader(dataset, dataset_opt, num_gpu=1, dist=False, sampler=None, seed=None): |
| 41 | """Build dataloader. |
| 42 | |
| 43 | Args: |
| 44 | dataset (torch.utils.data.Dataset): Dataset. |
| 45 | dataset_opt (dict): Dataset options. It contains the following keys: |
| 46 | phase (str): 'train' or 'val'. |
| 47 | num_worker_per_gpu (int): Number of workers for each GPU. |
| 48 | batch_size_per_gpu (int): Training batch size for each GPU. |
| 49 | num_gpu (int): Number of GPUs. Used only in the train phase. |
| 50 | Default: 1. |
| 51 | dist (bool): Whether in distributed training. Used only in the train |
| 52 | phase. Default: False. |
| 53 | sampler (torch.utils.data.sampler): Data sampler. Default: None. |
| 54 | seed (int | None): Seed. Default: None |
| 55 | """ |
| 56 | phase = dataset_opt['phase'] |
| 57 | rank, _ = get_dist_info() |
| 58 | if phase == 'train': |
| 59 | if dist: # distributed training |
| 60 | batch_size = dataset_opt['batch_size_per_gpu'] |
| 61 | num_workers = dataset_opt['num_worker_per_gpu'] |
| 62 | else: # non-distributed training |
| 63 | multiplier = 1 if num_gpu == 0 else num_gpu |
| 64 | batch_size = dataset_opt['batch_size_per_gpu'] * multiplier |
| 65 | num_workers = dataset_opt['num_worker_per_gpu'] * multiplier |
| 66 | dataloader_args = dict( |
| 67 | dataset=dataset, |
| 68 | batch_size=batch_size, |
| 69 | shuffle=False, |
| 70 | num_workers=num_workers, |
| 71 | sampler=sampler, |
| 72 | drop_last=True) |
| 73 | if sampler is None: |
| 74 | dataloader_args['shuffle'] = True |
| 75 | dataloader_args['worker_init_fn'] = partial( |
| 76 | worker_init_fn, num_workers=num_workers, rank=rank, seed=seed) if seed is not None else None |
| 77 | elif phase in ['val', 'test']: # validation |
| 78 | dataloader_args = dict(dataset=dataset, batch_size=1, shuffle=False, num_workers=0) |
| 79 | else: |
| 80 | raise ValueError(f"Wrong dataset phase: {phase}. Supported ones are 'train', 'val' and 'test'.") |
| 81 | |
| 82 | dataloader_args['pin_memory'] = dataset_opt.get('pin_memory', False) |
| 83 | dataloader_args['persistent_workers'] = dataset_opt.get('persistent_workers', False) |
| 84 | |
| 85 | prefetch_mode = dataset_opt.get('prefetch_mode') |
| 86 | if prefetch_mode == 'cpu': # CPUPrefetcher |
| 87 | num_prefetch_queue = dataset_opt.get('num_prefetch_queue', 1) |
| 88 | logger = get_root_logger() |
| 89 | logger.info(f'Use {prefetch_mode} prefetch dataloader: num_prefetch_queue = {num_prefetch_queue}') |
| 90 | return PrefetchDataLoader(num_prefetch_queue=num_prefetch_queue, **dataloader_args) |
| 91 | else: |
| 92 | # prefetch_mode=None: Normal dataloader |
| 93 | # prefetch_mode='cuda': dataloader for CUDAPrefetcher |
| 94 | return torch.utils.data.DataLoader(**dataloader_args) |
| 95 | |
| 96 | |
| 97 | def worker_init_fn(worker_id, num_workers, rank, seed): |