"""Generate publication-style figures from an eval-run's CSVs. Reads ----- {results_dir}/closed_set.csv {results_dir}/open_set.csv {results_dir}/roc.csv Writes ------ {results_dir}/figures/ ├── 1_roc_curves.png # open-set ROC, panel per n_refs ├── 2_closed_set_metrics.png # bar chart R@1 / R@5 / mAP × method × n_refs └── 3_n_refs_effect.png # R@1 + AUC vs n_refs If --results-dir is not passed, picks the most recent timestamped folder under /seed_data/eval_results/. Usage (one-time install of matplotlib inside the container): docker compose exec backend pip install matplotlib Then run: docker compose exec backend python -m scripts.plot_eval """ from __future__ import annotations import argparse import csv import logging from collections import defaultdict from pathlib import Path import matplotlib.pyplot as plt import numpy as np log = logging.getLogger("plot_eval") METHOD_ORDER = ("flat", "centroid", "max_sim", "max_sim_bonus") METHOD_COLORS = { "flat": "#888888", "centroid": "#4c72b0", "max_sim": "#dd8452", "max_sim_bonus": "#55a868", } def load_csv(path: Path) -> list[dict]: with path.open(encoding="utf-8") as f: return list(csv.DictReader(f)) def latest_results_dir(root: Path) -> Path | None: if not root.exists(): return None candidates = [d for d in root.iterdir() if d.is_dir() and (d / "closed_set.csv").exists()] if not candidates: return None return max(candidates, key=lambda d: d.name) def trapezoid_auc(points: list[tuple[float, float]]) -> float: pts = sorted(set(points)) pts = [(0.0, 0.0)] + pts + [(1.0, 1.0)] pts = sorted(set(pts)) auc = 0.0 for (x1, y1), (x2, y2) in zip(pts, pts[1:]): auc += (x2 - x1) * (y1 + y2) / 2.0 return auc # ---- Figure 1: ROC curves ------------------------------------------------ def plot_roc(roc_rows: list[dict], out_path: Path) -> None: n_refs_values = sorted({int(r["n_refs"]) for r in roc_rows}) methods = [m for m in METHOD_ORDER if any(r["method"] == m for r in roc_rows)] fig, axes = plt.subplots(1, len(n_refs_values), figsize=(5.2 * len(n_refs_values), 5), sharey=True) if len(n_refs_values) == 1: axes = [axes] for ax, n_refs in zip(axes, n_refs_values): for method in methods: pts = [ (float(r["fpr_mean"]), float(r["tpr_mean"])) for r in roc_rows if int(r["n_refs"]) == n_refs and r["method"] == method ] pts.sort() xs = [0.0] + [p[0] for p in pts] + [1.0] ys = [0.0] + [p[1] for p in pts] + [1.0] auc = trapezoid_auc(pts) ax.plot( xs, ys, label=f"{method} (AUC = {auc:.3f})", color=METHOD_COLORS[method], marker="o", markersize=4, linewidth=1.8, ) ax.plot([0, 1], [0, 1], color="black", linestyle="--", linewidth=0.8, alpha=0.4, label="random") ax.set_xlim(-0.02, 1.02) ax.set_ylim(-0.02, 1.02) ax.set_xlabel("False positive rate (1 − specificity)") ax.set_title(f"n_refs = {n_refs}") ax.grid(True, alpha=0.3, linewidth=0.5) ax.legend(loc="lower right", fontsize=8.5, frameon=True) axes[0].set_ylabel("True positive rate (sensitivity)") fig.suptitle("Open-set ROC — Find My Dog re-identification", fontsize=13.5, fontweight="bold") fig.tight_layout() fig.savefig(out_path, dpi=140, bbox_inches="tight") plt.close(fig) log.info("Wrote %s", out_path) # ---- Figure 2: Closed-set bar chart -------------------------------------- def plot_closed_set(closed_rows: list[dict], out_path: Path) -> None: methods = [m for m in METHOD_ORDER if any(r["method"] == m for r in closed_rows)] n_refs_values = sorted({int(r["n_refs"]) for r in closed_rows}) metrics = [("r1", "R@1"), ("r5", "R@5"), ("map", "mAP")] agg: dict[tuple[str, int], dict[str, list[float]]] = defaultdict( lambda: {m[0]: [] for m in metrics} ) for r in closed_rows: key = (r["method"], int(r["n_refs"])) for k, _ in metrics: agg[key][k].append(float(r[k])) fig, axes = plt.subplots(1, 3, figsize=(15, 4.8), sharey=True) n_groups = len(methods) width = 0.8 / max(len(n_refs_values), 1) n_refs_colors = plt.cm.viridis(np.linspace(0.25, 0.85, len(n_refs_values))) for ax, (mkey, mlabel) in zip(axes, metrics): x = np.arange(n_groups) for j, (n_refs, color) in enumerate(zip(n_refs_values, n_refs_colors)): offset = (j - (len(n_refs_values) - 1) / 2) * width heights = [ np.mean(agg[(m, n_refs)][mkey]) * 100 if agg[(m, n_refs)][mkey] else 0 for m in methods ] errs = [ np.std(agg[(m, n_refs)][mkey]) * 100 if agg[(m, n_refs)][mkey] else 0 for m in methods ] bars = ax.bar( x + offset, heights, width, yerr=errs, label=f"n_refs = {n_refs}", color=color, edgecolor="white", linewidth=0.5, capsize=3, ) # Bar value labels. for xi, h in zip(x + offset, heights): ax.text(xi, h + 1.5, f"{h:.0f}", ha="center", fontsize=7.5, color="#333") ax.set_xticks(x) ax.set_xticklabels(methods, rotation=18, ha="right", fontsize=9) ax.set_title(mlabel) ax.grid(axis="y", alpha=0.3, linewidth=0.5) ax.set_ylim(0, 110) if mkey == "r1": ax.legend(loc="lower right", fontsize=8.5) axes[0].set_ylabel("metric value (%)") fig.suptitle("Closed-set metrics — assumes correct dog is in gallery", fontsize=13.5, fontweight="bold") fig.tight_layout() fig.savefig(out_path, dpi=140, bbox_inches="tight") plt.close(fig) log.info("Wrote %s", out_path) # ---- Figure 3: n_refs effect -------------------------------------------- def plot_n_refs_effect(closed_rows: list[dict], roc_rows: list[dict], out_path: Path) -> None: methods = [m for m in METHOD_ORDER if any(r["method"] == m for r in closed_rows)] n_refs_values = sorted({int(r["n_refs"]) for r in closed_rows}) fig, axes = plt.subplots(1, 2, figsize=(12, 4.8)) # Left panel: closed-set R@1 ax = axes[0] for method in methods: means, stds = [], [] for n_refs in n_refs_values: vals = [ float(r["r1"]) for r in closed_rows if r["method"] == method and int(r["n_refs"]) == n_refs ] means.append(np.mean(vals) * 100 if vals else np.nan) stds.append(np.std(vals) * 100 if vals else 0) ax.errorbar( n_refs_values, means, yerr=stds, label=method, color=METHOD_COLORS[method], marker="o", capsize=4, linewidth=2, ) ax.set_xlabel("number of reference photos per dog") ax.set_ylabel("Closed-set R@1 (%)") ax.set_title("Closed-set R@1 vs. cluster size") ax.set_xticks(n_refs_values) ax.set_ylim(50, 105) ax.grid(alpha=0.3, linewidth=0.5) ax.legend(loc="lower right", fontsize=9) # Right panel: open-set AUC ax = axes[1] by_mn = defaultdict(list) for r in roc_rows: by_mn[(r["method"], int(r["n_refs"]))].append( (float(r["fpr_mean"]), float(r["tpr_mean"])) ) for method in methods: aucs = [] for n_refs in n_refs_values: pts = by_mn.get((method, n_refs), []) aucs.append(trapezoid_auc(pts) if pts else np.nan) ax.plot( n_refs_values, aucs, label=method, color=METHOD_COLORS[method], marker="o", linewidth=2, ) ax.axhline(y=0.5, color="black", linestyle="--", linewidth=0.8, alpha=0.4, label="random") ax.set_xlabel("number of reference photos per dog") ax.set_ylabel("Open-set ROC AUC") ax.set_title("Open-set AUC vs. cluster size") ax.set_xticks(n_refs_values) ax.set_ylim(0.4, 1.02) ax.grid(alpha=0.3, linewidth=0.5) ax.legend(loc="lower right", fontsize=9) fig.suptitle("Effect of reference cluster size", fontsize=13.5, fontweight="bold") fig.tight_layout() fig.savefig(out_path, dpi=140, bbox_inches="tight") plt.close(fig) log.info("Wrote %s", out_path) # ---- Main ---------------------------------------------------------------- def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument( "--results-dir", type=Path, default=None, help="Path to a specific eval result folder. Defaults to latest under /seed_data/eval_results/.", ) parser.add_argument( "--root", type=Path, default=Path("/seed_data/eval_results"), help="Root used to find latest result if --results-dir not given.", ) args = parser.parse_args() logging.basicConfig(level=logging.INFO, format="%(levelname)s %(name)s: %(message)s") results_dir = args.results_dir or latest_results_dir(args.root) if not results_dir: raise SystemExit(f"No eval results found under {args.root}.") log.info("Using results dir: %s", results_dir) closed = load_csv(results_dir / "closed_set.csv") roc = load_csv(results_dir / "roc.csv") figdir = results_dir / "figures" figdir.mkdir(parents=True, exist_ok=True) plot_roc(roc, figdir / "1_roc_curves.png") plot_closed_set(closed, figdir / "2_closed_set_metrics.png") plot_n_refs_effect(closed, roc, figdir / "3_n_refs_effect.png") log.info("All figures in: %s", figdir) if __name__ == "__main__": main()