Data loader. Note that batch-size is the local (per GPU) batch-size.
(dataset,
micro_batch_size,
num_workers,
drop_last,
task_collate_fn=None)
| 80 | |
| 81 | |
| 82 | def build_data_loader(dataset, |
| 83 | micro_batch_size, |
| 84 | num_workers, |
| 85 | drop_last, |
| 86 | task_collate_fn=None): |
| 87 | """Data loader. Note that batch-size is the local (per GPU) batch-size.""" |
| 88 | |
| 89 | # Sampler. |
| 90 | world_size = mpu.get_data_parallel_world_size() |
| 91 | rank = mpu.get_data_parallel_rank() |
| 92 | sampler = torch.utils.data.distributed.DistributedSampler( |
| 93 | dataset, num_replicas=world_size, rank=rank) |
| 94 | |
| 95 | # Data loader. Note that batch size is the per GPU batch size. |
| 96 | data_loader = torch.utils.data.DataLoader(dataset, |
| 97 | batch_size=micro_batch_size, |
| 98 | sampler=sampler, |
| 99 | shuffle=False, |
| 100 | num_workers=num_workers, |
| 101 | drop_last=drop_last, |
| 102 | pin_memory=True, |
| 103 | collate_fn=task_collate_fn) |
| 104 | |
| 105 | return data_loader |
| 106 | |
| 107 | |
| 108 | def _build_infinite_size_dataloader(dataloader): |
no outgoing calls
no test coverage detected