Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2"""3This file contains primitives for multi-gpu communication.4This is useful when doing distributed training.5"""6 7import functools8import numpy as np9import torch10import torch.distributed as dist11 12_LOCAL_PROCESS_GROUP = None13_MISSING_LOCAL_PG_ERROR = (14 "Local process group is not yet created! Please use detectron2's `launch()` "15 "to start processes and initialize pytorch process group. If you need to start "16 "processes in other ways, please call comm.create_local_process_group("17 "num_workers_per_machine) after calling torch.distributed.init_process_group()."18)19 20 21def get_world_size() -> int:22 if not dist.is_available():23 return 124 if not dist.is_initialized():25 return 126 return dist.get_world_size()27 28 29def get_rank() -> int:30 if not dist.is_available():31 return 032 if not dist.is_initialized():33 return 034 return dist.get_rank()35 36 37@functools.lru_cache()38def create_local_process_group(num_workers_per_machine: int) -> None:39 """40 Create a process group that contains ranks within the same machine.41 42 Detectron2's launch() in engine/launch.py will call this function. If you start43 workers without launch(), you'll have to also call this. Otherwise utilities44 like `get_local_rank()` will not work.45 46 This function contains a barrier. All processes must call it together.47 48 Args:49 num_workers_per_machine: the number of worker processes per machine. Typically50 the number of GPUs.51 """52 global _LOCAL_PROCESS_GROUP53 assert _LOCAL_PROCESS_GROUP is None54 assert get_world_size() % num_workers_per_machine == 055 num_machines = get_world_size() // num_workers_per_machine56 machine_rank = get_rank() // num_workers_per_machine57 for i in range(num_machines):58 ranks_on_i = list(range(i * num_workers_per_machine, (i + 1) * num_workers_per_machine))59 pg = dist.new_group(ranks_on_i)60 if i == machine_rank:61 _LOCAL_PROCESS_GROUP = pg62 63 64def get_local_process_group():65 """66 Returns:67 A torch process group which only includes processes that are on the same68 machine as the current process. This group can be useful for communication69 within a machine, e.g. a per-machine SyncBN.70 """71 assert _LOCAL_PROCESS_GROUP is not None, _MISSING_LOCAL_PG_ERROR72 return _LOCAL_PROCESS_GROUP73 74 75def get_local_rank() -> int:76 """77 Returns:78 The rank of the current process within the local (per-machine) process group.79 """80 if not dist.is_available():81 return 082 if not dist.is_initialized():83 return 084 assert _LOCAL_PROCESS_GROUP is not None, _MISSING_LOCAL_PG_ERROR85 return dist.get_rank(group=_LOCAL_PROCESS_GROUP)86 87 88def get_local_size() -> int:89 """90 Returns:91 The size of the per-machine process group,92 i.e. the number of processes per machine.93 """94 if not dist.is_available():95 return 196 if not dist.is_initialized():97 return 198 assert _LOCAL_PROCESS_GROUP is not None, _MISSING_LOCAL_PG_ERROR99 return dist.get_world_size(group=_LOCAL_PROCESS_GROUP)100 101 102def is_main_process() -> bool:103 return get_rank() == 0104 105 106def synchronize():107 """108 Helper function to synchronize (barrier) among all processes when109 using distributed training110 """111 if not dist.is_available():112 return113 if not dist.is_initialized():114 return115 world_size = dist.get_world_size()116 if world_size == 1:117 return118 if dist.get_backend() == dist.Backend.NCCL:119 # This argument is needed to avoid warnings.120 # It's valid only for NCCL backend.121 dist.barrier(device_ids=[torch.cuda.current_device()])122 else:123 dist.barrier()124 125 126@functools.lru_cache()127def _get_global_gloo_group():128 """129 Return a process group based on gloo backend, containing all the ranks130 The result is cached.131 """132 if dist.get_backend() == "nccl":133 return dist.new_group(backend="gloo")134 else:135 return dist.group.WORLD136 137 138def all_gather(data, group=None):139 """140 Run all_gather on arbitrary picklable data (not necessarily tensors).141 142 Args:143 data: any picklable object144 group: a torch process group. By default, will use a group which145 contains all ranks on gloo backend.146 147 Returns:148 list[data]: list of data gathered from each rank149 """150 if get_world_size() == 1:151 return [data]152 if group is None:153 group = _get_global_gloo_group() # use CPU group by default, to reduce GPU RAM usage.154 world_size = dist.get_world_size(group)155 if world_size == 1:156 return [data]157 158 output = [None for _ in range(world_size)]159 dist.all_gather_object(output, data, group=group)160 return output161 162 163def gather(data, dst=0, group=None):164 """165 Run gather on arbitrary picklable data (not necessarily tensors).166 167 Args:168 data: any picklable object169 dst (int): destination rank170 group: a torch process group. By default, will use a group which171 contains all ranks on gloo backend.172 173 Returns:174 list[data]: on dst, a list of data gathered from each rank. Otherwise,175 an empty list.176 """177 if get_world_size() == 1:178 return [data]179 if group is None:180 group = _get_global_gloo_group()181 world_size = dist.get_world_size(group=group)182 if world_size == 1:183 return [data]184 rank = dist.get_rank(group=group)185 186 if rank == dst:187 output = [None for _ in range(world_size)]188 dist.gather_object(data, output, dst=dst, group=group)189 return output190 else:191 dist.gather_object(data, None, dst=dst, group=group)192 return []193 194 195def shared_random_seed():196 """197 Returns:198 int: a random number that is the same across all workers.199 If workers need a shared RNG, they can use this shared seed to200 create one.201 202 All workers must call this function, otherwise it will deadlock.203 """204 ints = np.random.randint(2**31)205 all_ints = all_gather(ints)206 return all_ints[0]207 208 209def reduce_dict(input_dict, average=True):210 """211 Reduce the values in the dictionary from all processes so that process with rank212 0 has the reduced results.213 214 Args:215 input_dict (dict): inputs to be reduced. All the values must be scalar CUDA Tensor.216 average (bool): whether to do average or sum217 218 Returns:219 a dict with the same keys as input_dict, after reduction.220 """221 world_size = get_world_size()222 if world_size < 2:223 return input_dict224 with torch.no_grad():225 names = []226 values = []227 # sort the keys so that they are consistent across processes228 for k in sorted(input_dict.keys()):229 names.append(k)230 values.append(input_dict[k])231 values = torch.stack(values, dim=0)232 dist.reduce(values, dst=0)233 if dist.get_rank() == 0 and average:234 # only main process gets accumulated, so only divide by235 # world_size in this case236 values /= world_size237 reduced_dict = {k: v for k, v in zip(names, values)}238 return reduced_dict239 