Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
comm.py239 linesDownload Raw Back to utils
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