| """OCR-vs-visual-retrieval degradation benchmark (issue #16, docs/prd.md sections |
| 6, 8, and open risk #3). |
| |
| Open risk #3, verbatim: "Whether visual retrieval actually beats a strong OCR |
| baseline on this specific corpus is an open question, not a known result." |
| This module is the harness to answer it -- it is explicitly NOT built to |
| confirm a foregone conclusion. It compares any number of named extraction |
| approaches on the *same* held-out set with the *same* metric (reusing |
| evaluation.field_accuracy, issue #15), split by a real corpus degradation |
| axis, and reports whichever way the numbers come out. |
| |
| The two approaches the PRD contrasts: |
| - baseline: OCR text -> LLM extraction (llm_extraction). Section 6 insists |
| this be a *genuinely strong* baseline, not a strawman -- so it's the same |
| schema-constrained tool-use extractor the product uses, not a weakened one. |
| - visual: ColPali retrieves the relevant page -> multimodal LLM extraction |
| from that page image (the combined pass from resolution.py, issue #10). |
| Plugged in as one or more predictors so "zero-shot" and "lightly adapted" |
| (section 9) are just two entries in the approaches map. |
| |
| The degradation axis is REAL, not synthetic noise (acceptance criterion): it |
| reuses the exact OCR-damage text heuristic from |
| notebooks/01_dataset_exploration.ipynb -- a spurious mid-sentence period (a |
| lowercase letter, '.', whitespace, lowercase letter), which a real sentence |
| boundary, followed by a capital, wouldn't produce. |
| |
| IMPORTANT caveat, surfaced by notebooks/02_colpali_retrieval_exploration.ipynb |
| (cell 22) and carried here so the benchmark can't quietly overclaim: this is a |
| TEXT-level degradation proxy. Text-level and image-level scan degradation are |
| different axes -- that notebook found a case flagged "clean" by this text |
| heuristic whose scanned image was visibly degraded (redaction block, faded |
| print). A clean/degraded split on this proxy measures robustness to *OCR text* |
| damage, which is exactly the axis where visual retrieval is hypothesized to |
| help; it does not claim to measure visual scan quality. |
| |
| Blocked on real results until an LLM API key, ColPali compute, and a |
| ground-truth labeling pass (issue #15) all exist. The harness and the real |
| degraded-subset partition run today; the predictors are injected so the |
| comparison logic is tested without live models. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import re |
| from dataclasses import dataclass, field |
| from pathlib import Path |
| from typing import Callable, Mapping, Optional |
|
|
| from endopath import colpali_retrieval, consensus, llm_extraction, resolution |
| from endopath.evaluation import ( |
| AccuracyReport, |
| CasePredictions, |
| FieldPrediction, |
| GroundTruthCase, |
| field_accuracy, |
| ) |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| SPURIOUS_MID_SENTENCE_PERIOD = re.compile(r"[a-z]\.\s+[a-z]") |
|
|
| DEGRADATION_CAVEAT = ( |
| "Clean/degraded split is on a TEXT-level OCR-damage proxy (spurious " |
| "mid-sentence periods). Text-level and image-level scan degradation are " |
| "different axes (notebook 02 found a text-'clean' case with a visibly " |
| "degraded image); this measures robustness to OCR text damage, not visual " |
| "scan quality." |
| ) |
|
|
| Predictor = Callable[[str], CasePredictions] |
|
|
|
|
| def has_ocr_text_damage(text: str) -> bool: |
| return bool(SPURIOUS_MID_SENTENCE_PERIOD.search(text or "")) |
|
|
|
|
| def ocr_damage_density(text: str) -> float: |
| """Graded OCR-damage severity: spurious-period matches per 1000 chars. |
| The binary flag saturates (~91% of the corpus), so severity is what lets |
| the benchmark rank reports and take the genuinely worst ones. Empty/short |
| text scores 0.0 rather than dividing by zero.""" |
| text = text or "" |
| if not text: |
| return 0.0 |
| matches = len(SPURIOUS_MID_SENTENCE_PERIOD.findall(text)) |
| return matches / (len(text) / 1000.0) |
|
|
|
|
| def damage_flags_from_scores( |
| scores: Mapping[str, float], *, top_fraction: float = 0.2 |
| ) -> dict[str, bool]: |
| """Turn graded damage scores into a clean/degraded split by flagging the |
| worst `top_fraction` of cases as degraded. Ties at the cutoff are all |
| included on the degraded side, so the split never depends on dict order. |
| This is how the benchmark gets a balanced, meaningfully-worst degraded |
| subset out of a proxy that would otherwise flag almost everything.""" |
| if not 0.0 < top_fraction < 1.0: |
| raise ValueError("top_fraction must be in (0, 1)") |
| if not scores: |
| return {} |
| ordered = sorted(scores.values(), reverse=True) |
| k = max(1, round(len(ordered) * top_fraction)) |
| cutoff = ordered[k - 1] |
| return {barcode: score >= cutoff for barcode, score in scores.items()} |
|
|
|
|
| def load_ocr_damage_scores(reports_csv_path: Optional[Path | str] = None) -> dict[str, float]: |
| """case_barcode -> graded OCR-damage density over the real corpus CSV.""" |
| import pandas as pd |
|
|
| path = Path(reports_csv_path) if reports_csv_path else Path("data") / "raw" / "TCGA_Reports.csv" |
| reports = pd.read_csv(path) |
| return { |
| str(row["patient_filename"]).split(".")[0]: ocr_damage_density(str(row["text"])) |
| for _, row in reports.iterrows() |
| } |
|
|
|
|
| def load_ocr_damage_flags(reports_csv_path: Optional[Path | str] = None) -> dict[str, bool]: |
| """case_barcode -> whether its OCR text shows the spurious-period damage |
| signal. Runs today against the real corpus CSV -- this is what makes the |
| degraded subset real rather than synthetic. case_barcode is the part of |
| patient_filename before the first '.', matching ingestion.py.""" |
| import pandas as pd |
|
|
| path = Path(reports_csv_path) if reports_csv_path else Path("data") / "raw" / "TCGA_Reports.csv" |
| reports = pd.read_csv(path) |
| flags: dict[str, bool] = {} |
| for _, row in reports.iterrows(): |
| case_barcode = str(row["patient_filename"]).split(".")[0] |
| flags[case_barcode] = has_ocr_text_damage(str(row["text"])) |
| return flags |
|
|
|
|
| @dataclass |
| class BenchmarkReport: |
| |
| results: dict[str, dict[str, AccuracyReport]] |
| subset_case_counts: dict[str, int] = field(default_factory=dict) |
|
|
| def degradation_gap(self, approach: str) -> Optional[float]: |
| """clean accuracy minus degraded accuracy for one approach -- how much |
| that approach loses on OCR-damaged reports. Positive = worse on |
| degraded. None if either subset had nothing to score. Reported as a |
| number, not a verdict; the point of the benchmark is that the sign and |
| size are unknown until measured.""" |
| clean = self.results.get(approach, {}).get("clean") |
| degraded = self.results.get(approach, {}).get("degraded") |
| if clean is None or degraded is None: |
| return None |
| if clean.accuracy is None or degraded.accuracy is None: |
| return None |
| return clean.accuracy - degraded.accuracy |
|
|
|
|
| def run_benchmark( |
| ground_truth: Mapping[str, GroundTruthCase], |
| damage_by_case: Mapping[str, bool], |
| approaches: Mapping[str, Predictor], |
| *, |
| numeric_tolerance: float = 0.0, |
| numeric_tolerances: Optional[Mapping[str, float]] = None, |
| ) -> BenchmarkReport: |
| """Score every approach on the same held-out set with the same metric, |
| split into all / clean / degraded subsets. A case whose damage flag is |
| unknown lands only in 'all' -- it is not silently assumed clean.""" |
| barcodes = list(ground_truth) |
| clean = [b for b in barcodes if damage_by_case.get(b) is False] |
| degraded = [b for b in barcodes if damage_by_case.get(b) is True] |
| subsets = {"all": barcodes, "clean": clean, "degraded": degraded} |
|
|
| results: dict[str, dict[str, AccuracyReport]] = {} |
| for name, predictor in approaches.items(): |
| preds = {b: predictor(b) for b in barcodes} |
| results[name] = {} |
| for subset_name, subset_barcodes in subsets.items(): |
| gt_subset = {b: ground_truth[b] for b in subset_barcodes} |
| results[name][subset_name] = field_accuracy( |
| preds, |
| gt_subset, |
| numeric_tolerance=numeric_tolerance, |
| numeric_tolerances=numeric_tolerances, |
| ) |
|
|
| return BenchmarkReport( |
| results=results, |
| subset_case_counts={k: len(v) for k, v in subsets.items()}, |
| ) |
|
|
|
|
| def format_report(report: BenchmarkReport) -> str: |
| """Plain, neutral rendering -- states the numbers for every approach and |
| subset without framing a winner. Ends with the text-vs-image caveat so a |
| reader can't mistake the degradation axis for visual scan quality.""" |
| lines: list[str] = [] |
| counts = report.subset_case_counts |
| lines.append( |
| f"cases: all={counts.get('all', 0)}, clean={counts.get('clean', 0)}, " |
| f"degraded={counts.get('degraded', 0)}" |
| ) |
| for approach, by_subset in report.results.items(): |
| lines.append(f"\n{approach}:") |
| for subset in ("all", "clean", "degraded"): |
| rep = by_subset.get(subset) |
| if rep is None: |
| continue |
| acc = f"{rep.accuracy:.3f}" if rep.accuracy is not None else "n/a" |
| rec = f"{rep.recall:.3f}" if rep.recall is not None else "n/a" |
| lines.append( |
| f" {subset:8s} accuracy={acc} recall={rec} " |
| f"(answered {rep.n_answered}/{rep.n_applicable})" |
| ) |
| gap = report.degradation_gap(approach) |
| if gap is not None: |
| lines.append(f" clean-minus-degraded accuracy gap: {gap:+.3f}") |
| lines.append(f"\nCAVEAT: {DEGRADATION_CAVEAT}") |
| return "\n".join(lines) |
|
|
|
|
| |
| |
|
|
|
|
| def make_ocr_llm_predictor(text_by_case: Mapping[str, str], *, client=None) -> Predictor: |
| """The strong baseline: the same schema-constrained tool-use extractor the |
| product uses, over OCR text. Not weakened -- section 7 requires the OCR |
| baseline be genuinely strong before visual retrieval is trusted over it.""" |
|
|
| def predict(case_barcode: str) -> CasePredictions: |
| text = text_by_case[case_barcode] |
| raw = llm_extraction.extract_raw(text, client=client) |
| fields = llm_extraction.build_checklist_fields(raw, text) |
| return { |
| name: FieldPrediction(value=f.value, confidence=f.confidence) |
| for name, f in fields.items() |
| } |
|
|
| return predict |
|
|
|
|
| def _visual_predictions( |
| text: str, |
| image_paths: list[Path], |
| probe_for: Callable[[str], Callable[[], object]], |
| client, |
| ) -> dict[str, FieldPrediction]: |
| """Shared visual-pipeline loop: for each checklist field, run the field's |
| embedding probe (which page does the retriever flag?) then the combined |
| multimodal LLM extract on that page. `probe_for(name)` returns the field's |
| probe callable -- the one thing that differs between the live-ColPali and |
| precomputed-embeddings variants; everything downstream is identical.""" |
| preds: dict[str, FieldPrediction] = {} |
| for name in colpali_retrieval.FIELD_QUERIES: |
| hit = probe_for(name)() |
| extracted = resolution._build_combined_extract(name, text, image_paths, client=client)(hit) |
| if extracted is not None: |
| preds[name] = FieldPrediction(value=extracted.value, confidence=extracted.confidence) |
| else: |
| preds[name] = FieldPrediction(value=None, confidence=None) |
| return preds |
|
|
|
|
| def make_visual_llm_predictor( |
| image_paths_by_case: Mapping[str, list[Path]], |
| text_by_case: Mapping[str, str], |
| *, |
| client=None, |
| margin_threshold: float = consensus.DEFAULT_VISUAL_MARGIN_THRESHOLD, |
| ) -> Predictor: |
| """The visual approach: ColPali retrieves the best page per field, then a |
| multimodal LLM extracts from that page image (resolution.py's combined |
| pass). Zero-shot as written; a lightly-adapted retriever would be a second |
| predictor built the same way, not a change here. |
| |
| Embeds each case's pages live with ColPali (torch). When the corpus has |
| already been embedded offline (issue #19), prefer |
| make_precomputed_visual_llm_predictor, which scores the stored vectors in |
| numpy and needs no GPU or model load.""" |
|
|
| def predict(case_barcode: str) -> CasePredictions: |
| image_paths = list(image_paths_by_case.get(case_barcode, [])) |
| text = text_by_case.get(case_barcode, "") |
| page_embeddings = resolution._embed_pages_if_available(image_paths) |
| probe_for = lambda name: resolution._build_embedding_probe( |
| name, page_embeddings, margin_threshold |
| ) |
| return _visual_predictions(text, image_paths, probe_for, client) |
|
|
| return predict |
|
|
|
|
| def make_precomputed_visual_llm_predictor( |
| image_paths_by_case: Mapping[str, list[Path]], |
| text_by_case: Mapping[str, str], |
| page_embeddings_by_case: Mapping[str, list], |
| query_embeddings: Mapping[str, object], |
| *, |
| client=None, |
| margin_threshold: float = consensus.DEFAULT_VISUAL_MARGIN_THRESHOLD, |
| ) -> Predictor: |
| """Same visual pipeline as make_visual_llm_predictor, but the retrieval half |
| scores precomputed page + query vectors (issue #19/#21) in numpy instead of |
| embedding pages live. This is what makes the benchmark runnable on a corpus |
| that was embedded offline: no torch, no ~6GB model load, no per-case GPU |
| time. The combined multimodal LLM extraction still runs live (that is the |
| approach under test); only the page-retrieval step reads stored vectors. |
| |
| `page_embeddings_by_case` maps case_barcode -> the case's per-page vectors |
| (precomputed_retrieval.load_case_page_embeddings); `query_embeddings` maps |
| field name -> query vector (precomputed_retrieval.load_query_embeddings). A |
| case absent from page_embeddings_by_case falls back to a skipped embedding |
| pass (probe returns None), exactly as a missing image would live.""" |
|
|
| def predict(case_barcode: str) -> CasePredictions: |
| image_paths = list(image_paths_by_case.get(case_barcode, [])) |
| text = text_by_case.get(case_barcode, "") |
| page_embs = page_embeddings_by_case.get(case_barcode) |
| probe_for = lambda name: resolution._build_precomputed_embedding_probe( |
| name, page_embs, query_embeddings, margin_threshold |
| ) |
| return _visual_predictions(text, image_paths, probe_for, client) |
|
|
| return predict |
|
|