CoolFace
Apppublic

blanchon/Metric3D

sourceHugging Facegpl-3.0updated 2y agoView on Hugging Face
0likes
comm.py323 linesDownload Raw Back to utils
1import importlib 2import torch3import torch.distributed as dist4from .avg_meter import AverageMeter5from collections import defaultdict, OrderedDict6import os7import socket8from mmcv.utils import collect_env as collect_base_env9try:10    from mmcv.utils import get_git_hash11except:12    from mmengine.utils import get_git_hash13#import mono.mmseg as mmseg14# import mmseg15import time16import datetime17import logging18 19 20def main_process() -> bool:21    return get_rank() == 022    #return not cfg.distributed or \23    #       (cfg.distributed and cfg.local_rank == 0)24 25def get_world_size() -> int:26    if not dist.is_available():27        return 128    if not dist.is_initialized():29        return 130    return dist.get_world_size()31 32def get_rank() -> int:33    if not dist.is_available():34        return 035    if not dist.is_initialized():36        return 037    return dist.get_rank()38 39def _find_free_port():40    # refer to https://github.com/facebookresearch/detectron2/blob/main/detectron2/engine/launch.py # noqa: E50141    sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)42    # Binding to port 0 will cause the OS to find an available port for us43    sock.bind(('', 0))44    port = sock.getsockname()[1]45    sock.close()46    # NOTE: there is still a chance the port could be taken by other processes.47    return port 48 49def _is_free_port(port):50    ips = socket.gethostbyname_ex(socket.gethostname())[-1]51    ips.append('localhost')52    with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:53        return all(s.connect_ex((ip, port)) != 0 for ip in ips)54 55 56# def collect_env():57#     """Collect the information of the running environments."""58#     env_info = collect_base_env()59#     env_info['MMSegmentation'] = f'{mmseg.__version__}+{get_git_hash()[:7]}'60 61#     return env_info62 63def init_env(launcher, cfg):64    """Initialize distributed training environment.65    If argument ``cfg.dist_params.dist_url`` is specified as 'env://', then the master port will be system66    environment variable ``MASTER_PORT``. If ``MASTER_PORT`` is not in system67    environment variable, then a default port ``29500`` will be used.68    """69    if launcher == 'slurm':70        _init_dist_slurm(cfg)71    elif launcher == 'ror':72        _init_dist_ror(cfg)73    elif launcher == 'None':74        _init_none_dist(cfg)75    else:76        raise RuntimeError(f'{cfg.launcher} has not been supported!')77 78def _init_none_dist(cfg):79    cfg.dist_params.num_gpus_per_node = 180    cfg.dist_params.world_size = 181    cfg.dist_params.nnodes = 182    cfg.dist_params.node_rank = 083    cfg.dist_params.global_rank = 084    cfg.dist_params.local_rank = 085    os.environ["WORLD_SIZE"] = str(1)86 87def _init_dist_ror(cfg):88    from ac2.ror.comm import get_local_rank, get_world_rank, get_local_size, get_node_rank, get_world_size89    cfg.dist_params.num_gpus_per_node = get_local_size()90    cfg.dist_params.world_size = get_world_size()91    cfg.dist_params.nnodes = (get_world_size()) // (get_local_size())92    cfg.dist_params.node_rank = get_node_rank()93    cfg.dist_params.global_rank = get_world_rank()94    cfg.dist_params.local_rank = get_local_rank()95    os.environ["WORLD_SIZE"] = str(get_world_size())96 97 98def _init_dist_slurm(cfg):99    if 'NNODES' not in os.environ:100        os.environ['NNODES'] = str(cfg.dist_params.nnodes)101    if 'NODE_RANK' not in os.environ:102        os.environ['NODE_RANK'] = str(cfg.dist_params.node_rank)103 104    #cfg.dist_params.105    num_gpus = torch.cuda.device_count()106    world_size = int(os.environ['NNODES']) * num_gpus107    os.environ['WORLD_SIZE'] = str(world_size)108 109    # config port110    if 'MASTER_PORT' in os.environ:111        master_port = str(os.environ['MASTER_PORT'])  # use MASTER_PORT in the environment variable112    else:113        # if torch.distributed default port(29500) is available114        # then use it, else find a free port115        if _is_free_port(16500):116            master_port = '16500'117        else:118            master_port = str(_find_free_port())119        os.environ['MASTER_PORT'] = master_port120 121    # config addr122    if 'MASTER_ADDR' in os.environ:123        master_addr = str(os.environ['MASTER_PORT'])  # use MASTER_PORT in the environment variable124    # elif cfg.dist_params.dist_url is not None:125    #     master_addr = ':'.join(str(cfg.dist_params.dist_url).split(':')[:2])126    else:127        master_addr = '127.0.0.1' #'tcp://127.0.0.1'128        os.environ['MASTER_ADDR'] = master_addr129 130    # set dist_url to 'env://' 131    cfg.dist_params.dist_url =  'env://' #f"{master_addr}:{master_port}"132    133    cfg.dist_params.num_gpus_per_node = num_gpus134    cfg.dist_params.world_size = world_size135    cfg.dist_params.nnodes = int(os.environ['NNODES'])136    cfg.dist_params.node_rank = int(os.environ['NODE_RANK'])137        138    # if int(os.environ['NNODES']) > 1 and cfg.dist_params.dist_url.startswith("file://"):139    #     raise Warning("file:// is not a reliable init_method in multi-machine jobs. Prefer tcp://")140        141 142def get_func(func_name):143    """144        Helper to return a function object by name. func_name must identify 145        a function in this module or the path to a function relative to the base146        module.147        @ func_name: function name.148    """149    if func_name == '':150        return None151    try:152        parts = func_name.split('.')153        # Refers to a function in this module154        if len(parts) == 1:155            return globals()[parts[0]]156        # Otherwise, assume we're referencing a module under modeling157        module_name = '.'.join(parts[:-1])158        module = importlib.import_module(module_name)159        return getattr(module, parts[-1])160    except:161        raise RuntimeError(f'Failed to find function: {func_name}')162 163class Timer(object):164    """A simple timer."""165 166    def __init__(self):167        self.reset()168 169    def tic(self):170        # using time.time instead of time.clock because time time.clock171        # does not normalize for multithreading172        self.start_time = time.time()173 174    def toc(self, average=True):175        self.diff = time.time() - self.start_time176        self.total_time += self.diff177        self.calls += 1178        self.average_time = self.total_time / self.calls179        if average:180            return self.average_time181        else:182            return self.diff183 184    def reset(self):185        self.total_time = 0.186        self.calls = 0187        self.start_time = 0.188        self.diff = 0.189        self.average_time = 0.190 191class TrainingStats(object):192    """Track vital training statistics."""193    def __init__(self, log_period, tensorboard_logger=None):194        self.log_period = log_period195        self.tblogger = tensorboard_logger196        self.tb_ignored_keys = ['iter', 'eta', 'epoch', 'time']197        self.iter_timer = Timer()198        # Window size for smoothing tracked values (with median filtering)199        self.filter_size = log_period200        def create_smoothed_value():201            return AverageMeter()202        self.smoothed_losses = defaultdict(create_smoothed_value)203        #self.smoothed_metrics = defaultdict(create_smoothed_value)204        #self.smoothed_total_loss = AverageMeter()205 206 207    def IterTic(self):208        self.iter_timer.tic()209 210    def IterToc(self):211        return self.iter_timer.toc(average=False)212 213    def reset_iter_time(self):214        self.iter_timer.reset()215 216    def update_iter_stats(self, losses_dict):217        """Update tracked iteration statistics."""218        for k, v in losses_dict.items():219            self.smoothed_losses[k].update(float(v), 1)220 221    def log_iter_stats(self, cur_iter, optimizer, max_iters, val_err={}):222        """Log the tracked statistics."""223        if (cur_iter % self.log_period == 0):224            stats = self.get_stats(cur_iter, optimizer, max_iters, val_err)225            log_stats(stats)226            if self.tblogger:227                self.tb_log_stats(stats, cur_iter)228            for k, v in self.smoothed_losses.items():229                v.reset()230 231    def tb_log_stats(self, stats, cur_iter):232        """Log the tracked statistics to tensorboard"""233        for k in stats:234            # ignore some logs235            if k not in self.tb_ignored_keys:236                v = stats[k]237                if isinstance(v, dict):238                    self.tb_log_stats(v, cur_iter)239                else:240                    self.tblogger.add_scalar(k, v, cur_iter)241 242 243    def get_stats(self, cur_iter, optimizer, max_iters, val_err = {}):244        eta_seconds = self.iter_timer.average_time * (max_iters - cur_iter)245 246        eta = str(datetime.timedelta(seconds=int(eta_seconds)))247        stats = OrderedDict(248            iter=cur_iter,  # 1-indexed249            time=self.iter_timer.average_time,250            eta=eta,251        )252        optimizer_state_dict = optimizer.state_dict()253        lr = {}254        for i in range(len(optimizer_state_dict['param_groups'])):255            lr_name = 'group%d_lr' % i256            lr[lr_name] = optimizer_state_dict['param_groups'][i]['lr']257 258        stats['lr'] = OrderedDict(lr)259        for k, v in self.smoothed_losses.items():260            stats[k] = v.avg261 262        stats['val_err'] = OrderedDict(val_err)263        stats['max_iters'] = max_iters264        return stats265 266 267def reduce_dict(input_dict, average=True):268    """269    Reduce the values in the dictionary from all processes so that process with rank270    0 has the reduced results.271    Args:272        @input_dict (dict): inputs to be reduced. All the values must be scalar CUDA Tensor.273        @average (bool): whether to do average or sum274    Returns:275        a dict with the same keys as input_dict, after reduction.276    """277    world_size = get_world_size()278    if world_size < 2:279        return input_dict280    with torch.no_grad():281        names = []282        values = []283        # sort the keys so that they are consistent across processes284        for k in sorted(input_dict.keys()):285            names.append(k)286            values.append(input_dict[k])287        values = torch.stack(values, dim=0)288        dist.reduce(values, dst=0)289        if dist.get_rank() == 0 and average:290            # only main process gets accumulated, so only divide by291            # world_size in this case292            values /= world_size293        reduced_dict = {k: v for k, v in zip(names, values)}294    return reduced_dict295 296 297def log_stats(stats):298    logger = logging.getLogger()299    """Log training statistics to terminal"""300    lines = "[Step %d/%d]\n" % (301            stats['iter'], stats['max_iters'])302 303    lines += "\t\tloss: %.3f,    time: %.6f,    eta: %s\n" % (304        stats['total_loss'], stats['time'], stats['eta'])305 306    # log loss307    lines += "\t\t" 308    for k, v in stats.items():309        if 'loss' in k.lower() and 'total_loss' not in k.lower():310            lines += "%s: %.3f" % (k, v)  + ",  "311    lines = lines[:-3]312    lines += '\n'313 314    # validate criteria315    lines += "\t\tlast val err:" + ",  ".join("%s: %.6f" % (k, v) for k, v in stats['val_err'].items()) + ", "316    lines += '\n'317 318    # lr in different groups319    lines += "\t\t" +  ",  ".join("%s: %.8f" % (k, v) for k, v in stats['lr'].items())320    lines += '\n'321    logger.info(lines[:-1])  # remove last new linen_pxl322 323