MCPcopy Create free account
hub / github.com/OpenBMB/AgentCPM-GUI / GlobalDistributed0MQDataLoader

Class GlobalDistributed0MQDataLoader

rft/trainer/utils/dataloader.py:11–181  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

9import torch.distributed as dist
10
11class 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(

Callers 1

get_train_dataloaderMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected