#!/usr/bin/env python3
"""Reproduce held-out Jev vs lexical arithmetic from the sanitized attachment."""
from __future__ import annotations

import json
from pathlib import Path

HERE = Path(__file__).resolve().parent


def main() -> None:
    data = json.loads((HERE / "heldout-results.json").read_text())
    rows = data["held_out"]
    n = len(rows)
    inj = sum(r["jev"]["inject"] == r["gold"]["inject"] for r in rows)
    rel = sum(r["jev"]["relevant_yes"] == r["gold"]["relevant"] for r in rows)
    reln = sum(r["jev"]["relation"] == r["gold"]["relation"] for r in rows)
    tp = sum(r["gold"]["inject"] and r["jev"]["inject"] for r in rows)
    fp = sum((not r["gold"]["inject"]) and r["jev"]["inject"] for r in rows)
    fn = sum(r["gold"]["inject"] and (not r["jev"]["inject"]) for r in rows)
    gold_pos = sum(r["gold"]["inject"] for r in rows)
    brier = sum((r["jev"]["relevant_noul"] - (1.0 if r["gold"]["relevant"] else 0.0)) ** 2 for r in rows) / n
    base_inj = sum(r["baseline"]["inject"] == r["gold"]["inject"] for r in rows)
    base_rel = sum(r["baseline"]["relevant_yes"] == r["gold"]["relevant"] for r in rows)
    base_reln = sum(r["baseline"]["relation"] == r["gold"]["relation"] for r in rows)
    print(
        json.dumps(
            {
                "held_out_n": n,
                "jev_inject": f"{inj}/{n}",
                "jev_relevant": f"{rel}/{n}",
                "jev_relation": f"{reln}/{n}",
                "useful_kept": f"{tp}/{gold_pos}",
                "false_positives": fp,
                "false_exclusions": fn,
                "brier_relevant": round(brier, 6),
                "baseline_inject": f"{base_inj}/{n}",
                "baseline_relevant": f"{base_rel}/{n}",
                "baseline_relation": f"{base_reln}/{n}",
            },
            indent=2,
        )
    )


if __name__ == "__main__":
    main()
