Team Ai
Apppublic

radames/Text2Human-API

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
logger.py113 linesDownload Raw Back to utils
1import datetime2import logging3import time4 5 6class MessageLogger():7    """Message logger for printing.8 9    Args:10        opt (dict): Config. It contains the following keys:11            name (str): Exp name.12            logger (dict): Contains 'print_freq' (str) for logger interval.13            train (dict): Contains 'niter' (int) for total iters.14            use_tb_logger (bool): Use tensorboard logger.15        start_iter (int): Start iter. Default: 1.16        tb_logger (obj:`tb_logger`): Tensorboard logger. Default: None.17    """18 19    def __init__(self, opt, start_iter=1, tb_logger=None):20        self.exp_name = opt['name']21        self.interval = opt['print_freq']22        self.start_iter = start_iter23        self.max_iters = opt['max_iters']24        self.use_tb_logger = opt['use_tb_logger']25        self.tb_logger = tb_logger26        self.start_time = time.time()27        self.logger = get_root_logger()28 29    def __call__(self, log_vars):30        """Format logging message.31 32        Args:33            log_vars (dict): It contains the following keys:34                epoch (int): Epoch number.35                iter (int): Current iter.36                lrs (list): List for learning rates.37 38                time (float): Iter time.39                data_time (float): Data time for each iter.40        """41        # epoch, iter, learning rates42        epoch = log_vars.pop('epoch')43        current_iter = log_vars.pop('iter')44        lrs = log_vars.pop('lrs')45 46        message = (f'[{self.exp_name[:5]}..][epoch:{epoch:3d}, '47                   f'iter:{current_iter:8,d}, lr:(')48        for v in lrs:49            message += f'{v:.3e},'50        message += ')] '51 52        # time and estimated time53        if 'time' in log_vars.keys():54            iter_time = log_vars.pop('time')55            data_time = log_vars.pop('data_time')56 57            total_time = time.time() - self.start_time58            time_sec_avg = total_time / (current_iter - self.start_iter + 1)59            eta_sec = time_sec_avg * (self.max_iters - current_iter - 1)60            eta_str = str(datetime.timedelta(seconds=int(eta_sec)))61            message += f'[eta: {eta_str}, '62            message += f'time: {iter_time:.3f}, data_time: {data_time:.3f}] '63 64        # other items, especially losses65        for k, v in log_vars.items():66            message += f'{k}: {v:.4e} '67            # tensorboard logger68            if self.use_tb_logger and 'debug' not in self.exp_name:69                self.tb_logger.add_scalar(k, v, current_iter)70 71        self.logger.info(message)72 73 74def init_tb_logger(log_dir):75    from torch.utils.tensorboard import SummaryWriter76    tb_logger = SummaryWriter(log_dir=log_dir)77    return tb_logger78 79 80def get_root_logger(logger_name='base', log_level=logging.INFO, log_file=None):81    """Get the root logger.82 83    The logger will be initialized if it has not been initialized. By default a84    StreamHandler will be added. If `log_file` is specified, a FileHandler will85    also be added.86 87    Args:88        logger_name (str): root logger name. Default: base.89        log_file (str | None): The log filename. If specified, a FileHandler90            will be added to the root logger.91        log_level (int): The root logger level. Note that only the process of92            rank 0 is affected, while other processes will set the level to93            "Error" and be silent most of the time.94 95    Returns:96        logging.Logger: The root logger.97    """98    logger = logging.getLogger(logger_name)99    # if the logger has been initialized, just return it100    if logger.hasHandlers():101        return logger102 103    format_str = '%(asctime)s.%(msecs)03d - %(levelname)s: %(message)s'104    logging.basicConfig(format=format_str, level=log_level)105 106    if log_file is not None:107        file_handler = logging.FileHandler(log_file, 'w')108        file_handler.setFormatter(logging.Formatter(format_str))109        file_handler.setLevel(log_level)110        logger.addHandler(file_handler)111 112    return logger113