provinans / src /endopath /resolution.py
reversely's picture
Upload folder using huggingface_hub
eea689d verified
Raw
History Blame Contribute Delete
13.8 kB
"""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