(world_size, num_workers, device_batch_size)
| 78 | (3, 6, 4), |
| 79 | ]) |
| 80 | def test_distributed(world_size, num_workers, device_batch_size): |
| 81 | start = 0 |
| 82 | global_batch_size = device_batch_size*world_size |
| 83 | iterators = [] |
| 84 | datasets = [MockDataset("a", 5), MockDataset("b", 11)] |
| 85 | mixture_rates = [0.8, 0.2] |
| 86 | device_iterators = [] |
| 87 | bk = torch.utils.data.get_worker_info |
| 88 | for rank in range(world_size): |
| 89 | worker_iterators = [] |
| 90 | for worker_id in range(num_workers): |
| 91 | worker_iterators.append(iter(IterableDatasetMixture( |
| 92 | datasets, mixture_rates=mixture_rates, |
| 93 | rank=rank, world_size=world_size, |
| 94 | worker_info=MockWorkerInfo(worker_id, num_workers), |
| 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 |
| 107 | |
| 108 | grouped_by_dataset = defaultdict(list) |
| 109 | for i in range(100): |
| 110 | global_batch = [] |
| 111 | for it in device_iterators: |
| 112 | global_batch += next(it) |
| 113 | global_batch.sort(key=lambda x: x.epoch) |
| 114 | for ex in global_batch: |
| 115 | grouped_by_dataset[ex.dataset].append(ex) |
| 116 | |
| 117 | for dataset in datasets: |
| 118 | items = grouped_by_dataset[dataset.name] |
| 119 | ds_len = dataset.n |
| 120 | for epoch in range(len(items)//ds_len): |
| 121 | epoch_items = items[epoch*ds_len:(epoch+1)*ds_len] |
| 122 | assert all(x.epoch == epoch for x in epoch_items) |
| 123 | assert set(x.idx for x in epoch_items) == set(range(ds_len)) |
| 124 | |
| 125 | |
| 126 | @pytest.mark.parametrize("ns,start_index,world_size,rank", [ |
nothing calls this directly
no test coverage detected