Spaces:
Sleeping
Sleeping
File size: 3,074 Bytes
2e818da | 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 | """Adaptive local cross-encoder reranking for genuinely ambiguous retrievals."""
from __future__ import annotations
import os
import threading
from dataclasses import dataclass
from typing import Sequence
from app.rag.models import EvidenceRequest, EvidenceType, RetrievedEvidence
@dataclass(frozen=True)
class RerankDecision:
use: bool
reason: str
class EvidenceReranker:
_model = None
_lock = threading.Lock()
def decide(
self,
request: EvidenceRequest,
candidates: Sequence[RetrievedEvidence],
) -> RerankDecision:
if request.anchor_evidence_ids or request.selection_anchors:
return RerankDecision(False, "explicit_anchor")
if len(candidates) < 9:
return RerankDecision(False, "small_candidate_pool")
documents = {item.evidence.document_id for item in candidates[:24]}
lowered = request.query.casefold()
visual = any(
term in lowered for term in ("table", "figure", "plot", "chart", "diagram", "metric")
) or any(
item.evidence.element_type
in {EvidenceType.TABLE, EvidenceType.FIGURE, EvidenceType.PLOT, EvidenceType.DIAGRAM}
for item in candidates[:16]
)
comparative = len(documents) > 1 and any(
term in lowered
for term in ("compare", "contrast", "across", "difference", "versus", "synthesis")
)
if visual:
return RerankDecision(True, "visual_or_table_ambiguity")
if comparative:
return RerankDecision(True, "cross_document_comparison")
if len(candidates) >= 18 and len(documents) > 1:
return RerankDecision(True, "large_cross_document_pool")
return RerankDecision(False, "fused_ranking_sufficient")
def rerank(
self,
query: str,
candidates: Sequence[RetrievedEvidence],
*,
limit: int = 24,
) -> list[RetrievedEvidence]:
head = list(candidates[:limit])
if not head:
return list(candidates)
model = self._get_model()
scores = list(model.rerank(query, [item.evidence.index_text for item in head], batch_size=16))
for item, score in zip(head, scores):
item.rerank_score = float(score)
head.sort(
key=lambda item: (
-(item.rerank_score if item.rerank_score is not None else float("-inf")),
-item.fused_score,
)
)
return [*head, *candidates[limit:]]
@classmethod
def _get_model(cls):
if cls._model is not None:
return cls._model
with cls._lock:
if cls._model is None:
from fastembed.rerank.cross_encoder import TextCrossEncoder
cls._model = TextCrossEncoder(
model_name=os.getenv(
"RAG_RERANK_MODEL",
"Xenova/ms-marco-MiniLM-L-6-v2",
),
lazy_load=True,
)
return cls._model
|