aipluto-backend / scripts /plot_eval.py
andreskoenig's picture
Deploy v1.2.0 hub redesign — status, reject, further candidates, add photo, edit
b12d042 verified
Raw
History Blame Contribute Delete
10.2 kB
"""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()