blanchon/Metric3D
0
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 