File size: 15,631 Bytes
de1e3fc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
"""Per-dataset + combined re-ID evaluation and intra-subject similarity report.

Two things, over the loaded (known+unknown) dataset pairs, using embeddings already in the DB:

1. INTRA-SUBJECT SIMILARITY β€” for every dog, cosine between all pairs of its OWN photos
   (min / max / avg). Answers "do photos of the same dog read as wildly different?". A random
   different-dog baseline is sampled per dataset for context (intra should sit well above it).

2. RETRIEVAL / MATCHING QUALITY β€” for each found (unknown) dog, rank it against the gallery of
   known dogs with the production dog-level score (max cosine over photo pairs) and report
   recall@1/5/10 (+ counts), mean/median rank, MRR, and true-vs-best-wrong score separation.
   Evaluated for EACH dataset separately and then COMBINED (one shared gallery β€” the other
   dataset's dogs act as extra distractors). Identity is namespaced per source so colliding folder
   numbers across datasets (both have "Dog0") never cross-match.

This measures the MODEL (pure embedding retrieval over the full gallery). It intentionally does NOT
apply the production ZIP-radius / breed gates β€” on this synthetic data each dog's found+known share
a ZIP, so a ZIP gate would inflate recall as an artifact rather than reflect the model.

Dataset pairs are auto-detected by name: "<base> known" + "<base> unknown". Reuses embeddings for
the most-embedded model unless --model is given.

Usage (from backend/):
    python -m scripts.eval_datasets
    python -m scripts.eval_datasets --model hf-embed --out-dir eval_out --seed 0
"""
from __future__ import annotations

import argparse
import csv
from pathlib import Path

import numpy as np
from sqlalchemy import func, select

from app.db import SessionLocal
from app.models import Dataset, Embedding, KnownDog, Picture, UnknownDog
from app.models.base import SubjectType


# ---- identity parsing (baked in by scripts/prepare_dogfacenet.py) ----
def _known_folder(name: str | None) -> str | None:
    return name[3:] if name and name.startswith("Dog") else None


def _found_folder(desc: str | None) -> str | None:
    if desc and desc.lower().startswith("test found dog "):
        return desc.split()[-1]
    return None


def _base_name(name: str) -> tuple[str, str] | None:
    """('Foo known'|'Foo unknown') -> ('Foo', 'known'|'unknown'); else None."""
    low = name.lower()
    for suffix, kind in ((" known", "known"), (" unknown", "unknown")):
        if low.endswith(suffix):
            return name[: -len(suffix)].strip(), kind
    return None


def _pick_model(db, override: str | None) -> tuple[str, str]:
    rows = db.execute(
        select(Embedding.model_name, Embedding.model_version, func.count())
        .group_by(Embedding.model_name, Embedding.model_version)
        .order_by(func.count().desc())
    ).all()
    if not rows:
        raise SystemExit("No embeddings in the DB.")
    if override:
        for mn, mv, _ in rows:
            if override in mn:
                return mn, mv
        raise SystemExit(f"No embeddings match --model {override!r}. Available: {[r[0] for r in rows]}")
    return rows[0][0], rows[0][1]


def _vectors_for_dataset(db, subject_type, dataset_id, model_name, model_version):
    """{subject_id: (n_photos x dim) float32} for one dataset's dogs (active model only)."""
    subj = KnownDog if subject_type == SubjectType.known else UnknownDog
    ids = [i for (i,) in db.execute(select(subj.id).where(subj.dataset_id == dataset_id))]
    if not ids:
        return {}
    rows = db.execute(
        select(Picture.subject_id, Embedding.vector)
        .join(Picture, Embedding.picture_id == Picture.id)
        .where(
            Picture.subject_type == subject_type,
            Picture.subject_id.in_(ids),
            Embedding.model_name == model_name,
            Embedding.model_version == model_version,
        )
    ).all()
    out: dict[int, list[np.ndarray]] = {}
    for sid, blob in rows:
        out.setdefault(sid, []).append(np.frombuffer(blob, dtype=np.float32))
    return {sid: np.vstack(v).astype(np.float32) for sid, v in out.items()}


