study-buddy / app /rag /reranker.py
GitHub Actions
deploy d092bea3608b7a29952f16357fda39b7a29e399b
2e818da
Raw
History Blame Contribute Delete
3.07 kB
"""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