Wire token-budget packing into the overview/rerank fallback path
This commit is contained in:
@@ -0,0 +1,62 @@
|
||||
"""Bounded sibling expansion for split clinical sections."""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from .models import RetrievalDocument, SearchHit
|
||||
|
||||
|
||||
def _compatible(left: RetrievalDocument, right: RetrievalDocument) -> bool:
|
||||
"""Prevent expansion across two explicitly different clinical scopes."""
|
||||
if left.drug_id != right.drug_id or left.section_key != right.section_key:
|
||||
return False
|
||||
if left.context_labels and right.context_labels:
|
||||
return bool(set(left.context_labels) & set(right.context_labels))
|
||||
return not left.context_labels and not right.context_labels
|
||||
|
||||
|
||||
def expand_siblings(
|
||||
seeds: Sequence[SearchHit],
|
||||
documents: Sequence[RetrievalDocument],
|
||||
*,
|
||||
window: int = 1,
|
||||
limit: int | None = None,
|
||||
) -> list[SearchHit]:
|
||||
"""Add adjacent parts without crossing drug, section, or scope labels.
|
||||
|
||||
Each seed remains first in its group. Neighbours retain the seed score so
|
||||
later context packing treats them as context for that match, not as a new
|
||||
independently scored retrieval result.
|
||||
"""
|
||||
if window < 0:
|
||||
raise ValueError("window must be non-negative")
|
||||
if limit is not None and limit < 0:
|
||||
raise ValueError("limit must be non-negative or None")
|
||||
|
||||
by_position = {
|
||||
(doc.drug_id, doc.section_key, doc.part_index): doc
|
||||
for doc in documents
|
||||
if doc.part_index is not None
|
||||
}
|
||||
output: list[SearchHit] = []
|
||||
seen: set[str] = set()
|
||||
|
||||
for seed in seeds:
|
||||
doc = seed.document
|
||||
group = [doc]
|
||||
if doc.part_index is not None and window:
|
||||
group = []
|
||||
for index in range(doc.part_index - window, doc.part_index + window + 1):
|
||||
sibling = by_position.get((doc.drug_id, doc.section_key, index))
|
||||
if sibling is not None and _compatible(doc, sibling):
|
||||
group.append(sibling)
|
||||
group.sort(key=lambda item: item.part_index if item.part_index is not None else 0)
|
||||
|
||||
for item in group:
|
||||
if item.doc_id in seen:
|
||||
continue
|
||||
seen.add(item.doc_id)
|
||||
output.append(SearchHit(item, seed.score))
|
||||
if limit is not None and len(output) >= limit:
|
||||
return output
|
||||
return output
|
||||
Reference in New Issue
Block a user