Spaces:
Sleeping
Sleeping
| """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 | |
| 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:]] | |
| 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 | |