Coverage for hopwise/utils/logger.py: 82%
82 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
1# @Time : 2020/8/7
2# @Author : Zihan Lin
3# @Email : linzihan.super@foxmail.com
5# UPDATE
6# @Time : 2021/3/7
7# @Author : Jiawei Guan
8# @Email : guanjw@ruc.edu.cn
10# UPDATE:
11# @Time : 2022/07/10
12# @Author : Junjie Zhang
13# @Email : zjj001128@163.com
15"""hopwise.utils.logger
16###############################
17"""
19import hashlib
20import logging
21import os
22import re
23from functools import partial
25import colorama
26import colorlog
27import tqdm
28import tqdm.rich
30from hopwise.utils.utils import ensure_dir, get_local_time
32_progress_bar = None
35class ProgressBar:
36 def __init__(self, progress_bar_rich=True):
37 """
38 Initialize the progress bar with the configuration settings.
39 """
40 progress_bar = tqdm.rich.tqdm if progress_bar_rich else tqdm.tqdm
41 self.progress_bar = partial(progress_bar, disable=os.environ.get("DISABLE_TQDM", False)) # noqa: PLW1508
43 def __call__(self, *args, **kwargs):
44 return self.progress_bar(*args, **kwargs)
47def progress_bar(*args, **kwargs):
48 if _progress_bar is None:
49 return ProgressBar({"progress_bar_rich": True})(*args, **kwargs)
50 return _progress_bar(*args, **kwargs)
53log_colors_config = {
54 "DEBUG": "cyan",
55 "WARNING": "yellow",
56 "ERROR": "red",
57 "CRITICAL": "red",
58}
61class RemoveColorFilter(logging.Filter):
62 def filter(self, record):
63 if record:
64 ansi_escape = re.compile(r"\x1B(?:[@-Z\\-_]|\[[0-?]*[ -/]*[@-~])")
65 record.msg = ansi_escape.sub("", str(record.msg))
66 return True
69def set_color(log, color, highlight=True, progress=False):
70 if not progress or _progress_bar.progress_bar.func is tqdm.tqdm:
71 color_set = ["black", "red", "green", "yellow", "blue", "magenta", "cyan", "white"]
72 try:
73 index = color_set.index(color)
74 except IndexError:
75 index = len(color_set) - 1
76 prev_log = "\033["
77 if highlight:
78 prev_log += "1;3"
79 else:
80 prev_log += "0;3"
81 prev_log += str(index) + "m"
82 return prev_log + log + "\033[0m"
83 elif _progress_bar.progress_bar.func is tqdm.rich.tqdm:
84 return f"[{color}]{log}[/{color}]"
87def init_logger(config):
88 """A logger that can show a message on standard output and write it into the
89 file named `filename` simultaneously.
90 All the message that you want to log MUST be str.
92 Args:
93 config (Config): An instance object of Config, used to record parameter information.
95 Example:
96 >>> logger = logging.getLogger(config)
97 >>> logger.debug(train_state)
98 >>> logger.info(train_result)
99 """
100 colorama.init(autoreset=True)
101 LOGROOT = "./log/"
102 dir_name = os.path.dirname(LOGROOT)
103 ensure_dir(dir_name)
104 model_name = os.path.join(dir_name, config["model"])
105 ensure_dir(model_name)
106 if config["proc_title"] is None:
107 config_str = "".join([str(key) for key in config.final_config_dict.values()])
108 md5 = hashlib.md5(config_str.encode(encoding="utf-8")).hexdigest()[:6]
109 logfilename = "{}/{}-{}-{}-{}.log".format(
110 config["model"], config["model"], config["dataset"], get_local_time(), md5
111 )
112 else:
113 logfilename = "{}/{}-{}.log".format(config["model"], config["proc_title"], get_local_time())
115 logfilepath = os.path.join(LOGROOT, logfilename)
117 filefmt = "%(asctime)-15s %(levelname)s %(message)s"
118 filedatefmt = "%a %d %b %Y %H:%M:%S"
119 fileformatter = logging.Formatter(filefmt, filedatefmt)
121 sfmt = "%(log_color)s%(asctime)-15s %(levelname)s %(message)s"
122 sdatefmt = "%d %b %H:%M"
123 sformatter = colorlog.ColoredFormatter(sfmt, sdatefmt, log_colors=log_colors_config)
124 if config["state"] is None or config["state"].lower() == "info":
125 level = logging.INFO
126 elif config["state"].lower() == "debug":
127 level = logging.DEBUG
128 elif config["state"].lower() == "error":
129 level = logging.ERROR
130 elif config["state"].lower() == "warning":
131 level = logging.WARNING
132 elif config["state"].lower() == "critical":
133 level = logging.CRITICAL
134 else:
135 level = logging.INFO
137 fh = logging.FileHandler(logfilepath)
138 fh.setLevel(level)
139 fh.setFormatter(fileformatter)
140 remove_color_filter = RemoveColorFilter()
141 fh.addFilter(remove_color_filter)
143 sh = logging.StreamHandler()
144 sh.setLevel(level)
145 sh.setFormatter(sformatter)
147 logging.basicConfig(level=level, handlers=[sh, fh])