#!/usr/bin/env python3
"""Score text-to-image retrieval for every <provider>.npz in a results directory.

For each query i, rank all N images by cosine similarity and record where the
paired image i lands. Reports Recall@1/5/10 and MRR, writes scores.json.

Usage:
  python score.py --results results/2026-09-10
"""
import argparse
import glob
import json
import os

import numpy as np


def score(img: np.ndarray, txt: np.ndarray) -> dict:
    img = img / np.linalg.norm(img, axis=1, keepdims=True)
    txt = txt / np.linalg.norm(txt, axis=1, keepdims=True)
    sims = txt @ img.T
    target = np.diag(sims)[:, None]
    rank = (sims > target).sum(axis=1) + 1  # 1 = paired image ranked first
    return {
        "recall_at_1": round(float((rank <= 1).mean() * 100), 1),
        "recall_at_5": round(float((rank <= 5).mean() * 100), 1),
        "recall_at_10": round(float((rank <= 10).mean() * 100), 1),
        "mrr": round(float((1.0 / rank).mean()), 3),
        "median_rank": int(np.median(rank)),
    }


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--results", required=True)
    a = ap.parse_args()
    out = {}
    for path in sorted(glob.glob(os.path.join(a.results, "*.npz"))):
        name = os.path.splitext(os.path.basename(path))[0]
        z = np.load(path)
        meta = json.load(open(os.path.join(a.results, f"{name}.json")))
        s = score(z["img"], z["txt"])
        n = meta["n_items"]
        s["images_per_second"] = round(n / meta["image_encode_seconds"], 1)
        s["queries_per_second"] = round(n / meta["text_encode_seconds"], 1)
        out[name] = {**meta, **s}
        print(
            f"{name:18s} R@1 {s['recall_at_1']:5.1f}  R@5 {s['recall_at_5']:5.1f}  "
            f"R@10 {s['recall_at_10']:5.1f}  MRR {s['mrr']:.3f}  "
            f"{s['images_per_second']} img/s"
        )
    json.dump(out, open(os.path.join(a.results, "scores.json"), "w"), indent=1)


if __name__ == "__main__":
    main()
