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