MCPcopy Create free account
hub / github.com/allenai/molmo / test_distributed

Function test_distributed

tests/data/test_data_iterator.py:80–123  ·  view source on GitHub ↗
(world_size, num_workers, device_batch_size)

Source from the content-addressed store, hash-verified

78 (3, 6, 4),
79])
80def 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", [

Callers

nothing calls this directly

Calls 4

MockDatasetClass · 0.85
MockWorkerInfoClass · 0.85
get_device_batchFunction · 0.85

Tested by

no test coverage detected