# ---- intra-subject similarity ----
def _intra_stats(mat: np.ndarray) -> tuple[float, float, float] | None:
    """min/max/avg cosine over all distinct photo pairs of one dog (vectors are L2-normalized)."""
    n = mat.shape[0]
    if n < 2:
        return None
    sims = mat @ mat.T
    iu = np.triu_indices(n, k=1)
    pair = sims[iu]
    return float(pair.min()), float(pair.max()), float(pair.mean())


# ---- retrieval ----
def _eval_retrieval(gallery: dict, queries: dict):
    """gallery/queries: {identity_key: [(dog_label, mat), ...]}. Returns metrics dict.

    Each query dog is scored against every gallery dog (dog-level max cosine over photo pairs);
    rank of the query's own identity is recorded. Gallery may hold several dogs per identity
    (shouldn't here) β€” any same-identity dog counts as correct.
    """
    gal_rows, gal_owner, owner_identity = [], [], []
    for ident, dogs in gallery.items():
        for _label, mat in dogs:
            oidx = len(owner_identity)
            owner_identity.append(ident)
            for v in mat:
                gal_rows.append(v)
                gal_owner.append(oidx)
    if not gal_rows:
        return None
    G = np.vstack(gal_rows).astype(np.float32)
    gal_owner = np.asarray(gal_owner)
    n_owners = len(owner_identity)
    owner_ident_arr = np.asarray(owner_identity, dtype=object)

    ranks, true_scores, best_wrong, skipped = [], [], [], 0
    gallery_idents = set(owner_identity)
    for ident, dogs in queries.items():
        if ident not in gallery_idents:
            skipped += 1
            continue
        for _label, q in dogs:
            sims = q @ G.T
            per_gal = sims.max(axis=0)
            dog_scores = np.full(n_owners, -1.0, dtype=np.float32)
            np.maximum.at(dog_scores, gal_owner, per_gal)
            order = np.argsort(-dog_scores)
            ordered_idents = owner_ident_arr[order]
            correct_mask = ordered_idents == ident
            rank = int(np.argmax(correct_mask)) + 1  # first matching-identity position
            ranks.append(rank)
            true_scores.append(float(dog_scores[order[rank - 1]]))
            wrong = dog_scores[order][~correct_mask]
            best_wrong.append(float(wrong.max()) if wrong.size else -1.0)

    if not ranks:
        return None
    r = np.asarray(ranks)
    ts = np.asarray(true_scores)
    bw = np.asarray(best_wrong)
    n = len(r)
    return {
        "n_queries": n,
        "skipped_no_gallery": skipped,
        "n_gallery_dogs": n_owners,
        "recall@1": int(np.sum(r <= 1)),
        "recall@5": int(np.sum(r <= 5)),
        "recall@10": int(np.sum(r <= 10)),
        "mean_rank": float(r.mean()),
        "median_rank": float(np.median(r)),
        "mrr": float(np.mean(1.0 / r)),
        "true_mean": float(ts.mean()),
        "true_min": float(ts.min()),
        "bestwrong_mean": float(bw.mean()),
        "margin_mean": float(np.mean(ts - bw)),
        "pct_margin_pos": float(np.mean(ts > bw) * 100),
    }


def _print_retrieval(title: str, m: dict) -> None:
    n = m["n_queries"]
    print(f"\n### {title}")
    print(f"  queries: {n} found dogs  |  gallery: {m['n_gallery_dogs']} known dogs"
          f"  |  skipped (no counterpart): {m['skipped_no_gallery']}")
    for k in (1, 5, 10):
        c = m[f"recall@{k}"]
        print(f"  recall@{k:<2} {c/n*100:6.2f}%   ({c}/{n})")
    print(f"  mean rank {m['mean_rank']:.2f}   median {m['median_rank']:.0f}   MRR {m['mrr']:.4f}")
    print(f"  score: true mean {m['true_mean']:.4f} (min {m['true_min']:.4f})  "
          f"best-wrong mean {m['bestwrong_mean']:.4f}  "
          f"margin {m['margin_mean']:+.4f}  (true>wrong {m['pct_margin_pos']:.1f}%)")


