Data loader. Note that batch-size is the local (per GPU) batch-size.
(dataset, batch_size, num_workers, drop_last)
| 72 | |
| 73 | |
| 74 | def build_data_loader(dataset, batch_size, num_workers, drop_last): |
| 75 | """Data loader. Note that batch-size is the local (per GPU) batch-size.""" |
| 76 | |
| 77 | # Sampler. |
| 78 | world_size = mpu.get_data_parallel_world_size() |
| 79 | rank = mpu.get_data_parallel_rank() |
| 80 | sampler = torch.utils.data.distributed.DistributedSampler( |
| 81 | dataset, num_replicas=world_size, rank=rank) |
| 82 | |
| 83 | # Data loader. Note that batch size is the per GPU batch size. |
| 84 | data_loader = torch.utils.data.DataLoader(dataset, |
| 85 | batch_size=batch_size, |
| 86 | sampler=sampler, |
| 87 | shuffle=False, |
| 88 | num_workers=num_workers, |
| 89 | drop_last=drop_last, |
| 90 | pin_memory=True) |
| 91 | |
| 92 | return data_loader |
| 93 | |
| 94 | |
| 95 | def _build_infinite_size_dataloader(dataloader): |
no test coverage detected