feat: add render-algo subcommand for single-algorithm GT+detection PNG

This commit is contained in:
张宗平
2026-06-11 21:19:09 +08:00
parent 45fc062627
commit ab31b11fb5
3 changed files with 226 additions and 0 deletions
+159
View File
@@ -153,3 +153,162 @@ def render_contrast(
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