Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2import logging3from datetime import timedelta4import torch5import torch.distributed as dist6import torch.multiprocessing as mp7 8from detectron2.utils import comm9 10__all__ = ["DEFAULT_TIMEOUT", "launch"]11 12DEFAULT_TIMEOUT = timedelta(minutes=30)13 14 15def _find_free_port():16 import socket17 18 sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)19 # Binding to port 0 will cause the OS to find an available port for us20 sock.bind(("", 0))21 port = sock.getsockname()[1]22 sock.close()23 # NOTE: there is still a chance the port could be taken by other processes.24 return port25 26 27def launch(28 main_func,29 # Should be num_processes_per_machine, but kept for compatibility.30 num_gpus_per_machine,31 num_machines=1,32 machine_rank=0,33 dist_url=None,34 args=(),35 timeout=DEFAULT_TIMEOUT,36):37 """38 Launch multi-process or distributed training.39 This function must be called on all machines involved in the training.40 It will spawn child processes (defined by ``num_gpus_per_machine``) on each machine.41 42 Args:43 main_func: a function that will be called by `main_func(*args)`44 num_gpus_per_machine (int): number of processes per machine. When45 using GPUs, this should be the number of GPUs.46 num_machines (int): the total number of machines47 machine_rank (int): the rank of this machine48 dist_url (str): url to connect to for distributed jobs, including protocol49 e.g. "tcp://127.0.0.1:8686".50 Can be set to "auto" to automatically select a free port on localhost51 timeout (timedelta): timeout of the distributed workers52 args (tuple): arguments passed to main_func53 """54 world_size = num_machines * num_gpus_per_machine55 if world_size > 1:56 # https://github.com/pytorch/pytorch/pull/1439157 # TODO prctl in spawned processes58 59 if dist_url == "auto":60 assert num_machines == 1, "dist_url=auto not supported in multi-machine jobs."61 port = _find_free_port()62 dist_url = f"tcp://127.0.0.1:{port}"63 if num_machines > 1 and dist_url.startswith("file://"):64 logger = logging.getLogger(__name__)65 logger.warning(66 "file:// is not a reliable init_method in multi-machine jobs. Prefer tcp://"67 )68 69 mp.start_processes(70 _distributed_worker,71 nprocs=num_gpus_per_machine,72 args=(73 main_func,74 world_size,75 num_gpus_per_machine,76 machine_rank,77 dist_url,78 args,79 timeout,80 ),81 daemon=False,82 )83 else:84 main_func(*args)85 86 87def _distributed_worker(88 local_rank,89 main_func,90 world_size,91 num_gpus_per_machine,92 machine_rank,93 dist_url,94 args,95 timeout=DEFAULT_TIMEOUT,96):97 has_gpu = torch.cuda.is_available()98 if has_gpu:99 assert num_gpus_per_machine <= torch.cuda.device_count()100 global_rank = machine_rank * num_gpus_per_machine + local_rank101 try:102 dist.init_process_group(103 backend="NCCL" if has_gpu else "GLOO",104 init_method=dist_url,105 world_size=world_size,106 rank=global_rank,107 timeout=timeout,108 )109 except Exception as e:110 logger = logging.getLogger(__name__)111 logger.error("Process group URL: {}".format(dist_url))112 raise e113 114 # Setup the local process group.115 comm.create_local_process_group(num_gpus_per_machine)116 if has_gpu:117 torch.cuda.set_device(local_rank)118 119 # synchronize is needed here to prevent a possible timeout after calling init_process_group120 # See: https://github.com/facebookresearch/maskrcnn-benchmark/issues/172121 comm.synchronize()122 123 main_func(*args)124 