def run(model_override: str | None, out_dir: str, seed: int) -> None:
    db = SessionLocal()
    rng = np.random.default_rng(seed)
    try:
        model_name, model_version = _pick_model(db, model_override)
        print(f"Model: {model_name}/{model_version}")

        # Group datasets into (base -> {known: id, unknown: id}) pairs.
        pairs: dict[str, dict[str, int]] = {}
        for d in db.execute(select(Dataset)).scalars():
            parsed = _base_name(d.name)
            if parsed:
                base, kind = parsed
                pairs.setdefault(base, {})[kind] = d.id
        pairs = {b: v for b, v in pairs.items() if "known" in v and "unknown" in v}
        if not pairs:
            raise SystemExit("No '<base> known' + '<base> unknown' dataset pairs found.")
        print("Dataset pairs:", ", ".join(f"{b} (known #{v['known']}, unknown #{v['unknown']})"
                                          for b, v in pairs.items()))

        out = Path(out_dir)
        out.mkdir(parents=True, exist_ok=True)

        # ---- load vectors per dataset, keyed by namespaced identity ----
        # gallery/query dicts: {identity_key: [(label, mat)]}, identity_key = f"{base}:{folder}"
        per_pair_gallery: dict[str, dict] = {}
        per_pair_query: dict[str, dict] = {}
        intra_rows: list[dict] = []
        inter_baseline: dict[str, float] = {}

        for base, v in pairs.items():
            kv = _vectors_for_dataset(db, SubjectType.known, v["known"], model_name, model_version)
            uv = _vectors_for_dataset(db, SubjectType.unknown, v["unknown"], model_name, model_version)
            kname = dict(db.execute(select(KnownDog.id, KnownDog.name).where(KnownDog.dataset_id == v["known"])).all())
            udesc = dict(db.execute(select(UnknownDog.id, UnknownDog.description).where(UnknownDog.dataset_id == v["unknown"])).all())

            gal: dict = {}
            for sid, mat in kv.items():
                folder = _known_folder(kname.get(sid))
                if folder is None:
                    continue
                gal.setdefault(f"{base}:{folder}", []).append((f"{base}/Dog{folder}", mat))
            qry: dict = {}
            for sid, mat in uv.items():
                folder = _found_folder(udesc.get(sid))
                if folder is None:
                    continue
                qry.setdefault(f"{base}:{folder}", []).append((f"{base}/found{folder}", mat))
            per_pair_gallery[base] = gal
            per_pair_query[base] = qry

            # intra-subject stats for every dog in this pair (known + unknown)
            for kind, vecs, names in (("known", kv, kname), ("unknown", uv, udesc)):
                for sid, mat in vecs.items():
                    st = _intra_stats(mat)
                    intra_rows.append({
                        "dataset": base, "kind": kind, "dog_id": sid,
                        "identity": names.get(sid), "n_photos": int(mat.shape[0]),
                        "min": None if st is None else round(st[0], 4),
                        "max": None if st is None else round(st[1], 4),
                        "avg": None if st is None else round(st[2], 4),
                    })

            # inter-subject baseline: sample random cross-dog photo pairs within the pair's known set
            all_photos, all_owner = [], []
            for oi, (sid, mat) in enumerate(kv.items()):
                for row in mat:
                    all_photos.append(row)
                    all_owner.append(oi)
            if len(all_photos) > 2:
                P = np.vstack(all_photos).astype(np.float32)
                own = np.asarray(all_owner)
                sample = min(5000, len(P) * 4)
                ia = rng.integers(0, len(P), sample)
                ib = rng.integers(0, len(P), sample)
                diff = own[ia] != own[ib]
                if diff.any():
                    cos = np.sum(P[ia[diff]] * P[ib[diff]], axis=1)
                    inter_baseline[base] = float(cos.mean())

        # ---- write intra-subject CSV ----
        intra_csv = out / "intra_subject_similarity.csv"
        with intra_csv.open("w", newline="", encoding="utf-8") as fh:
            w = csv.DictWriter(fh, fieldnames=["dataset", "kind", "dog_id", "identity",
                                               "n_photos", "min", "max", "avg"])
            w.writeheader()
            w.writerows(intra_rows)

        # ---- intra-subject report ----
        print("\n" + "=" * 70)
        print("INTRA-SUBJECT SIMILARITY (same dog, photo-to-photo cosine)")
        print("=" * 70)
        for base in pairs:
            rows = [r for r in intra_rows if r["dataset"] == base and r["avg"] is not None]
            singles = sum(1 for r in intra_rows if r["dataset"] == base and r["avg"] is None)
            if not rows:
                continue
            avgs = np.array([r["avg"] for r in rows])
            mins = np.array([r["min"] for r in rows])
            print(f"\n### {base}  ({len(rows)} dogs with >=2 photos, {singles} single-photo skipped)")
            print(f"  per-dog AVG pairwise sim : mean {avgs.mean():.4f}  "
                  f"p10 {np.percentile(avgs,10):.4f}  min {avgs.min():.4f}")
            print(f"  per-dog MIN pairwise sim : mean {mins.mean():.4f}  "
                  f"p10 {np.percentile(mins,10):.4f}  min {mins.min():.4f}")
            for thr in (0.5, 0.7):
                lo = int(np.sum(mins < thr))
                print(f"  dogs with a photo-pair < {thr}: {lo} ({lo/len(rows)*100:.1f}%) "
                      f"(their photos read quite differently)")
            if base in inter_baseline:
                print(f"  different-dog baseline (sampled): {inter_baseline[base]:.4f}  "
                      f"<- intra AVG should sit well above this")

        # ---- retrieval: each dataset, then combined ----
        print("\n" + "=" * 70)
        print("RETRIEVAL / MATCHING QUALITY  (found dog -> ranked known gallery)")
        print("=" * 70)
        summary_rows = []
        for base in pairs:
            m = _eval_retrieval(per_pair_gallery[base], per_pair_query[base])
            if m:
                _print_retrieval(base, m)
                summary_rows.append({"scope": base, **m})

        combined_gal: dict = {}
        combined_qry: dict = {}
        for base in pairs:
            for k, v in per_pair_gallery[base].items():
                combined_gal.setdefault(k, []).extend(v)
            for k, v in per_pair_query[base].items():
                combined_qry.setdefault(k, []).extend(v)
        mc = _eval_retrieval(combined_gal, combined_qry)
        if mc:
            _print_retrieval("COMBINED (both datasets, shared gallery)", mc)
            summary_rows.append({"scope": "COMBINED", **mc})

        # ---- write retrieval summary CSV ----
        ret_csv = out / "retrieval_summary.csv"
        if summary_rows:
            with ret_csv.open("w", newline="", encoding="utf-8") as fh:
                w = csv.DictWriter(fh, fieldnames=list(summary_rows[0].keys()))
                w.writeheader()
                w.writerows(summary_rows)

        print(f"\nWrote per-dog stats  -> {intra_csv}")
        print(f"Wrote retrieval table -> {ret_csv}")
    finally:
        db.close()


def main() -> None:
    ap = argparse.ArgumentParser(description="Per-dataset + combined re-ID & intra-subject metrics.")
    ap.add_argument("--model", default=None, help="Embedding model name substring (default: most-embedded)")
    ap.add_argument("--out-dir", default="eval_out", help="Directory for CSV outputs")
    ap.add_argument("--seed", type=int, default=0, help="RNG seed for the different-dog baseline sample")
    args = ap.parse_args()
    run(args.model, args.out_dir, args.seed)


if __name__ == "__main__":
    main()