Spaces:
Runtime error
Runtime error
| """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() | |