| 9 | import torch.distributed as dist |
| 10 | |
| 11 | class GlobalDistributed0MQDataLoader: |
| 12 | def __init__( |
| 13 | self, |
| 14 | dataset: Any, |
| 15 | global_sync_address: str, |
| 16 | batch_size: int, |
| 17 | collate_fn: Callable, |
| 18 | num_workers: int, |
| 19 | sampler: Sampler, |
| 20 | worker_init_fn: Callable, |
| 21 | prefetch_factor: int, |
| 22 | world_size:int, |
| 23 | **kwargs: Any |
| 24 | ): |
| 25 | self.dataset = dataset |
| 26 | self.global_sync_address = global_sync_address |
| 27 | self.batch_size = batch_size |
| 28 | self.collate_fn = collate_fn |
| 29 | self.num_workers = num_workers |
| 30 | self.sampler = sampler |
| 31 | self.worker_init_fn = worker_init_fn |
| 32 | self.prefetch_factor = prefetch_factor |
| 33 | self.kwargs = kwargs |
| 34 | self._init_kwargs = kwargs |
| 35 | self.world_size = world_size |
| 36 | self.rank = dist.get_rank() |
| 37 | |
| 38 | self.index_queue = multiprocessing.Queue(self.num_workers) |
| 39 | self.result_queue = multiprocessing.Queue(self.prefetch_factor) |
| 40 | |
| 41 | self.workers = [ multiprocessing.Process( |
| 42 | target=GlobalDistributed0MQDataLoader._load_data, |
| 43 | args=( |
| 44 | self.dataset, |
| 45 | self.index_queue, |
| 46 | self.result_queue, |
| 47 | self.collate_fn, |
| 48 | self.worker_init_fn |
| 49 | ), |
| 50 | daemon=True |
| 51 | ) for _ in range(self.num_workers) ] |
| 52 | |
| 53 | for p in self.workers: |
| 54 | p.start() |
| 55 | |
| 56 | if self.rank == 0: |
| 57 | self.master_proc = multiprocessing.Process( |
| 58 | target=GlobalDistributed0MQDataLoader._master_loop, |
| 59 | args=(self.global_sync_address, self.batch_size,self.sampler), |
| 60 | daemon=True |
| 61 | ) |
| 62 | self.master_proc.start() |
| 63 | |
| 64 | def __len__(self): |
| 65 | return len(self.sampler) // self.world_size |
| 66 | |
| 67 | @staticmethod |
| 68 | def _master_loop( |
no outgoing calls
no test coverage detected