File size: 1,289 Bytes
8db761b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

from typing import Protocol


class Reranker(Protocol):
    def scores(self, query: str, texts: list[str]) -> list[float]:
        """Relevance score per text for the query; higher = more relevant."""
        ...


class FakeReranker:
    """Deterministic reranker for tests: lexical token-overlap with the query."""

    def scores(self, query: str, texts: list[str]) -> list[float]:
        q = set(query.lower().split())
        return [float(len(q & set(t.lower().split()))) for t in texts]


class BGEReranker:
    """Cross-encoder reranker (BAAI/bge-reranker-v2-m3). Lazy import; uses GPU + fp16 when available.

    Re-scores retrieved passages so the most relevant ones reach the model — sharper,
    better-grounded answers and citations.
    """

    def __init__(self, model_name: str = "BAAI/bge-reranker-v2-m3"):
        from FlagEmbedding import FlagReranker
        import torch
        self._model = FlagReranker(model_name, use_fp16=torch.cuda.is_available())

    def scores(self, query: str, texts: list[str]) -> list[float]:
        if not texts:
            return []
        out = self._model.compute_score([[query, t] for t in texts], normalize=True)
        return [float(x) for x in (out if isinstance(out, list) else [out])]