#!/usr/bin/env python3 """ WHY DID THE READER SCORE F1 0.36 / EM 0.004? There are exactly three ways an extractive reader can fail, and they demand three completely different fixes. Reporting one aggregate F1 cannot tell them apart, so this script separates them: 1. THE SPAN DOES NOT EXIST. The translated answer is a paraphrase, not a quote. No extractor can win. -> the architecture must generate, not extract. Retrain the plan, not the model. 2. THE SPAN EXISTS BUT IN A PASSAGE THE READER DID NOT PICK. -> fix passage selection (retrieval / reranking), leave the span scorer alone. 3. THE SPAN EXISTS IN THE PASSAGE THE READER PICKED, AND IT PICKED THE WRONG WORDS. -> fix the span scorer. This is the one a trained neural reader repairs. THE LADDER (each rung is an upper bound on the one below it) oracle_any best token-F1 over every contiguous span in ANY passage oracle_picked ... restricted to the passage the reader actually chose reader what the reader returned oracle_any - oracle_picked = passage-selection loss (failure mode 2) oracle_picked - reader = span-selection loss (failure mode 3) 1 - oracle_any = irreducible / abstractive (failure mode 1) Run it on a few languages first; the span enumeration is O(words x span_len). python src/diagnose_reader.py --per-lang 400 --langs hi,ta,bn python src/diagnose_reader.py --per-lang 300 # all 14 """ from __future__ import annotations import argparse import json import sys from collections import Counter, defaultdict from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from src.evaluate_reader import iter_eval_rows, token_f1 # noqa: E402 from src.extractability import normalise # noqa: E402 from src.chunkers.base import split_sentences, words_of # noqa: E402 from src.reader import LexicalSpanReader # noqa: E402 from src.router import Passage # noqa: E402 from src.schema_utils import LANG_NAMES, default_root # noqa: E402 # ------------------------------------------------------------------ oracle def best_span_f1(passage: str, gold_tokens: list[str], max_extra: int = 8) -> tuple[float, str]: """ Highest token-F1 achievable by ANY contiguous word span of `passage`, plus the span itself. Computed with an incremental overlap counter rather than by scoring every span from scratch: for a fixed start, extending the window by one word changes the overlap by at most one, so the whole sweep is O(W * L) cheap increments instead of O(W * L * |gold|) Counter intersections. max_extra bounds the window at |gold| + max_extra words. A span much longer than the answer cannot beat a tighter one -- precision falls faster than recall rises -- so the bound costs nothing and saves the quadratic term. """ words = normalise(passage, True).split() if not words or not gold_tokens: return 0.0, "" need = Counter(gold_tokens) G = len(gold_tokens) limit = min(len(words), G + max_extra) best, best_at = 0.0, (0, 0) for i in range(len(words)): have: Counter = Counter() same = 0 hi = min(i + limit, len(words)) for j in range(i, hi): w = words[j] have[w] += 1 if have[w] <= need.get(w, 0): same += 1 if same: n = j - i + 1 f1 = 2 * same / (n + G) # == 2PR/(P+R) with P=same/n, R=same/G if f1 > best: best, best_at = f1, (i, j + 1) if best >= 1.0: return 1.0, " ".join(words[i:j + 1]) return best, " ".join(words[best_at[0]:best_at[1]]) # ------------------------------------------------------------------ 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=400) ap.add_argument("--examples", type=int, default=6) ap.add_argument("--prior-weight", type=float, default=0.0, help="how hard to trust Passage.score when choosing a sentence") ap.add_argument("--ablate", action="store_true", help="sweep prior-weight instead of running one config") ap.add_argument("--real-scores", action="store_true", help="re-score passages with bge-m3 instead of the gold label " "(REQUIRED for an honest prior ablation -- see below)") ap.add_argument("--model", default=None) ap.add_argument("--answer-mode", default="span", choices=["span", "sentence", "both"], help="'span' trims inside the winning sentence; 'sentence' returns " "the whole sentence. Gold answers average 19.2 words -- about a " "sentence -- so trimming may be costing F1. 'both' measures it.") args = ap.parse_args() root = args.root.expanduser().resolve() if args.root else default_root() langs = set(args.langs.split(",")) if args.langs else None print(f"==> data root: {root}") # ---------------------------------------------------------- ablation # The retrieval prior is either worth using or it is not; guessing which is # how the last two days got spent. Sweep it once, on cached rows, and read # the answer off the table. if args.ablate: rows = list(iter_eval_rows(root, langs, args.per_lang)) rows = [r for r in rows if r[4]] # answerable only # ------------------------------------------------------------------ # WHY --real-scores IS NOT OPTIONAL FOR A HONEST NUMBER # # iter_eval_rows builds Passage(score = 1.0 if is_selected else 0.5). # That score IS THE GOLD LABEL. Feeding it to the reader as a "retrieval # prior" and reporting the improvement would be measuring an ORACLE # reranker -- a number no deployed system can reach, and exactly the # kind of leak that makes a leaderboard entry collapse on the held-out # set. # # With the gold label -> the ablation is an UPPER BOUND (headroom). # With --real-scores -> bge-m3 cosine, i.e. what the pipeline # actually produces. That is the number to # build on. # Run both. The gap between them is what better reranking is worth. # ------------------------------------------------------------------ if args.real_scores: import torch from src.evaluate_retrieval import Embedder model = args.model if model is None: hits = list((root / "hf_cache" / "hub").glob("models--BAAI--bge-m3/snapshots/*")) model = str(hits[0]) if hits else "BAAI/bge-m3" device = "cuda" if torch.cuda.is_available() else "cpu" print(f"==> re-scoring passages with {model} on {device}") emb = Embedder(model, device, batch=64, max_len=192) rescored = [] for lang, qid, q, gold, a, psgs in rows: qv = emb.encode([q]) pv = emb.encode([p.text for p in psgs]) sims = (pv @ qv.T).squeeze(-1).float().tolist() rescored.append((lang, qid, q, gold, a, [Passage(p.chunk_id, p.text, float(s), p.lang) for p, s in zip(psgs, sims)])) rows = rescored print(f"==> scores are bge-m3 cosine (REAL). " f"Numbers below are achievable, not an oracle.") else: print("\n !! Passage.score here is the GOLD LABEL, not a retrieval score.") print(" These numbers are an UPPER BOUND on what reranking can buy.") print(" Re-run with --real-scores for the achievable figure.\n") print(f"\n==> ablating prior_weight over {len(rows):,} answerable queries") print(f"\n {'prior_w':>9}{'reader F1':>12}{'exact match':>13}" f"{'read gold psg':>15}{'span len':>10}") print(" " + "-" * 57) modes = (["span", "sentence"] if args.answer_mode == "both" else [args.answer_mode]) modes = ["span", "sentence"] if args.answer_mode == "both" else [args.answer_mode] best = {} for mode in modes: if len(modes) > 1: print(f"\n --- answer_mode={mode} ---") print(f" {'prior_w':>9}{'reader F1':>12}{'exact match':>13}" f"{'read gold psg':>15}{'span len':>10}") print(" " + "-" * 57) best_w, best_f1, base_f1 = 0.0, -1.0, None for w in (0.0, 0.25, 0.50, 0.75, 1.0): rd = LexicalSpanReader(prior_weight=w, answer_mode=mode) f1s = em = gold_hit = ratio = 0.0 for lang, qid, q, gold, _a, psgs in rows: gt = normalise(gold, True).split() if not gt: continue sp = rd.read(q, psgs) f = token_f1(sp.text, gold) f1s += f em += float(f >= 0.999) top = max(psgs, key=lambda p: p.score) gold_hit += float(sp.chunk_id == top.chunk_id) ratio += len(normalise(sp.text, True).split()) / max(1, len(gt)) n_r = max(1, len(rows)) print(f" {w:>9.2f}{f1s/n_r:>12.4f}{em/n_r:>13.4f}" f"{gold_hit/n_r:>14.1%}{ratio/n_r:>10.2f}x") if base_f1 is None: base_f1 = f1s / n_r # the w=0 row of THIS run if f1s / n_r > best_f1: best_w, best_f1 = w, f1s / n_r # Compare against the w=0 row of the SAME subset. A hardcoded # all-language baseline mixes populations and overstates the gain. print(f"\n best prior_weight = {best_w:.2f} (reader F1 {best_f1:.4f}, " f"{100*(best_f1/max(1e-9, base_f1)-1):+.0f}% vs prior_weight 0 " f"on these same {len(rows):,} queries = {base_f1:.4f})") best[mode] = (best_w, best_f1) if len(best) > 1: # Gold answers average 19.2 words -- about a sentence. Trimming to a # sub-span may be cutting sentences down to fragments to match # targets that are sentences. This is the line that settles it. sp_f1, se_f1 = best["span"][1], best["sentence"][1] print(f"\n {'='*56}") print(f" span best F1 {sp_f1:.4f} @ prior {best['span'][0]:.2f}") print(f" sentence best F1 {se_f1:.4f} @ prior {best['sentence'][0]:.2f}") win = "sentence" if se_f1 > sp_f1 else "span" print(f" -> {win} wins by {abs(se_f1-sp_f1):.4f} F1") if win == "span": print(" NOTE: serving still uses answer_mode=sentence. A trimmed span") print(" starts mid-clause and stops mid-number; through TTS that is") print(" unusable regardless of F1. Report BOTH numbers.") print("\n 'read gold psg' = share of queries where the reader read the") print(" top-RETRIEVED passage. At prior_weight 0 this is what the reader") print(" happened to agree with; at 1.0 it is forced.") return 0 mode = "span" if args.answer_mode == "both" else args.answer_mode reader = LexicalSpanReader(prior_weight=args.prior_weight, answer_mode=mode) print(f"==> reader : {reader.name} (prior_weight {args.prior_weight}, " f"answer_mode {mode})\n") agg: dict[str, dict[str, float]] = defaultdict(lambda: defaultdict(float)) examples: list[dict] = [] n_total = 0 for lang, qid, q, gold, answerable, psgs in iter_eval_rows(root, langs, args.per_lang): if not answerable: continue gold_tokens = normalise(gold, True).split() if not gold_tokens: continue span = reader.read(q, psgs) r_f1 = token_f1(span.text, gold) # oracle over every passage, and over the one the reader chose o_any, o_any_text = 0.0, "" o_picked = 0.0 picked_id = span.chunk_id matched = False picked_psg = None for p in psgs: f1, txt = best_span_f1(p.text, gold_tokens) if f1 > o_any: o_any, o_any_text = f1, txt if p.chunk_id == picked_id: matched = True picked_psg = p o_picked = max(o_picked, f1) if not matched: # reader abstained or returned an unattributed span o_picked = o_any # Split the span-selection loss one level further. "Wrong words" is two # different bugs wearing one number: # picked the wrong SENTENCE -> a ranking problem (query overlap picks # the sentence that restates the question) # right sentence, wrong CUT -> a boundary problem (trim/extend) # Only the second is cheap. Measuring which one dominates costs one # best_span_f1 call over ~40 words. o_sent = o_picked if picked_psg is not None and span.start >= 0: sent_text, cursor = None, 0 for sent in split_sentences(picked_psg.text): n = len(words_of(sent)) if n and cursor <= span.start < cursor + n: sent_text = sent break cursor += n if sent_text: o_sent = best_span_f1(sent_text, gold_tokens)[0] a = agg[lang] a["n"] += 1 a["reader_f1"] += r_f1 a["oracle_picked"] += o_picked a["oracle_any"] += o_any a["oracle_sent"] += o_sent a["oracle_em"] += float(o_any >= 0.999) a["reader_em"] += float(r_f1 >= 0.999) # Span-length ratio separates "found the wrong place" from "found the # right place and did not know where to stop". A reader that returns # whole sentences sits at 3-5x with decent recall and dead precision -- # which is EM ~0 with F1 ~0.4, exactly the shape we saw. n_pred, n_gold = len(normalise(span.text, True).split()), len(gold_tokens) a["len_ratio"] += n_pred / max(1, n_gold) a["pred_words"] += n_pred a["gold_words"] += n_gold n_total += 1 if len(examples) < args.examples and o_any - r_f1 > 0.4: examples.append({"lang": lang, "query": q[:90], "gold": gold[:90], "reader": span.text[:90], "oracle": o_any_text[:90], "reader_f1": round(r_f1, 3), "oracle_f1": round(o_any, 3)}) if not n_total: raise SystemExit("no answerable rows — check --langs / --root") # ---------------------------------------------------------- report W = 78 print(f"\n{'='*W}\nPER-LANGUAGE LADDER (answerable queries only)\n{'='*W}") print(f" {'lang':6s}{'n':>7}{'reader':>9}{'oracle@':>9}{'oracle':>9}" f"{'passage':>10}{'span':>9}{'absent':>9}") print(f" {'':6s}{'':>7}{'F1':>9}{'picked':>9}{'any':>9}{'loss':>10}{'loss':>9}{'':>9}") print(" " + "-" * (W - 2)) tot = defaultdict(float) for lang in sorted(agg): a = agg[lang] n = a["n"] rf, op, oa = a["reader_f1"] / n, a["oracle_picked"] / n, a["oracle_any"] / n print(f" {lang:6s}{int(n):>7,}{rf:>9.3f}{op:>9.3f}{oa:>9.3f}" f"{oa-op:>10.3f}{op-rf:>9.3f}{1-oa:>9.3f}") for k in ("n", "reader_f1", "oracle_picked", "oracle_sent", "oracle_any", "oracle_em", "reader_em", "len_ratio", "pred_words", "gold_words"): tot[k] += a[k] n = tot["n"] rf, op, oa = tot["reader_f1"] / n, tot["oracle_picked"] / n, tot["oracle_any"] / n o_em, r_em = tot["oracle_em"] / n, tot["reader_em"] / n print(" " + "-" * (W - 2)) print(f" {'ALL':6s}{int(n):>7,}{rf:>9.3f}{op:>9.3f}{oa:>9.3f}" f"{oa-op:>10.3f}{op-rf:>9.3f}{1-oa:>9.3f}") os_ = tot["oracle_sent"] / n print(f"\n{'='*W}\nWHERE THE F1 GOES\n{'='*W}") absent, psg_loss, span_loss = 1 - oa, oa - op, op - rf total_loss = max(1e-9, absent + psg_loss + span_loss) for label, val, fix in ( ("answer not in ANY passage (abstractive)", absent, "no extractor can fix this -- must generate"), ("in a passage the reader did not pick", psg_loss, "fix passage selection / reranking"), ("in the picked passage, wrong words", span_loss, "fix the span scorer -- a trained reader repairs this"), ): print(f" {label:42s} {val:6.3f} {100*val/total_loss:5.1f}% of the loss") print(f" {'':42s} -> {fix}") # Split the "wrong words" bar into its two very different halves. sent_loss, bound_loss = op - os_, os_ - rf sub_tot = max(1e-9, sent_loss + bound_loss) print(f"\n ...and 'wrong words' ({span_loss:.3f}) splits into:") print(f" {'wrong SENTENCE in the right passage':42s} {sent_loss:6.3f} " f"{100*sent_loss/sub_tot:5.1f}% of it") print(f" {'':42s} -> ranking: query overlap picks the sentence that") print(f" {'':42s} RESTATES the question, not the one that answers it") print(f" {'right sentence, wrong cut':42s} {bound_loss:6.3f} " f"{100*bound_loss/sub_tot:5.1f}% of it") print(f" {'':42s} -> boundary: trim/extend inside the sentence. CHEAP.") print(f"\n exact-match ceiling (oracle) : {o_em:.3f}") print(f" exact-match achieved (reader): {r_em:.3f}") if o_em > 0: print(f" reader captures {100*r_em/o_em:.1f}% of the achievable exact matches") ratio = tot["len_ratio"] / n print(f"\n span length vs gold : {ratio:.2f}x " f"({tot['pred_words']/n:.1f} words returned vs {tot['gold_words']/n:.1f} in the answer)") if ratio > 2.0: print(" ^ THE READER IS RETURNING SENTENCES, NOT SPANS. Precision is capped at\n" f" ~{1/ratio:.2f} before the span is even in the right place, and exact match\n" " is structurally near zero. Fix the STOP boundary, not the search.") elif ratio < 0.6: print(" ^ spans are TRUNCATED -- the scorer stops too early and drops recall.") print(f"\n{'='*W}\nVERDICT\n{'='*W}") bars = {"absent (must generate)": absent, "wrong passage (rerank)": psg_loss, "wrong sentence (ranking)": sent_loss, "wrong cut (boundary)": bound_loss} top = max(bars, key=bars.get) fixable = {k: v for k, v in bars.items() if not k.startswith("absent")} top_fix = max(fixable, key=fixable.get) spread = max(fixable.values()) - min(fixable.values()) print(f" largest bar overall : {top} ({bars[top]:.3f})") print(f" largest FIXABLE bar : {top_fix} ({fixable[top_fix]:.3f})") if absent >= max(fixable.values()): print(f"\n THE BIGGEST SINGLE LOSS IS UNFIXABLE BY EXTRACTION ({absent:.3f}).") print(f" The extraction ceiling is {oa:.3f} and no reader beats it. Before") print(" spending days on any bar below, decide whether the answering path") print(" should be extractive at all.") if spread < 0.06: print(f"\n The three fixable bars are within {spread:.3f} of each other -- there is") print(" NO dominant fix. Sequential work on any one of them buys little on its") print(" own. Treat this as a signal to change the approach, not to grind.") else: print(f"\n -> work {top_fix} first ({fixable[top_fix]:.3f}).") if examples: print(f"\n{'='*W}\nEXAMPLES WHERE THE ORACLE BEATS THE READER BY >0.4 F1\n{'='*W}") for e in examples: print(f"\n [{e['lang']}] reader {e['reader_f1']} oracle {e['oracle_f1']}") print(f" Q : {e['query']}") print(f" gold : {e['gold']}") print(f" reader : {e['reader']}") print(f" oracle : {e['oracle']}") out = root / "results" / "reader_diagnosis.json" out.parent.mkdir(parents=True, exist_ok=True) out.write_text(json.dumps({ "n_answerable": int(n), "per_language": { lg: {"n": int(a["n"]), "reader_f1": round(a["reader_f1"] / a["n"], 4), "oracle_picked": round(a["oracle_picked"] / a["n"], 4), "oracle_sent": round(a["oracle_sent"] / a["n"], 4), "oracle_any": round(a["oracle_any"] / a["n"], 4), "oracle_em": round(a["oracle_em"] / a["n"], 4), "language": LANG_NAMES.get(lg, lg)} for lg, a in agg.items()}, "overall": {"reader_f1": round(rf, 4), "oracle_picked": round(op, 4), "oracle_sent": round(os_, 4), "oracle_any": round(oa, 4), "oracle_em": round(o_em, 4), "reader_em": round(r_em, 4), "span_len_ratio": round(tot["len_ratio"] / n, 3), "mean_pred_words": round(tot["pred_words"] / n, 2), "mean_gold_words": round(tot["gold_words"] / n, 2)}, "loss_decomposition": {"absent": round(absent, 4), "passage_selection": round(psg_loss, 4), "span_selection": round(span_loss, 4), "sentence_selection": round(sent_loss, 4), "boundary": round(bound_loss, 4)}, "examples": examples, }, indent=2, ensure_ascii=False)) print(f"\n==> wrote {out}") return 0 if __name__ == "__main__": raise SystemExit(main())