Buld dataloader given an input dataset.
(dataset)
| 90 | |
| 91 | |
| 92 | def make_data_loader(dataset): |
| 93 | """Buld dataloader given an input dataset.""" |
| 94 | if dataset is None: |
| 95 | return None |
| 96 | args = get_args() |
| 97 | |
| 98 | # Data parallel arguments. |
| 99 | world_size = mpu.get_data_parallel_world_size() |
| 100 | rank = mpu.get_data_parallel_rank() |
| 101 | global_batch_size = args.batch_size * world_size |
| 102 | num_workers = args.num_workers |
| 103 | |
| 104 | # Use a simple sampler with distributed batch sampler. |
| 105 | sampler = torch.utils.data.SequentialSampler(dataset) |
| 106 | batch_sampler = DistributedBatchSampler(sampler=sampler, |
| 107 | batch_size=global_batch_size, |
| 108 | drop_last=True, |
| 109 | rank=rank, |
| 110 | world_size=world_size) |
| 111 | # Torch dataloader. |
| 112 | return torch.utils.data.DataLoader(dataset, |
| 113 | batch_sampler=batch_sampler, |
| 114 | num_workers=num_workers, |
| 115 | pin_memory=True) |
| 116 | |
| 117 | def average_losses_across_data_parallel_group(losses): |
| 118 | """Reduce a tensor of losses across all GPUs.""" |
no test coverage detected