CoolFace
Apppublic

faisalhr1997/codeformer

sourceHugging Faceupdated 4y agoView on Hugging Face
0likes
logger.py169 linesDownload Raw Back to utils
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