Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
launch.py124 linesDownload Raw Back to engine
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