faisalhr1997/codeformer
0
1import datetime2import logging3import time4 5from .dist_util import get_dist_info, master_only6 7initialized_logger = {}8 9 10class MessageLogger():11 """Message logger for printing.12 Args:13 opt (dict): Config. It contains the following keys:14 name (str): Exp name.15 logger (dict): Contains 'print_freq' (str) for logger interval.16 train (dict): Contains 'total_iter' (int) for total iters.17 use_tb_logger (bool): Use tensorboard logger.18 start_iter (int): Start iter. Default: 1.19 tb_logger (obj:`tb_logger`): Tensorboard logger. Default: None.20 """21 22 def __init__(self, opt, start_iter=1, tb_logger=None):23 self.exp_name = opt['name']24 self.interval = opt['logger']['print_freq']25 self.start_iter = start_iter26 self.max_iters = opt['train']['total_iter']27 self.use_tb_logger = opt['logger']['use_tb_logger']28 self.tb_logger = tb_logger29 self.start_time = time.time()30 self.logger = get_root_logger()31 32 @master_only33 def __call__(self, log_vars):34 """Format logging message.35 Args:36 log_vars (dict): It contains the following keys:37 epoch (int): Epoch number.38 iter (int): Current iter.39 lrs (list): List for learning rates.40 time (float): Iter time.41 data_time (float): Data time for each iter.42 """43 # epoch, iter, learning rates44 epoch = log_vars.pop('epoch')45 current_iter = log_vars.pop('iter')46 lrs = log_vars.pop('lrs')47 48 message = (f'[{self.exp_name[:5]}..][epoch:{epoch:3d}, ' f'iter:{current_iter:8,d}, lr:(')49 for v in lrs:50 message += f'{v:.3e},'51 message += ')] '52 53 # time and estimated time54 if 'time' in log_vars.keys():55 iter_time = log_vars.pop('time')56 data_time = log_vars.pop('data_time')57 58 total_time = time.time() - self.start_time59 time_sec_avg = total_time / (current_iter - self.start_iter + 1)60 eta_sec = time_sec_avg * (self.max_iters - current_iter - 1)61 eta_str = str(datetime.timedelta(seconds=int(eta_sec)))62 message += f'[eta: {eta_str}, '63 message += f'time (data): {iter_time:.3f} ({data_time:.3f})] '64 65 # other items, especially losses66 for k, v in log_vars.items():67 message += f'{k}: {v:.4e} '68 # tensorboard logger69 if self.use_tb_logger:70 if k.startswith('l_'):71 self.tb_logger.add_scalar(f'losses/{k}', v, current_iter)72 else:73 self.tb_logger.add_scalar(k, v, current_iter)74 self.logger.info(message)75 76 77@master_only78def init_tb_logger(log_dir):79 from torch.utils.tensorboard import SummaryWriter80 tb_logger = SummaryWriter(log_dir=log_dir)81 return tb_logger82 83 84@master_only85def init_wandb_logger(opt):86 """We now only use wandb to sync tensorboard log."""87 import wandb88 logger = logging.getLogger('basicsr')89 90 project = opt['logger']['wandb']['project']91 resume_id = opt['logger']['wandb'].get('resume_id')92 if resume_id:93 wandb_id = resume_id94 resume = 'allow'95 logger.warning(f'Resume wandb logger with id={wandb_id}.')96 else:97 wandb_id = wandb.util.generate_id()98 resume = 'never'99 100 wandb.init(id=wandb_id, resume=resume, name=opt['name'], config=opt, project=project, sync_tensorboard=True)101 102 logger.info(f'Use wandb logger with id={wandb_id}; project={project}.')103 104 105def get_root_logger(logger_name='basicsr', log_level=logging.INFO, log_file=None):106 """Get the root logger.107 The logger will be initialized if it has not been initialized. By default a108 StreamHandler will be added. If `log_file` is specified, a FileHandler will109 also be added.110 Args:111 logger_name (str): root logger name. Default: 'basicsr'.112 log_file (str | None): The log filename. If specified, a FileHandler113 will be added to the root logger.114 log_level (int): The root logger level. Note that only the process of115 rank 0 is affected, while other processes will set the level to116 "Error" and be silent most of the time.117 Returns:118 logging.Logger: The root logger.119 """120 logger = logging.getLogger(logger_name)121 # if the logger has been initialized, just return it122 if logger_name in initialized_logger:123 return logger124 125 format_str = '%(asctime)s %(levelname)s: %(message)s'126 stream_handler = logging.StreamHandler()127 stream_handler.setFormatter(logging.Formatter(format_str))128 logger.addHandler(stream_handler)129 logger.propagate = False130 rank, _ = get_dist_info()131 if rank != 0:132 logger.setLevel('ERROR')133 elif log_file is not None:134 logger.setLevel(log_level)135 # add file handler136 # file_handler = logging.FileHandler(log_file, 'w')137 file_handler = logging.FileHandler(log_file, 'a') #Shangchen: keep the previous log138 file_handler.setFormatter(logging.Formatter(format_str))139 file_handler.setLevel(log_level)140 logger.addHandler(file_handler)141 initialized_logger[logger_name] = True142 return logger143 144 145def get_env_info():146 """Get environment information.147 Currently, only log the software version.148 """149 import torch150 import torchvision151 152 from basicsr.version import __version__153 msg = r"""154 ____ _ _____ ____155 / __ ) ____ _ _____ (_)_____/ ___/ / __ \156 / __ |/ __ `// ___// // ___/\__ \ / /_/ /157 / /_/ // /_/ /(__ )/ // /__ ___/ // _, _/158 /_____/ \__,_//____//_/ \___//____//_/ |_|159 ______ __ __ __ __160 / ____/____ ____ ____/ / / / __ __ _____ / /__ / /161 / / __ / __ \ / __ \ / __ / / / / / / // ___// //_/ / /162 / /_/ // /_/ // /_/ // /_/ / / /___/ /_/ // /__ / /< /_/163 \____/ \____/ \____/ \____/ /_____/\____/ \___//_/|_| (_)164 """165 msg += ('\nVersion Information: '166 f'\n\tBasicSR: {__version__}'167 f'\n\tPyTorch: {torch.__version__}'168 f'\n\tTorchVision: {torchvision.__version__}')169 return msg