MCPcopy Create free account
hub / github.com/togethercomputer/OpenChatKit / __init__

Method __init__

training/data_parallel/dist_dp_local.py:8–52  ·  view source on GitHub ↗
(self, args, device, module: torch.nn.Module, optimizer: torch.optim.Optimizer = None, flatten=True)

Source from the content-addressed store, hash-verified

6
7class LocalDP:
8 def __init__(self, args, device, module: torch.nn.Module, optimizer: torch.optim.Optimizer = None, flatten=True):
9 flatten = True
10 self.flatten = flatten
11 self.global_rank = args.rank
12 self.dp_group_size = args.data_group_size
13 self.enable_tidy_profiling = (args.profiling == 'tidy_profiling')
14 self.dp_comm = get_data_parallel_comm()
15 self.dp_rank = get_data_parallel_rank()
16 self.dp_comm_stream = torch.cuda.Stream(device=device, priority=-1)
17 self.torch_optim_comp_stream = torch.cuda.default_stream(device=device)
18 self.backward_ready_event = torch.cuda.Event(enable_timing=self.enable_tidy_profiling, blocking=False)
19 self.allreduce_gradients_start_event = torch.cuda.Event(enable_timing=self.enable_tidy_profiling, blocking=False)
20 self.allreduce_grad_ready_event = torch.cuda.Event(enable_timing=self.enable_tidy_profiling, blocking=False)
21 self.optimizer_step_ready_event = torch.cuda.Event(enable_timing=self.enable_tidy_profiling, blocking=False)
22
23 self.module = module
24 num_paras, element_size = self._compute_total_para_num()
25 print("Total number of parameters: {}, element size: {}, total size {} MB."
26 .format(num_paras, element_size, num_paras * element_size // 1024 // 1024))
27
28 if self.flatten:
29 self.flatten_para = flatten_params(self.module.parameters())
30 print("Flattened parameter number: {}, element size: {}."
31 .format(self.flatten_para.data.numel(), self.flatten_para.data.element_size()))
32 print("Flattened parameter grad number: {}, element size: {}."
33 .format(self.flatten_para.grad.numel(), self.flatten_para.grad.element_size()))
34
35 assert optimizer is not None
36 self.optimizer = optimizer
37
38 if self.enable_tidy_profiling:
39 self.global_rank = args.rank
40 self.init_event = None
41 self.init_time_stamp = None
42 if self.flatten:
43 self.allreduce_gradients_start_event = torch.cuda.Event(enable_timing=True, blocking=False)
44 else:
45 self.allreduce_gradients_start_events = dict()
46 self.allreduce_gradients_end_events = dict()
47 for name, _ in self.module.named_parameters():
48 self.allreduce_gradients_start_events[name] = torch.cuda.Event(enable_timing=True, blocking=False)
49 self.allreduce_gradients_end_events[name] = torch.cuda.Event(enable_timing=True, blocking=False)
50
51 self.optimizer_step_start_event = torch.cuda.Event(enable_timing=self.enable_tidy_profiling,
52 blocking=False)
53
54 def _compute_total_para_num(self):
55 total_count = 0

Callers

nothing calls this directly

Calls 4

get_data_parallel_commFunction · 0.85
get_data_parallel_rankFunction · 0.85
flatten_paramsFunction · 0.85

Tested by

no test coverage detected