File size: 4,147 Bytes
1e05592 | 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 143 144 145 146 | """
日志系统
提供统一的日志接口,支持文件和终端输出
"""
import logging
import sys
from pathlib import Path
from datetime import datetime
import json
class ColoredFormatter(logging.Formatter):
"""带颜色的日志格式化器"""
COLORS = {
'DEBUG': '\033[36m', # 青色
'INFO': '\033[32m', # 绿色
'WARNING': '\033[33m', # 黄色
'ERROR': '\033[31m', # 红色
'CRITICAL': '\033[35m', # 紫色
}
RESET = '\033[0m'
def format(self, record):
log_color = self.COLORS.get(record.levelname, self.RESET)
record.levelname = f"{log_color}{record.levelname}{self.RESET}"
return super().format(record)
def setup_logger(
name: str,
log_file: str = None,
level: int = logging.INFO,
console: bool = True
):
"""
设置logger
Args:
name: logger名称
log_file: 日志文件路径(可选)
level: 日志级别
console: 是否输出到控制台
Returns:
logger实例
"""
logger = logging.getLogger(name)
logger.setLevel(level)
logger.handlers.clear() # 清除已有的handlers
# 格式
fmt = '%(asctime)s - %(name)s - %(levelname)s - %(message)s'
datefmt = '%Y-%m-%d %H:%M:%S'
# 控制台handler
if console:
console_handler = logging.StreamHandler(sys.stdout)
console_handler.setLevel(level)
console_formatter = ColoredFormatter(fmt, datefmt=datefmt)
console_handler.setFormatter(console_formatter)
logger.addHandler(console_handler)
# 文件handler
if log_file:
log_path = Path(log_file)
log_path.parent.mkdir(parents=True, exist_ok=True)
file_handler = logging.FileHandler(log_file, encoding='utf-8')
file_handler.setLevel(level)
file_formatter = logging.Formatter(fmt, datefmt=datefmt)
file_handler.setFormatter(file_formatter)
logger.addHandler(file_handler)
return logger
class MetricsLogger:
"""
指标记录器
记录训练/验证指标到JSON文件
"""
def __init__(self, log_dir: str, exp_name: str):
self.log_dir = Path(log_dir)
self.log_dir.mkdir(parents=True, exist_ok=True)
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
self.log_file = self.log_dir / f"{exp_name}_{timestamp}.json"
self.metrics = {
'train': [],
'val': [],
'config': {}
}
def log_config(self, config: dict):
"""记录配置"""
self.metrics['config'] = config
self._save()
def log_train(self, step: int, metrics: dict):
"""记录训练指标"""
metrics['step'] = step
metrics['timestamp'] = datetime.now().isoformat()
self.metrics['train'].append(metrics)
self._save()
def log_val(self, epoch: int, metrics: dict):
"""记录验证指标"""
metrics['epoch'] = epoch
metrics['timestamp'] = datetime.now().isoformat()
self.metrics['val'].append(metrics)
self._save()
def _save(self):
"""保存到文件"""
with open(self.log_file, 'w', encoding='utf-8') as f:
json.dump(self.metrics, f, indent=2, ensure_ascii=False)
def get_best_metric(self, metric_name: str, mode: str = 'min'):
"""获取最佳指标"""
if not self.metrics['val']:
return None
values = [m[metric_name] for m in self.metrics['val'] if metric_name in m]
if not values:
return None
if mode == 'min':
best_val = min(values)
best_epoch = values.index(best_val)
else:
best_val = max(values)
best_epoch = values.index(best_val)
return {
'value': best_val,
'epoch': self.metrics['val'][best_epoch]['epoch']
}
# 创建全局logger
def get_logger(name: str = "lkalert"):
"""获取或创建logger"""
return logging.getLogger(name) |