File size: 15,185 Bytes
eea689d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 | """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,
)
# The real OCR-damage heuristic from notebook 01: a spurious mid-sentence
# period. A genuine sentence boundary is followed by a capital letter; a
# lowercase letter after "word. word" is an OCR artifact.
#
# Measured over the real corpus, this binary flag fires on ~91% of reports
# (500/545 ingested UCEC image cases) -- one spurious period anywhere trips
# it, so as a binary split it separates "essentially all reports" from a tiny
# clean remainder, not worst-from-typical. ocr_damage_density() below grades
# severity (matches per 1000 chars) so the benchmark can select a genuinely
# *worst-N%* degraded subset -- still real corpus OCR damage, but a
# discriminating axis rather than a near-constant flag.
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] # case_barcode -> predictions
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[approach_name][subset] -> AccuracyReport; subset in {all, clean, degraded}
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)
# --- Real predictor wirings (thin glue; the comparison logic above is the
# tested part). ------------------------------------------------------------
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( # noqa: E731
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( # noqa: E731
name, page_embs, query_embeddings, margin_threshold
)
return _visual_predictions(text, image_paths, probe_for, client)
return predict
|