diff --git a/scripts/run_compare.sh b/scripts/run_compare.sh index 8e1472c..8ceaed1 100755 --- a/scripts/run_compare.sh +++ b/scripts/run_compare.sh @@ -1,5 +1,5 @@ #!/usr/bin/env bash -# run_compare.sh — 对所有内置数据集 × 6 算法跑对比可视化 +# run_compare.sh — 对所有内置数据集跑纵向对比可视化(render-column) # # 用法: # bash scripts/run_compare.sh [数据集目录] [输出目录] @@ -8,23 +8,21 @@ # 数据集目录: ../../datasets/builtin # 输出目录: render_compare # -# 管道: inject → detect-batch → render-algo → montage +# 管道: inject → render-column(自动检测 + 渲染) set -uo pipefail BUILTIN_DIR="${1:-../../datasets/builtin}" OUTPUT="${2:-render_compare}" -ALGOS="ksigma iqr grubbs shesd lof sample_ad_model" URL="${TS_ANA_TD_URL:-ws://root:taosdata@localhost:6041}" # 多列数据集映射: dataset_name → value_col -# 未在此映射中的数据集默认使用 "value" 列 declare -A VALUE_COLS=( ["ecg_002"]="ECG1" ["sensor_007"]="ankle_horiz_fwd" ) -mkdir -p "$OUTPUT/det" "$OUTPUT/algo" +mkdir -p "$OUTPUT" total=0 ok=0 @@ -41,48 +39,22 @@ for dataset_dir in "$BUILTIN_DIR"/*/; do echo "=== $name (value_col=$value_col) ===" # 1. 注入 - echo " [1/3] inject..." + echo " [1/2] inject..." uv run python -m ts_anomaly_td inject \ --csv "$csv" --stable "$name" --value-col "$value_col" --url "$URL" \ || { echo " inject FAILED, skipping $name"; ((total++)); continue; } - # 2. 逐算法检测 + 渲染 - panels=() - for algo in $ALGOS; do - det_json="$OUTPUT/det/${name}_${algo}.json" - algo_png="$OUTPUT/algo/${name}_${algo}.png" + # 2. 纵向对比渲染(内部自动检测 6 算法) + echo " [2/2] render-column..." + uv run python -m ts_anomaly_td render-column \ + --gt-csv "$csv" --stable "$name" --value-col "$value_col" \ + --output "$OUTPUT/${name}_compare.png" --url "$URL" \ + && { echo " → $OUTPUT/${name}_compare.png"; ((ok++)); } \ + || { echo " render-column FAILED"; } - echo " [2/3] detect $algo..." - uv run python -m ts_anomaly_td detect-batch \ - --stable "$name" --algo "$algo" --json --url "$URL" \ - > "$det_json" 2>/dev/null || true - - echo " [2/3] render $algo..." - uv run python -m ts_anomaly_td render-algo \ - --gt-csv "$csv" --det-json "$det_json" --algo "$algo" \ - --output "$algo_png" --value-col "$value_col" \ - --title "$name / $algo" \ - || { echo " render $algo FAILED"; continue; } - - panels+=("$algo_png") - done - - # 3. 拼图 - if [[ ${#panels[@]} -gt 0 ]]; then - echo " [3/3] montage (${#panels[@]} panels)..." - uv run python -m ts_anomaly_td montage \ - --inputs "${panels[@]}" \ - --output "$OUTPUT/${name}_compare.png" \ - --cols 2 \ - --title "$name — 算法对比" - echo " → $OUTPUT/${name}_compare.png" - ((ok++)) - else - echo " no panels, skip montage for $name" - fi ((total++)) done echo "" echo "=== 完成: $ok/$total 成功, $skip 跳过 ===" -ls -la "$OUTPUT"/*_compare.png 2>/dev/null || echo "(无拼图输出)" +ls -la "$OUTPUT"/*_compare.png 2>/dev/null || echo "(无输出)" diff --git a/ts_anomaly_td/cli.py b/ts_anomaly_td/cli.py index f3387f9..3b0d9dc 100644 --- a/ts_anomaly_td/cli.py +++ b/ts_anomaly_td/cli.py @@ -225,6 +225,15 @@ def build_parser() -> argparse.ArgumentParser: p_mt.add_argument("--cols", type=int, default=2, help="列数 (default: 2)") p_mt.add_argument("--title", default="", help="总标题") + # render-column + p_rc = sub.add_parser("render-column", help="GT 顶部 + 各算法纵向对比") + p_rc.add_argument("--gt-csv", required=True, help="GT CSV 文件路径") + p_rc.add_argument("--stable", required=True, help="supertable 名称") + p_rc.add_argument("--db", default="ts_anomaly", help="数据库名") + p_rc.add_argument("--value-col", default="value", help="数值列名 (default: value)") + p_rc.add_argument("--output", required=True, help="PNG 输出路径") + p_rc.add_argument("--url", default="ws://root:taosdata@localhost:6041", help="TDengine WebSocket URL") + return parser @@ -258,6 +267,27 @@ def cmd_montage(args): logger.info("montage complete: %s", args.output) +def cmd_render_column(args): + """render-column 子命令:GT 顶部 + 各算法纵向对比。""" + logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s") + + from ts_anomaly_td.io_csv import read_csv + from ts_anomaly_td.visualize import render_column + + ts, vals, labels = read_csv(args.gt_csv, value_col=args.value_col) + + conn = TDConnection(args.url) + try: + db = args.db or "ts_anomaly" + conn.execute_no_result(f"USE {db}") + results = detect_all_algos(conn, args.stable) + finally: + conn.close() + + title = f"{args.stable} — 算法对比" + render_column(ts, vals, labels, results, args.output, title=title) + + def main(): parser = build_parser() args = parser.parse_args() @@ -271,6 +301,7 @@ def main(): "watch": cmd_watch, "render-algo": cmd_render_algo, "montage": cmd_montage, + "render-column": cmd_render_column, } fn = dispatch.get(args.command) diff --git a/ts_anomaly_td/io_csv.py b/ts_anomaly_td/io_csv.py index 819e55b..9223114 100644 --- a/ts_anomaly_td/io_csv.py +++ b/ts_anomaly_td/io_csv.py @@ -26,7 +26,21 @@ def read_csv(path: str | Path, value_col: str = "value") -> tuple[list[int], lis if "timestamp" not in df.columns or value_col not in df.columns: raise ValueError(f"CSV 缺少 timestamp 或 {value_col} 列: {path}") - timestamps = pd.to_datetime(df["timestamp"]) + timestamps = pd.to_datetime(df["timestamp"], format="mixed", errors="coerce") + if timestamps.isna().any(): + # 对无法解析的行尝试修正(如 00:60:00 → 01:00:00) + import re + na_mask = timestamps.isna() + raw_strs = df.loc[na_mask, "timestamp"].astype(str) + + def _fix_hour60(s: str) -> str: + def _repl(m): + h = (int(m.group(1)) + 1) % 24 + return f"{h:02d}:00:00" + return re.sub(r"(\d{2}):60:00", _repl, s) + + fixed = raw_strs.apply(_fix_hour60) + timestamps.loc[na_mask] = pd.to_datetime(fixed, errors="coerce") timestamps_ms = (timestamps.astype("int64") // 10**6).tolist() values = df[value_col].astype(float).tolist() diff --git a/ts_anomaly_td/visualize.py b/ts_anomaly_td/visualize.py index 938ed6a..7fdbaea 100644 --- a/ts_anomaly_td/visualize.py +++ b/ts_anomaly_td/visualize.py @@ -233,6 +233,144 @@ def render_algo( return output_path +def render_column( + timestamps: list[int], + values: list[float], + gt_labels: list[int], + det_results: dict[str, dict], + output_path: str, + title: str = "", +) -> str: + """单列纵向对比图:GT 顶部 + 各算法向下排布,共享 x 轴。 + + 异常区域用鲜明颜色 + 竖线边框标记。 + + Args: + timestamps: 毫秒时间戳列表 + values: 数值列表 + gt_labels: GT 标签(0/1) + det_results: {algo: {"windows": [[start,end],...], "error": str|None}} + output_path: PNG 输出路径 + title: 图表总标题 + + Returns: + output_path + """ + font_name = _find_cjk_font() + plt.rcParams["font.family"] = font_name + plt.rcParams["font.size"] = 9 + + # 过滤有效算法 + valid_algos = [(a, r) for a, r in det_results.items() if not r.get("error")] + n_panels = 1 + len(valid_algos) # GT + 各算法 + + fig, axes = plt.subplots(n_panels, 1, sharex=True, figsize=(16, 3 * n_panels + 1)) + if n_panels == 1: + axes = [axes] + + t0 = timestamps[0] + hours = [(t - t0) / 3_600_000 for t in timestamps] + + # ── 提取 GT windows ── + gt_windows_h = [] + in_anom = False + anom_start = None + for i, label in enumerate(gt_labels): + if label == 1 and not in_anom: + in_anom = True + anom_start = hours[i] + elif label == 0 and in_anom: + in_anom = False + gt_windows_h.append((anom_start, hours[i])) + if in_anom: + gt_windows_h.append((anom_start, hours[-1])) + + # ── Panel 0: Ground Truth ── + ax = axes[0] + has_gt = any(l == 1 for l in gt_labels) + ax.plot(hours, values, color="#333333", linewidth=0.8, alpha=0.9) + if has_gt: + for start_h, end_h in gt_windows_h: + ax.axvspan(start_h, end_h, alpha=0.35, color="#FF4444", + edgecolor="#CC0000", linewidth=1.5) + # 异常点散点标记 + anom_hours = [hours[i] for i, l in enumerate(gt_labels) if l == 1] + anom_vals = [values[i] for i, l in enumerate(gt_labels) if l == 1] + ax.scatter(anom_hours, anom_vals, color="red", s=8, zorder=5, label="GT 异常点") + ax.legend(loc="upper right", fontsize=7) + ax.set_ylabel("Value", fontsize=9) + gt_title = "Ground Truth" + if has_gt: + gt_title += f"({len(gt_windows_h)} 个异常区间)" + else: + gt_title += "(无异常标注)" + ax.set_title(gt_title, fontsize=10, fontweight="bold") + ax.grid(True, alpha=0.3) + + # ── Panels 1..N: 各算法 ── + for idx, (algo, result) in enumerate(valid_algos): + ax = axes[idx + 1] + color = ALGO_COLORS.get(algo, "#888888") + + # 原始序列(浅色背景) + ax.plot(hours, values, color="#AAAAAA", linewidth=0.5, alpha=0.6) + + # GT 红色竖线边框参考 + for start_h, end_h in gt_windows_h: + ax.axvspan(start_h, end_h, alpha=0.12, color="#FF4444", + linestyle="--", edgecolor="#FF6666", linewidth=0.8) + + # 算法检测窗口 — 鲜明颜色 + 边框 + windows = result.get("windows", []) + det_hours_list = [] + for wstart, wend in windows: + ws_h = (wstart - t0) / 3_600_000 + we_h = (wend - t0) / 3_600_000 + ax.axvspan(ws_h, we_h, alpha=0.4, color=color, + edgecolor=color, linewidth=1.5) + det_hours_list.append((ws_h, we_h)) + + # 在检测区间内的高亮数据点 + det_points_h = set() + for ws_h, we_h in det_hours_list: + for i, h in enumerate(hours): + if ws_h <= h <= we_h: + det_points_h.add(i) + if det_points_h: + dp_hours = [hours[i] for i in det_points_h] + dp_vals = [values[i] for i in det_points_h] + ax.scatter(dp_hours, dp_vals, color=color, s=6, zorder=5, alpha=0.7) + + # IoU 计算 + iou = _compute_iou(gt_windows_h, det_hours_list) + stats = f"检测窗口: {len(windows)} | IoU: {iou:.3f}" + + # 图例方块 + from matplotlib.patches import Patch + handles = [Patch(color=color, alpha=0.6, label=algo)] + if has_gt: + handles.insert(0, Patch(color="#FF4444", alpha=0.3, label="GT")) + ax.legend(handles=handles, loc="upper right", fontsize=7) + + ax.set_title(f"{algo} — {stats}", fontsize=10) + ax.set_ylabel("Value", fontsize=9) + ax.grid(True, alpha=0.3) + + # 底部 x 轴标签 + axes[-1].set_xlabel("Time (hours)", fontsize=9) + + # 总标题 + if title: + fig.suptitle(title, fontsize=13, fontweight="bold", y=0.995) + + plt.tight_layout(rect=[0, 0, 1, 0.98] if title else [0, 0, 1, 1]) + 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_column saved: %s (%d panels)", output_path, n_panels) + return output_path + + def montage_pngs( image_paths: list[str], output_path: str, diff --git a/uv.lock b/uv.lock index 376d721..f26a3ea 100644 --- a/uv.lock +++ b/uv.lock @@ -607,6 +607,7 @@ dependencies = [ { name = "matplotlib" }, { name = "numpy" }, { name = "pandas" }, + { name = "pillow" }, { name = "taos-ws-py" }, ] @@ -615,6 +616,7 @@ requires-dist = [ { name = "matplotlib", specifier = ">=3.7" }, { name = "numpy", specifier = ">=1.24" }, { name = "pandas", specifier = ">=2.0" }, + { name = "pillow", specifier = ">=10.0" }, { name = "taos-ws-py", specifier = ">=0.2.0" }, ]