(_worker_its)
| 95 | global_batch_size=global_batch_size, seed=32, start_index=start))) |
| 96 | |
| 97 | def get_device_batch(_worker_its): |
| 98 | while True: |
| 99 | for it in _worker_its: |
| 100 | batch = [] |
| 101 | for _ in range(device_batch_size): |
| 102 | batch.append(next(it)) |
| 103 | yield batch |
| 104 | |
| 105 | device_iterators.append(get_device_batch(worker_iterators)) |
| 106 | torch.utils.data.get_worker_info = bk |