File size: 13,804 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 334 335 336 337 338 339 340 341 342 343 344 345 | """Escalating not-found resolution (docs/prd.md section 7, issue #10).
"A field should only be marked not found after an escalating search, an LLM
pass, an embedding pass, and the two combined, come back empty, not after a
single failed prompt. A missing field that was actually searched for and
confirmed absent is honest. A missing field from one lazy pass is a silent
gap wearing an honest label."
This module owns the *not-found* path only. Fields the first LLM pass already
extracted go to the consensus gate (consensus.py, issue #9), not here. For a
field the LLM pass left empty, the escalation is:
1. LLM pass -- already run by llm_extraction.extract_case_with_llm; its
empty result is the reason we're escalating at all.
2. Embedding pass -- ColPali retrieval for the field's query. Produces no
value (sub-page localization is not built yet: #49), only "which page
looks most relevant, and how distinctively" (rank_pages_with_margin).
That page is a hint for the combined pass, not an answer on its own.
3. Combined pass -- a targeted, single-field LLM re-extraction given the
report text *and* the page image the embedding pass surfaced. This is
the literal "two combined": the text model reading the exact scan
region the visual model flagged, a genuinely different query from the
bulk first pass (which reads text only, all fields at once).
Only when all three come back empty is the field flagged not-found -- via
schema.ChecklistField.flag_not_found, which itself refuses a search_log with
fewer than three recorded passes, so a lazy single-pass miss cannot wear the
"searched and not found" label.
The orchestration core (resolve_field) takes the embedding and combined
passes as injected callables so it's testable without live models; the real
wirings to colpali_retrieval and llm_extraction are thin builders below.
"""
from __future__ import annotations
import base64
from dataclasses import dataclass
from enum import Enum
from pathlib import Path
from typing import Any, Callable, Optional
from endopath import colpali_retrieval, consensus, llm_extraction, precomputed_retrieval
from endopath.schema import Case, ChecklistField, EvidenceSpan, FieldStatus
class ResolutionOutcome(str, Enum):
RESOLVED = "resolved" # a value exists (LLM pass) or was recovered (combined pass)
NOT_FOUND = "not_found" # all three passes empty -> flagged, honestly earned
NOT_APPLICABLE = "not_applicable" # short-circuited, e.g. molecular pre-2023
@dataclass
class VisualHit:
"""What the embedding pass found: the best page and how distinctive it
was. margin is rank_pages_with_margin's best-minus-second-best gap, which
(unlike raw score) is comparable across field queries -- see #18/#9."""
page_number: int
margin: float
top_score: float
@dataclass
class ExtractedValue:
"""What the combined pass recovered for a single field."""
value: Any
confidence: Optional[float] = None
evidence: Optional[EvidenceSpan] = None
def resolve_field(
field_name: str,
field: ChecklistField,
*,
embedding_probe: Callable[[], Optional[VisualHit]],
combined_extract: Callable[[Optional[VisualHit]], Optional[ExtractedValue]],
margin_threshold: float = consensus.DEFAULT_VISUAL_MARGIN_THRESHOLD,
) -> ResolutionOutcome:
"""Run the escalating search for one field and mutate it in place.
embedding_probe() returns the field's VisualHit, or None if the pass
couldn't run (no page images / ColPali unavailable). combined_extract(hit)
returns the recovered value, or None if the combined pass came back empty.
Both are injected so the orchestration is testable without live models.
"""
# A field the checklist itself says doesn't apply (molecular
# classification on a pre-2023 report, FIGO grade on a non-endometrioid
# histotype) already carries its reason from schema.apply_applicability_
# rules. Running the full LLM/embedding/combined search on it would burn
# a model call to "find" something the protocol says can't be there --
# short-circuit instead, per #10's molecular-classification requirement.
if not field.applicable:
return ResolutionOutcome.NOT_APPLICABLE
# The first LLM pass already found a value: not a not-found case. The
# consensus gate (#9) owns found fields; leave this one untouched.
if field.value is not None:
return ResolutionOutcome.RESOLVED
search_log = ["llm pass: no value extracted"]
hit = embedding_probe()
if hit is None:
search_log.append("embedding pass: skipped (no page images or ColPali unavailable)")
elif hit.margin >= margin_threshold:
search_log.append(
f"embedding pass: distinctive match on page {hit.page_number} "
f"(margin {hit.margin:.2f}) -- field may be present, first LLM pass missed it"
)
else:
search_log.append(
f"embedding pass: no distinctive match "
f"(best page {hit.page_number}, margin {hit.margin:.2f})"
)
extracted = combined_extract(hit)
if extracted is not None and extracted.value is not None:
search_log.append(f"combined pass: recovered value {extracted.value!r}")
field.value = extracted.value
field.confidence = extracted.confidence
field.evidence = extracted.evidence
field.status = FieldStatus.NEEDS_REVIEW
field.search_log.extend(search_log)
return ResolutionOutcome.RESOLVED
search_log.append("combined pass: no value recovered")
# Exactly the three passes above -- satisfies flag_not_found's guard that
# a not-found label be earned by a recorded LLM + embedding + combined
# search, not a single failed prompt.
field.flag_not_found(search_log)
return ResolutionOutcome.NOT_FOUND
# --- Real wirings (thin glue to the model modules; the logic above is the
# tested part). ------------------------------------------------------------
_COMBINED_PROMPT = """\
You are re-checking a single CAP/ICCR endometrial carcinoma checklist field \
that a first extraction pass left empty. The field is: {field_name}.
{image_note}Read the report text below (and the page image, if provided) and \
call record_field. Only set a value if it is genuinely supported; leave value \
and evidence_quote null rather than guessing. evidence_quote must be an exact, \
verbatim substring of the report text -- do not paraphrase.
REPORT TEXT:
{report_text}
"""
def _embed_pages_if_available(image_paths: list[Path]):
"""Embed a case's pages once, so a per-field escalation loop doesn't
re-embed the same images for every field. Returns None when the embedding
pass can't run at all (no images, or ColPali not installed)."""
if not image_paths or not colpali_retrieval.is_available():
return None
return colpali_retrieval.embed_pages(image_paths)
def _build_embedding_probe(
field_name: str, page_embeddings, margin_threshold: float
) -> Callable[[], Optional[VisualHit]]:
def probe() -> Optional[VisualHit]:
if page_embeddings is None or field_name not in colpali_retrieval.FIELD_QUERIES:
return None
query = colpali_retrieval.FIELD_QUERIES[field_name]
page_number, top_score, margin = colpali_retrieval.rank_pages_with_margin(
query, page_embeddings
)
return VisualHit(page_number=page_number, margin=margin, top_score=top_score)
return probe
def _build_precomputed_embedding_probe(
field_name: str,
page_embeddings: Optional[list],
query_embeddings: dict,
margin_threshold: float,
) -> Callable[[], Optional[VisualHit]]:
"""Like _build_embedding_probe, but scores precomputed page vectors against
precomputed field-query vectors (precomputed_retrieval, issue #21) with
numpy MaxSim -- no torch, no live ColPali. `page_embeddings` is the case's
list of per-page arrays; `query_embeddings` maps field name -> query vector.
Both come from data/embeddings (issue #19's precompute). Returns None when
the case has no precomputed pages, or this field has no precomputed query.
"""
def probe() -> Optional[VisualHit]:
if page_embeddings is None or field_name not in query_embeddings:
return None
page_number, top_score, margin = precomputed_retrieval.rank_pages_with_margin(
query_embeddings[field_name], page_embeddings
)
return VisualHit(page_number=page_number, margin=margin, top_score=top_score)
return probe
def _build_combined_extract(
field_name: str,
report_text: str,
image_paths: list[Path],
client=None,
) -> Callable[[Optional[VisualHit]], Optional[ExtractedValue]]:
def extract(hit: Optional[VisualHit]) -> Optional[ExtractedValue]:
if field_name not in llm_extraction._FIELD_VALUE_SCHEMAS:
return None
page_image: Optional[Path] = None
page_number: Optional[int] = None
if hit is not None and image_paths:
idx = hit.page_number - 1
if 0 <= idx < len(image_paths):
page_image = image_paths[idx]
page_number = hit.page_number
return _combined_llm_extract(
field_name, report_text, page_image, page_number, client=client
)
return extract
def _combined_llm_extract(
field_name: str,
report_text: str,
page_image: Optional[Path],
page_number: Optional[int],
*,
client=None,
model: str = llm_extraction.MODEL,
) -> Optional[ExtractedValue]:
"""Targeted, single-field, optionally-multimodal re-extraction: the report
text plus the one page image the embedding pass surfaced. Reuses
llm_extraction's per-field value schemas and enum coercion so the combined
pass can't accept a value the first pass' schema would have rejected."""
client = client or llm_extraction.get_client()
value_schema = llm_extraction._FIELD_VALUE_SCHEMAS[field_name]
tool = {
"name": "record_field",
"description": f"Record the extracted value for the endometrial checklist field '{field_name}'.",
"input_schema": llm_extraction._field_schema(value_schema),
}
content: list[dict] = []
image_note = ""
if page_image is not None:
img_bytes = Path(page_image).read_bytes()
content.append(
{
"type": "image",
"source": {
"type": "base64",
"media_type": "image/png",
"data": base64.standard_b64encode(img_bytes).decode("ascii"),
},
}
)
image_note = (
"A visual-retrieval model flagged the attached page image as the "
"most likely location for this field; read it carefully.\n\n"
)
content.append(
{
"type": "text",
"text": _COMBINED_PROMPT.format(
field_name=field_name, image_note=image_note, report_text=report_text
),
}
)
response = client.messages.create(
model=model,
max_tokens=1024,
tools=[tool],
tool_choice={"type": "tool", "name": "record_field"},
messages=[{"role": "user", "content": content}],
)
tool_call = next(block for block in response.content if block.type == "tool_use")
raw = tool_call.input
value = raw.get("value")
if value is None:
return None
if field_name in llm_extraction._ENUM_FIELDS:
try:
value = llm_extraction._ENUM_FIELDS[field_name](value)
except ValueError:
return None
confidence = raw.get("confidence")
quote = raw.get("evidence_quote")
evidence = None
if quote:
verified = quote in report_text
char_start = report_text.find(quote) if verified else None
evidence = EvidenceSpan(
quote=quote,
char_start=char_start,
char_end=(char_start + len(quote)) if char_start is not None else None,
page_number=page_number,
source="combined" if verified else "combined_unverified_quote",
)
if not verified and confidence is not None:
confidence = min(confidence, 0.2)
return ExtractedValue(value=value, confidence=confidence, evidence=evidence)
def resolve_case_not_found(
case: Case,
report_text: str,
image_paths: Optional[list[Path]] = None,
*,
client=None,
margin_threshold: float = consensus.DEFAULT_VISUAL_MARGIN_THRESHOLD,
) -> dict[str, ResolutionOutcome]:
"""Run the escalating not-found search across every field of a case whose
first LLM pass came back empty. Embeds the case's pages once, then hands
each field its own embedding + combined pass. Mutates case.checklist in
place; returns the per-field outcome.
Assumes case.checklist has already had the LLM extraction pass and
schema.apply_applicability_rules run (llm_extraction.extract_case_with_llm
does both) -- resolve_field short-circuits any field already marked
not-applicable rather than searching for it.
"""
image_paths = image_paths or []
page_embeddings = _embed_pages_if_available(image_paths)
outcomes: dict[str, ResolutionOutcome] = {}
for name, field in case.checklist.all_fields().items():
probe = _build_embedding_probe(name, page_embeddings, margin_threshold)
extract = _build_combined_extract(name, report_text, image_paths, client=client)
outcomes[name] = resolve_field(
name,
field,
embedding_probe=probe,
combined_extract=extract,
margin_threshold=margin_threshold,
)
return outcomes
|