Build PyTorch DataLoader. In distributed training, each GPU/process has a dataloader. In non-distributed training, there is only one dataloader for all GPUs. Args: dataset (:obj:`Dataset`): A PyTorch dataset. samples_per_gpu (int): Number of training samples on each GPU,
(dataset: Dataset,
samples_per_gpu: int,
workers_per_gpu: int,
num_gpus: Optional[int] = 1,
dist: Optional[bool] = True,
shuffle: Optional[bool] = True,
round_up: Optional[bool] = True,
seed: Optional[Union[int, None]] = None,
persistent_workers: Optional[bool] = True,
**kwargs)
| 44 | |
| 45 | |
| 46 | def build_dataloader(dataset: Dataset, |
| 47 | samples_per_gpu: int, |
| 48 | workers_per_gpu: int, |
| 49 | num_gpus: Optional[int] = 1, |
| 50 | dist: Optional[bool] = True, |
| 51 | shuffle: Optional[bool] = True, |
| 52 | round_up: Optional[bool] = True, |
| 53 | seed: Optional[Union[int, None]] = None, |
| 54 | persistent_workers: Optional[bool] = True, |
| 55 | **kwargs): |
| 56 | """Build PyTorch DataLoader. |
| 57 | In distributed training, each GPU/process has a dataloader. |
| 58 | In non-distributed training, there is only one dataloader for all GPUs. |
| 59 | Args: |
| 60 | dataset (:obj:`Dataset`): A PyTorch dataset. |
| 61 | samples_per_gpu (int): Number of training samples on each GPU, i.e., |
| 62 | batch size of each GPU. |
| 63 | workers_per_gpu (int): How many subprocesses to use for data loading |
| 64 | for each GPU. |
| 65 | num_gpus (int, optional): Number of GPUs. Only used in non-distributed |
| 66 | training. |
| 67 | dist (bool, optional): Distributed training/test or not. Default: True. |
| 68 | shuffle (bool, optional): Whether to shuffle the data at every epoch. |
| 69 | Default: True. |
| 70 | round_up (bool, optional): Whether to round up the length of dataset by |
| 71 | adding extra samples to make it evenly divisible. Default: True. |
| 72 | kwargs: any keyword argument to be used to initialize DataLoader |
| 73 | Returns: |
| 74 | DataLoader: A PyTorch dataloader. |
| 75 | """ |
| 76 | rank, world_size = get_dist_info() |
| 77 | if dist: |
| 78 | sampler = DistributedSampler( |
| 79 | dataset, world_size, rank, shuffle=shuffle, round_up=round_up) |
| 80 | shuffle = False |
| 81 | batch_size = samples_per_gpu |
| 82 | num_workers = workers_per_gpu |
| 83 | else: |
| 84 | sampler = None |
| 85 | batch_size = num_gpus * samples_per_gpu |
| 86 | num_workers = num_gpus * workers_per_gpu |
| 87 | |
| 88 | init_fn = partial( |
| 89 | worker_init_fn, num_workers=num_workers, rank=rank, |
| 90 | seed=seed) if seed is not None else None |
| 91 | |
| 92 | data_loader = DataLoader( |
| 93 | dataset, |
| 94 | batch_size=batch_size, |
| 95 | sampler=sampler, |
| 96 | num_workers=num_workers, |
| 97 | collate_fn=partial(collate, samples_per_gpu=samples_per_gpu), |
| 98 | pin_memory=False, |
| 99 | shuffle=shuffle, |
| 100 | worker_init_fn=init_fn, |
| 101 | persistent_workers=persistent_workers, |
| 102 | **kwargs) |
| 103 |