#!/usr/bin/env python3 """ Reader evaluation: accuracy, calibration, and the conformal thresholds. WHAT THIS PRODUCES ------------------ 1. EM / token-F1 on answerable queries -- is the reader any good? 2. (confidence, correct) pairs -- the ONLY input calibration needs 3. Conformal thresholds, per language -- the guaranteed error bound 4. Risk-coverage curve + AURC -- the headline guardrail plot 5. Expected Calibration Error -- does confidence mean anything? WHY UNANSWERABLE QUERIES ARE INCLUDED ------------------------------------- ~45% of queries have no selected passage. Answering one is ALWAYS wrong, so they enter calibration as (confidence, correct=False). That is what forces the conformal threshold high enough to exclude them, and it is why the resulting bound is meaningful rather than a bound over the easy half of the data. Evaluating only on answerable queries would produce a guarantee that silently excludes the cases requirement 6 actually grades. WHY F1 >= 0.5 IS "CORRECT" -------------------------- Extractive QA convention (SQuAD). Exact match alone is too strict for spans that are right but include a trailing clause; raw overlap is too lenient. Both are reported so the choice is visible rather than buried. SCOPE ----- The reader is given the query's own candidate passages, so this measures READER quality in isolation. End-to-end numbers will be lower, bounded by retrieval -- src/evaluate_retrieval.py measured Hit@5 around 0.80 on the corpus pool. Stated plainly here rather than discovered later. CPU-only. Runs on either host in about a minute. python src/evaluate_reader.py --per-lang 2000 python src/evaluate_reader.py --per-lang 2000 --langs hi,ta,bn """ from __future__ import annotations import argparse import json import math import sys from collections import Counter, defaultdict from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from src.extractability import normalise # noqa: E402 from src.reader import LexicalSpanReader # noqa: E402 from src.router import Passage, calibrate_threshold, risk_coverage_curve # noqa: E402 from src.schema_utils import LANG_NAMES, default_root, iter_passages, load_report, norm_lang # noqa: E402 NO_ANSWER = "no answer present" F1_CORRECT = 0.50 # ------------------------------------------------------------------ metrics def token_f1(pred: str, gold: str) -> float: """SQuAD-style token F1 on normalised tokens.""" p, g = normalise(pred, True).split(), normalise(gold, True).split() if not p or not g: return float(p == g) common = Counter(p) & Counter(g) same = sum(common.values()) if same == 0: return 0.0 prec, rec = same / len(p), same / len(g) return 2 * prec * rec / (prec + rec) def exact_match(pred: str, gold: str) -> float: return float(normalise(pred, True) == normalise(gold, True)) def ece(confidences: list[float], correct: list[bool], bins: int = 10) -> float: """Expected Calibration Error. If confidence 0.8 does not mean 80% correct, the conformal threshold still holds but coverage suffers -- so this is the number to watch when deciding whether to temperature-scale.""" n = len(confidences) if n == 0: return 0.0 total = 0.0 for b in range(bins): lo, hi = b / bins, (b + 1) / bins idx = [i for i, c in enumerate(confidences) if (lo <= c < hi) or (b == bins - 1 and c == 1.0)] if not idx: continue acc = sum(correct[i] for i in idx) / len(idx) conf = sum(confidences[i] for i in idx) / len(idx) total += (len(idx) / n) * abs(acc - conf) return total # ------------------------------------------------------------------ data def iter_eval_rows(root: Path, langs: set[str] | None, per_lang: int): """Yield (lang, query, gold_answer, answerable, passages).""" import polars as pl rep = load_report(root) fmap, pmap = rep["field_mapping"], rep["passage_mapping"] pcol = fmap["passages"] qid_c, q_c, lang_c = fmap["query_id"], fmap["query"], fmap.get("lang") ans_c = fmap.get("answer") t_key, en_key, sel_key = pmap["text"], pmap.get("text_en"), pmap.get("is_selected") for fp in rep["files"]: if "val" not in Path(fp).name: continue try: df = pl.read_parquet(fp, columns=[c for c in (qid_c, q_c, ans_c, lang_c, pcol) if c], n_rows=per_lang * 3) except Exception as exc: print(f" skip {Path(fp).name}: {exc}") continue lang = norm_lang(df[lang_c][0]) if lang_c and len(df) else "?" if langs and lang not in langs: continue kept = 0 get = lambda c: df[c].to_list() if c and c in df.columns else [None] * len(df) # noqa: E731 for qid, q, ans, plist in zip(df[qid_c].to_list(), get(q_c), get(ans_c), df[pcol].to_list()): if kept >= per_lang: break if not isinstance(q, str) or not q.strip(): continue psgs, has_gold = [], False for idx, text, _en, sel, _u in iter_passages(plist, t_key, en_key, sel_key, None): if not isinstance(text, str) or not text.strip(): continue psgs.append(Passage(f"{qid}:{idx}", text, 1.0 if sel == 1 else 0.5, lang)) has_gold |= (sel == 1) if not psgs: continue gold = ans if isinstance(ans, str) else "" answerable = has_gold and bool(gold.strip()) and \ not gold.strip().lower().startswith(NO_ANSWER) yield lang, str(qid), q, gold, answerable, psgs kept += 1 print(f" {Path(fp).name:22s} {lang:3s} {kept:,} queries") # ------------------------------------------------------------------ run def main() -> int: ap = argparse.ArgumentParser() ap.add_argument("--root", type=Path, default=None) ap.add_argument("--langs", default=None) ap.add_argument("--per-lang", type=int, default=2000) ap.add_argument("--alpha", type=float, default=0.10, help="target error rate among ANSWERED queries") ap.add_argument("--delta", type=float, default=0.10, help="confidence for the bound") args = ap.parse_args() root = args.root.expanduser().resolve() if args.root else default_root() print(f"==> data root: {root}") langs = set(args.langs.split(",")) if args.langs else None reader = LexicalSpanReader() print(f"==> reader: {reader.name}\n") rows = [] for lang, qid, q, gold, answerable, psgs in iter_eval_rows(root, langs, args.per_lang): span = reader.read(q, psgs) f1 = token_f1(span.text, gold) if answerable else 0.0 em = exact_match(span.text, gold) if answerable else 0.0 rows.append({ "lang": lang, "query_id": qid, "answerable": answerable, "conf": float(span.score), "f1": f1, "em": em, # Answering an unanswerable query is ALWAYS wrong. "correct": bool(answerable and f1 >= F1_CORRECT), }) if not rows: raise SystemExit("no rows evaluated — check --langs and the validation split") ans = [r for r in rows if r["answerable"]] print(f"\n{'='*72}\nREADER QUALITY (n={len(rows):,}, answerable {len(ans):,} " f"= {100*len(ans)/len(rows):.0f}%)\n{'='*72}") print(f" on ANSWERABLE queries: EM {sum(r['em'] for r in ans)/len(ans):.4f} " f"F1 {sum(r['f1'] for r in ans)/len(ans):.4f} " f"correct(F1>={F1_CORRECT}) {sum(r['correct'] for r in ans)/len(ans):.4f}") print(f" over ALL queries: correct {sum(r['correct'] for r in rows)/len(rows):.4f}") print(" ^ the second number is the one conformal bounds, because answering an") print(" unanswerable query counts as an error.") print(f"\n{'='*72}\nPER LANGUAGE\n{'='*72}") print(f" {'lang':12s}{'n':>7}{'EM':>8}{'F1':>8}{'acc':>8}{'mean conf':>11}{'ECE':>8}") print(" " + "-" * 60) per_lang_rows = defaultdict(list) for r in rows: per_lang_rows[r["lang"]].append(r) for lang in sorted(per_lang_rows): rs = per_lang_rows[lang] a = [r for r in rs if r["answerable"]] or rs e = ece([r["conf"] for r in rs], [r["correct"] for r in rs]) print(f" {LANG_NAMES.get(lang, lang):12s}{len(rs):>7,}" f"{sum(r['em'] for r in a)/len(a):>8.3f}{sum(r['f1'] for r in a)/len(a):>8.3f}" f"{sum(r['correct'] for r in rs)/len(rs):>8.3f}" f"{sum(r['conf'] for r in rs)/len(rs):>11.3f}{e:>8.3f}") # ---------------- calibration ---------------- conf = [r["conf"] for r in rows] corr = [r["correct"] for r in rows] print(f"\n{'='*72}\nCALIBRATION\n{'='*72}") print(f" ECE (all queries): {ece(conf, corr):.4f} " f"{'well calibrated' if ece(conf, corr) < 0.10 else 'MISCALIBRATED — temperature-scale on dev'}") rc = risk_coverage_curve(conf, corr) print(f" AURC: {rc['aurc']:.4f} (lower is better; risk at full coverage " f"{rc['full_coverage_risk']:.3f})") print("\n coverage reachable at each error target:") for k, v in rc["coverage_at_alpha"].items(): a = float(k.split("_")[1]) print(f" alpha {a:.0%} -> answer {v:.1%} of queries") tau, diag = calibrate_threshold(conf, corr, alpha=args.alpha, delta=args.delta) print(f"\n GLOBAL threshold at alpha={args.alpha:.0%}, delta={args.delta:.0%}") print(f" tau {tau:.4f} coverage {diag['coverage']:.1%} " f"empirical error {diag['empirical_error']:.2%} bound {diag['error_ucb']:.2%}") print(f" {diag['note']}") print(f"\n PER-LANGUAGE (Mondrian) thresholds at alpha={args.alpha:.0%}:") print(f" {'lang':12s}{'tau':>8}{'coverage':>10}{'emp.err':>9}{'bound':>8}") print(" " + "-" * 47) per_lang_cfg = {} for lang in sorted(per_lang_rows): rs = per_lang_rows[lang] t, dg = calibrate_threshold([r["conf"] for r in rs], [r["correct"] for r in rs], alpha=args.alpha, delta=args.delta) per_lang_cfg[lang] = {"tau_extract": round(t, 4)} print(f" {LANG_NAMES.get(lang, lang):12s}{t:>8.4f}{dg['coverage']:>10.1%}" f"{dg['empirical_error']:>9.2%}{dg['error_ucb']:>8.2%}") spread = max(v["tau_extract"] for v in per_lang_cfg.values()) - \ min(v["tau_extract"] for v in per_lang_cfg.values()) print(f"\n threshold spread across languages: {spread:.4f}") print(" " + ("^ per-language calibration is worth it" if spread > 0.05 else "^ narrow — a single global threshold would do")) # ---------------- outputs ---------------- res = root / "results" res.mkdir(parents=True, exist_ok=True) (res / "reader_eval.json").write_text(json.dumps({ "reader": reader.name, "n": len(rows), "n_answerable": len(ans), "alpha": args.alpha, "delta": args.delta, "em": round(sum(r["em"] for r in ans) / len(ans), 4), "f1": round(sum(r["f1"] for r in ans) / len(ans), 4), "accuracy_all": round(sum(r["correct"] for r in rows) / len(rows), 4), "ece": round(ece(conf, corr), 4), "aurc": rc["aurc"], "coverage_at_alpha": rc["coverage_at_alpha"], "global_threshold": {"tau": round(tau, 4), **{k: (round(v, 4) if isinstance(v, float) else v) for k, v in diag.items()}}, "per_lang": per_lang_cfg, "risk_coverage": {"coverage": rc["coverage"], "risk": rc["risk"]}, }, indent=2)) print(f"\n==> wrote {res/'reader_eval.json'}") cfg = res / "router_config.json" cfg.write_text(json.dumps({"tau_extract": round(tau, 4), "per_lang": per_lang_cfg}, indent=2)) print(f"==> wrote {cfg} (feed straight into RouterConfig)") try: import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(13, 5.5), constrained_layout=True) ax1.plot(rc["coverage"], rc["risk"], lw=2, color="#4C78A8") ax1.axhline(args.alpha, ls="--", color="#E45756", label=f"alpha = {args.alpha:.0%}") cov_a = rc["coverage_at_alpha"][f"alpha_{args.alpha:.2f}"] ax1.axvline(cov_a, ls=":", color="#54A24B", label=f"coverage {cov_a:.1%}") ax1.set_xlabel("coverage (fraction of queries answered)") ax1.set_ylabel("risk (error rate among answered)") ax1.set_title(f"Risk-coverage AURC = {rc['aurc']:.4f}", weight="bold") ax1.legend(frameon=False) ax1.grid(alpha=.3) bins = 10 xs, ys, ns = [], [], [] for b in range(bins): lo, hi = b / bins, (b + 1) / bins idx = [i for i, c in enumerate(conf) if lo <= c < hi or (b == bins - 1 and c == 1.0)] if idx: xs.append((lo + hi) / 2) ys.append(sum(corr[i] for i in idx) / len(idx)) ns.append(len(idx)) ax2.plot([0, 1], [0, 1], ls="--", color="#999", label="perfect calibration") ax2.scatter(xs, ys, s=[max(20, 400 * n / max(ns)) for n in ns], color="#54A24B", zorder=3, label="observed") ax2.set_xlabel("predicted confidence") ax2.set_ylabel("observed accuracy") ax2.set_title(f"Reliability ECE = {ece(conf, corr):.4f}", weight="bold") ax2.legend(frameon=False) ax2.grid(alpha=.3) out = res / "reader_calibration.png" fig.savefig(out, dpi=150) print(f"==> wrote {out}") except Exception as exc: print(f" (plot skipped: {exc})") return 0 if __name__ == "__main__": raise SystemExit(main())