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

1# @Time : 2020/8/7 

2# @Author : Zihan Lin 

3# @Email : linzihan.super@foxmail.com 

4 

5# UPDATE 

6# @Time : 2021/3/7 

7# @Author : Jiawei Guan 

8# @Email : guanjw@ruc.edu.cn 

9 

10# UPDATE: 

11# @Time : 2022/07/10 

12# @Author : Junjie Zhang 

13# @Email : zjj001128@163.com 

14 

15"""hopwise.utils.logger 

16############################### 

17""" 

18 

19import hashlib 

20import logging 

21import os 

22import re 

23from functools import partial 

24 

25import colorama 

26import colorlog 

27import tqdm 

28import tqdm.rich 

29 

30from hopwise.utils.utils import ensure_dir, get_local_time 

31 

32_progress_bar = None 

33 

34 

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 

42 

43 def __call__(self, *args, **kwargs): 

44 return self.progress_bar(*args, **kwargs) 

45 

46 

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) 

51 

52 

53log_colors_config = { 

54 "DEBUG": "cyan", 

55 "WARNING": "yellow", 

56 "ERROR": "red", 

57 "CRITICAL": "red", 

58} 

59 

60 

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 

67 

68 

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}]" 

85 

86 

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. 

91 

92 Args: 

93 config (Config): An instance object of Config, used to record parameter information. 

94 

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()) 

114 

115 logfilepath = os.path.join(LOGROOT, logfilename) 

116 

117 filefmt = "%(asctime)-15s %(levelname)s %(message)s" 

118 filedatefmt = "%a %d %b %Y %H:%M:%S" 

119 fileformatter = logging.Formatter(filefmt, filedatefmt) 

120 

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 

136 

137 fh = logging.FileHandler(logfilepath) 

138 fh.setLevel(level) 

139 fh.setFormatter(fileformatter) 

140 remove_color_filter = RemoveColorFilter() 

141 fh.addFilter(remove_color_filter) 

142 

143 sh = logging.StreamHandler() 

144 sh.setLevel(level) 

145 sh.setFormatter(sformatter) 

146 

147 logging.basicConfig(level=level, handlers=[sh, fh])