File size: 5,115 Bytes
2680bd5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
import logging
import json
import sys
import torch
from collections import defaultdict
from typing import Union, List, Optional
from . import dist as dist_utils
from .utils.mics import get_device

class _ExperimentLogger:
    setup_count = 0

    def setup(
        self,
        logger_name: str = 'ml_experiment',
        file_log_level: int = logging.INFO,
        stream_log_level: int = logging.INFO,
        log_file: Optional[str] = 'experiment.log',
        formatter_str: str = '%(asctime)s - [%(name)s] - [%(levelname)s] - %(message)s'
    ):
        if self.setup_count > 0:
            raise RuntimeError("Logger 已经被初始化过了,请不要重复调用 setup 方法。")

        self.logger = logging.getLogger(logger_name)
        self.logger.setLevel(logging.DEBUG)

        if not self.logger.handlers:
            formatter = logging.Formatter(formatter_str)

            stream_handler = logging.StreamHandler(sys.stdout)
            stream_handler.setLevel(stream_log_level)
            stream_handler.setFormatter(formatter)
            self.logger.addHandler(stream_handler)

            if log_file:
                file_handler = logging.FileHandler(log_file, mode='a', encoding='utf-8')
                file_handler.setLevel(file_log_level)
                file_handler.setFormatter(formatter)
                self.logger.addHandler(file_handler)

        self._metrics = {}
        self.info(f"Logger '{logger_name}' 初始化成功。文件日志级别:{file_log_level}, 控制台日志级别:{stream_log_level}, 日志文件:{log_file if log_file else '无'}")
        _ExperimentLogger.setup_count += 1

    def debug(self, msg: str):
        """记录一条DEBUG级别的日志。"""
        self.logger.debug(msg)

    def info(self, msg: str):
        """记录一条INFO级别的日志。"""
        self.logger.info(msg)

    def warning(self, msg: str):
        """记录一条WARNING级别的日志。"""
        self.logger.warning(msg)

    def error(self, msg: str):
        """记录一条ERROR级别的日志。"""
        self.logger.error(msg)

    def critical(self, msg: str):
        """记录一条CRITICAL级别的日志。"""
        self.logger.critical(msg)


    def register_metric(self, name: str):
        if name in self._metrics:
            self.warning(f"指标 '{name}' 已存在,无需重复注册。")
        else:
            self._metrics[name] = []
            self.debug(f"指标 '{name}' 注册成功。")

    def record_metric(self, name: str, value):
        if name not in self._metrics:
            # self.warning(f"指标 '{name}' 未注册,将自动注册并记录数值。")
            self._metrics[name] = []

        self._metrics[name].append(value)

    def get_local_average(self, name: str) -> float:
        if name not in self._metrics:
            self.error(f"查询平均值失败: 指标 '{name}' 不存在。")
            return 0.0

        values = self._metrics[name]
        if not values:
            self.warning(f"指标 '{name}' 尚未记录任何数据,平均值为0.0。")
            return 0.0

        return sum(values) / len(values)

    def get_average(self, name: str) -> float:
        if not dist_utils.is_dist_avail_and_initialized():
            return self.get_local_average(name)

        local_avg = self.get_local_average(name)
        local_tensor = torch.tensor(local_avg, device=get_device())

        global_avg_tensor = dist_utils.reduce_tensor(local_tensor)
        return global_avg_tensor.item()

    def clear_metrics(self, names: Optional[List[str]] = None):
        if names is None:
            for name in self._metrics:
                self._metrics[name].clear()
        else:
            for name in names:
                if name in self._metrics:
                    self._metrics[name].clear()
                else:
                    self.warning(f"尝试清空失败: 指标 '{name}' 不存在。")

    def clear_metric(self, name: str):
        if name in self._metrics:
            self._metrics[name].clear()
        else:
            self.warning(f"尝试清空失败: 指标 '{name}' 不存在。")

    def get_all_metrics(self) -> dict:
        return self._metrics

    def save_metrics(self, filepath: str = 'metrics_summary.json'):
        self.info(f"准备将指标平均值摘要保存到 '{filepath}'...")

        average_metrics = {}

        for name in self._metrics.keys():
            average_value = self.get_average(name)
            average_metrics[name] = round(average_value, 6)

        self.info(f"计算出的平均值摘要: {average_metrics}")

        try:
            with open(filepath, 'w', encoding='utf-8') as f:
                json.dump(average_metrics, f, ensure_ascii=False, indent=4)
            self.info(f"指标平均值摘要已成功保存到 '{filepath}'。")
        except IOError as e:
            self.error(f"保存指标摘要到 '{filepath}' 时发生IO错误: {e}")
        except Exception as e:
            self.error(f"保存指标摘要时发生未知错误: {e}")

mylogger = _ExperimentLogger()