315 lines
9.5 KiB
Python
315 lines
9.5 KiB
Python
"""可视化模块:双栏对比图(GT 上、TDengine 检测下)。"""
|
||
|
||
import logging
|
||
from pathlib import Path
|
||
import matplotlib
|
||
matplotlib.use("Agg")
|
||
import matplotlib.pyplot as plt
|
||
import matplotlib.font_manager as fm
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
ALGO_COLORS = {
|
||
"ksigma": "#1f77b4",
|
||
"iqr": "#2ca02c",
|
||
"grubbs": "#9467bd",
|
||
"shesd": "#ff7f0e",
|
||
"lof": "#17becf",
|
||
"sample_ad_model": "#e377c2",
|
||
}
|
||
|
||
|
||
def _find_cjk_font():
|
||
"""查找可用 CJK 字体。"""
|
||
candidates = ["Noto Serif CJK SC", "Noto Sans CJK SC", "WenQuanYi Micro Hei",
|
||
"DejaVu Sans", "sans-serif"]
|
||
available = {f.name for f in fm.fontManager.ttflist}
|
||
for name in candidates:
|
||
if name in available:
|
||
return name
|
||
return "sans-serif"
|
||
|
||
|
||
def _compute_iou(gt_windows, det_windows):
|
||
"""计算 GT 与检测窗口的 IoU。
|
||
|
||
窗口格式: [(start_ts, end_ts), ...]
|
||
IoU = 交集时长 / 并集时长
|
||
"""
|
||
if not gt_windows or not det_windows:
|
||
return 0.0
|
||
|
||
def overlap(a_start, a_end, b_start, b_end):
|
||
o_start = max(a_start, b_start)
|
||
o_end = min(a_end, b_end)
|
||
return max(0, o_end - o_start)
|
||
|
||
intersection = 0.0
|
||
union = 0.0
|
||
|
||
for gw in gt_windows:
|
||
union += gw[1] - gw[0]
|
||
for dw in det_windows:
|
||
intersection += overlap(gw[0], gw[1], dw[0], dw[1])
|
||
|
||
for dw in det_windows:
|
||
union += dw[1] - dw[0]
|
||
|
||
if union == 0:
|
||
return 0.0
|
||
return min(intersection / union, 1.0)
|
||
|
||
|
||
def render_contrast(
|
||
timestamps: list[int],
|
||
values: list[float],
|
||
gt_labels: list[int],
|
||
detection_results: dict,
|
||
output_path: str,
|
||
title: str = "",
|
||
) -> str:
|
||
"""渲染双栏对比图。
|
||
|
||
Args:
|
||
timestamps: 毫秒时间戳列表
|
||
values: 数值列表
|
||
gt_labels: GT 标签(0/1)
|
||
detection_results: {algo: {"windows": [...]}} 来自 detect_all_algos
|
||
output_path: PNG 输出路径
|
||
title: 图表标题
|
||
|
||
Returns:
|
||
output_path
|
||
"""
|
||
font_name = _find_cjk_font()
|
||
plt.rcParams["font.family"] = font_name
|
||
plt.rcParams["font.size"] = 9
|
||
|
||
fig, (ax_gt, ax_det) = plt.subplots(2, 1, sharex=True, figsize=(14, 8))
|
||
|
||
t0 = timestamps[0]
|
||
hours = [(t - t0) / 3_600_000 for t in timestamps]
|
||
|
||
ax_gt.plot(hours, values, color="black", linewidth=0.6, alpha=0.8)
|
||
ax_gt.set_ylabel("Value")
|
||
ax_gt.set_title(f"Ground Truth — {title}" if title else "Ground Truth")
|
||
|
||
in_anomaly = False
|
||
anom_start = None
|
||
gt_windows = []
|
||
for i, label in enumerate(gt_labels):
|
||
if label == 1 and not in_anomaly:
|
||
in_anomaly = True
|
||
anom_start = hours[i]
|
||
elif label == 0 and in_anomaly:
|
||
in_anomaly = False
|
||
ax_gt.axvspan(anom_start, hours[i], alpha=0.25, color="red")
|
||
gt_windows.append((anom_start, hours[i]))
|
||
|
||
if in_anomaly:
|
||
ax_gt.axvspan(anom_start, hours[-1], alpha=0.25, color="red")
|
||
gt_windows.append((anom_start, hours[-1]))
|
||
|
||
ax_det.plot(hours, values, color="black", linewidth=0.6, alpha=0.3)
|
||
ax_det.set_ylabel("Value")
|
||
ax_det.set_xlabel("Time (hours)")
|
||
|
||
all_det_windows = []
|
||
for algo, result in detection_results.items():
|
||
if result.get("error"):
|
||
continue
|
||
color = ALGO_COLORS.get(algo, "#888888")
|
||
for wstart, wend in result.get("windows", []):
|
||
wstart_h = (wstart - timestamps[0]) / 3_600_000 if wstart > timestamps[0] else 0
|
||
wend_h = (wend - timestamps[0]) / 3_600_000
|
||
ax_det.axvspan(wstart_h, wend_h, alpha=0.2, color=color)
|
||
all_det_windows.append((wstart_h, wend_h))
|
||
|
||
legend_handles = []
|
||
for algo in detection_results:
|
||
if algo in ALGO_COLORS and not detection_results[algo].get("error"):
|
||
from matplotlib.patches import Patch
|
||
legend_handles.append(Patch(color=ALGO_COLORS[algo], alpha=0.5, label=algo))
|
||
if legend_handles:
|
||
ax_det.legend(handles=legend_handles, loc="upper right", fontsize=7)
|
||
|
||
ax_det.set_title(f"TDengine Detection — {title}" if title else "TDengine Detection (6 algos)")
|
||
|
||
iou = _compute_iou(gt_windows, all_det_windows)
|
||
total_det = len(all_det_windows)
|
||
stats_text = f"GT windows: {len(gt_windows)}\nDetected: {total_det}\nIoU: {iou:.3f}"
|
||
ax_det.text(
|
||
0.98, 0.97, stats_text,
|
||
transform=ax_det.transAxes,
|
||
verticalalignment="top",
|
||
horizontalalignment="right",
|
||
bbox=dict(boxstyle="round", facecolor="white", alpha=0.8),
|
||
fontsize=8,
|
||
)
|
||
|
||
plt.tight_layout()
|
||
Path(output_path).parent.mkdir(parents=True, exist_ok=True)
|
||
fig.savefig(output_path, dpi=150, bbox_inches="tight")
|
||
plt.close(fig)
|
||
logger.info("chart saved: %s", output_path)
|
||
return output_path
|
||
|
||
|
||
def render_algo(
|
||
timestamps: list[int],
|
||
values: list[float],
|
||
gt_labels: list[int],
|
||
det_windows: list[list[int]],
|
||
algo: str,
|
||
output_path: str,
|
||
title: str = "",
|
||
) -> str:
|
||
"""单算法渲染:GT(红)+ 该算法检测窗口(算法色)→ PNG。
|
||
|
||
Args:
|
||
timestamps: 毫秒时间戳列表
|
||
values: 数值列表
|
||
gt_labels: GT 标签(0/1)
|
||
det_windows: [[start_ts, end_ts], ...] 该算法检测窗口
|
||
algo: 算法名称
|
||
output_path: PNG 输出路径
|
||
title: 图表标题
|
||
|
||
Returns:
|
||
output_path
|
||
"""
|
||
font_name = _find_cjk_font()
|
||
plt.rcParams["font.family"] = font_name
|
||
plt.rcParams["font.size"] = 9
|
||
|
||
fig, ax = plt.subplots(1, 1, figsize=(14, 5))
|
||
|
||
t0 = timestamps[0]
|
||
hours = [(t - t0) / 3_600_000 for t in timestamps]
|
||
|
||
# 原始序列
|
||
ax.plot(hours, values, color="black", linewidth=0.6, alpha=0.8)
|
||
|
||
# GT 红色 axvspan
|
||
in_anomaly = False
|
||
anom_start = None
|
||
gt_windows = []
|
||
for i, label in enumerate(gt_labels):
|
||
if label == 1 and not in_anomaly:
|
||
in_anomaly = True
|
||
anom_start = hours[i]
|
||
elif label == 0 and in_anomaly:
|
||
in_anomaly = False
|
||
ax.axvspan(anom_start, hours[i], alpha=0.15, color="red", label="GT" if len(gt_windows) == 0 else None)
|
||
gt_windows.append((anom_start, hours[i]))
|
||
if in_anomaly:
|
||
ax.axvspan(anom_start, hours[-1], alpha=0.15, color="red", label="GT" if len(gt_windows) == 0 else None)
|
||
gt_windows.append((anom_start, hours[-1]))
|
||
|
||
# 该算法检测窗口
|
||
color = ALGO_COLORS.get(algo, "#888888")
|
||
det_hours = []
|
||
for i, (wstart, wend) in enumerate(det_windows):
|
||
ws_h = (wstart - t0) / 3_600_000
|
||
we_h = (wend - t0) / 3_600_000
|
||
ax.axvspan(ws_h, we_h, alpha=0.2, color=color, label=algo if i == 0 else None)
|
||
det_hours.append((ws_h, we_h))
|
||
|
||
# IoU + 统计
|
||
iou = _compute_iou(gt_windows, det_hours)
|
||
stats_text = f"GT: {len(gt_windows)} Det: {len(det_windows)} IoU: {iou:.3f}"
|
||
ax.set_title(f"{title} — {stats_text}" if title else stats_text)
|
||
ax.set_ylabel("Value")
|
||
ax.set_xlabel("Time (hours)")
|
||
|
||
if gt_windows or det_windows:
|
||
ax.legend(loc="upper right", fontsize=7)
|
||
|
||
plt.tight_layout()
|
||
Path(output_path).parent.mkdir(parents=True, exist_ok=True)
|
||
fig.savefig(output_path, dpi=150, bbox_inches="tight")
|
||
plt.close(fig)
|
||
logger.info("render_algo saved: %s", output_path)
|
||
return output_path
|
||
|
||
|
||
def montage_pngs(
|
||
image_paths: list[str],
|
||
output_path: str,
|
||
cols: int = 2,
|
||
title: str = "",
|
||
) -> str:
|
||
"""将多张等宽 PNG 按网格拼成一张对比图。
|
||
|
||
Args:
|
||
image_paths: PNG 文件路径列表
|
||
output_path: 拼图输出路径
|
||
cols: 列数
|
||
title: 总标题
|
||
|
||
Returns:
|
||
output_path
|
||
"""
|
||
from PIL import Image, ImageDraw, ImageFont
|
||
|
||
if not image_paths:
|
||
raise ValueError("image_paths 为空")
|
||
|
||
images = []
|
||
for p in image_paths:
|
||
img = Image.open(p)
|
||
images.append(img)
|
||
|
||
# 统一宽度为第一张图的宽度
|
||
target_w = images[0].width
|
||
resized = []
|
||
for img in images:
|
||
if img.width != target_w:
|
||
ratio = target_w / img.width
|
||
new_h = int(img.height * ratio)
|
||
img = img.resize((target_w, new_h), Image.LANCZOS)
|
||
resized.append(img)
|
||
|
||
n = len(resized)
|
||
rows_count = (n + cols - 1) // cols
|
||
|
||
# 计算每行高度
|
||
row_heights = []
|
||
for r in range(rows_count):
|
||
max_h = max(resized[r * cols + c].height for c in range(cols) if r * cols + c < n)
|
||
row_heights.append(max_h)
|
||
|
||
padding = 10
|
||
title_h = 40 if title else 0
|
||
total_w = target_w * cols + padding * (cols + 1)
|
||
total_h = sum(row_heights) + padding * (rows_count + 1) + title_h
|
||
|
||
canvas = Image.new("RGB", (total_w, total_h), "white")
|
||
draw = ImageDraw.Draw(canvas)
|
||
|
||
if title:
|
||
try:
|
||
font = ImageFont.truetype("/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf", 24)
|
||
except Exception:
|
||
font = ImageFont.load_default()
|
||
draw.text((total_w // 2, 15), title, fill="black", anchor="mm", font=font)
|
||
|
||
y_offset = title_h + padding
|
||
for r in range(rows_count):
|
||
x_offset = padding
|
||
for c in range(cols):
|
||
idx = r * cols + c
|
||
if idx >= n:
|
||
break
|
||
img = resized[idx]
|
||
# 垂直居中
|
||
y_pad = (row_heights[r] - img.height) // 2
|
||
canvas.paste(img, (x_offset, y_offset + y_pad))
|
||
x_offset += target_w + padding
|
||
y_offset += row_heights[r] + padding
|
||
|
||
Path(output_path).parent.mkdir(parents=True, exist_ok=True)
|
||
canvas.save(output_path)
|
||
logger.info("montage saved: %s (%d images, %dx%d)", output_path, n, cols, rows_count)
|
||
return output_path
|