(dataset, sampler, total_batch_size, *, num_workers=0)
| 231 | |
| 232 | |
| 233 | def build_batch_data_loader(dataset, sampler, total_batch_size, *, num_workers=0): |
| 234 | world_size = get_world_size() |
| 235 | assert ( |
| 236 | total_batch_size > 0 and total_batch_size % world_size == 0 |
| 237 | ), "Total batch size ({}) must be divisible by the number of gpus ({}).".format( |
| 238 | total_batch_size, world_size |
| 239 | ) |
| 240 | |
| 241 | batch_size = total_batch_size // world_size |
| 242 | batch_sampler = torch.utils.data.sampler.BatchSampler( |
| 243 | sampler, batch_size, drop_last=True |
| 244 | ) # drop_last so the batch always have the same size |
| 245 | return torch.utils.data.DataLoader( |
| 246 | dataset, |
| 247 | num_workers=num_workers, |
| 248 | batch_sampler=batch_sampler, |
| 249 | collate_fn=trivial_batch_collator, |
| 250 | worker_init_fn=worker_init_reset_seed, |
| 251 | ) |
| 252 | |
| 253 | |
| 254 | def trivial_batch_collator(batch): |
no outgoing calls
no test coverage detected