Spaces:
Running on Zero
Running on Zero
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() |