| |
| """ |
| 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 |
| from src.reader import LexicalSpanReader |
| from src.router import Passage, calibrate_threshold, risk_coverage_curve |
| from src.schema_utils import LANG_NAMES, default_root, iter_passages, load_report, norm_lang |
|
|
| NO_ANSWER = "no answer present" |
| F1_CORRECT = 0.50 |
|
|
|
|
| |
|
|
| 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 |
|
|
|
|
| |
|
|
| 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) |
| 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") |
|
|
|
|
| |
|
|
| 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, |
| |
| "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}") |
|
|
| |
| 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")) |
|
|
| |
| 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()) |
|
|