-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathretrieval.py
More file actions
97 lines (78 loc) · 2.98 KB
/
Copy pathretrieval.py
File metadata and controls
97 lines (78 loc) · 2.98 KB
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
92
93
94
95
96
97
"""
Hybrid retrieval: dense (Qdrant) + sparse (BM25) candidates, merged and
reranked with a cross-encoder. Loaded once at API startup and reused.
"""
import pickle
from dataclasses import dataclass
from qdrant_client import QdrantClient
from sentence_transformers import CrossEncoder, SentenceTransformer
import config
@dataclass
class Chunk:
text: str
source: str
page: int
score: float = 0.0
class Retriever:
def __init__(self):
print("Loading embedding model ...")
self.embedder = SentenceTransformer(config.EMBEDDING_MODEL)
print("Loading Qdrant collection ...")
self.qdrant = QdrantClient(path=config.QDRANT_PATH)
print("Loading BM25 index ...")
with open(config.BM25_INDEX_PATH, "rb") as f:
data = pickle.load(f)
self.bm25 = data["bm25"]
self.bm25_chunks = data["chunks"]
self.reranker = None
if config.USE_RERANKER:
print("Loading reranker ...")
self.reranker = CrossEncoder(config.RERANKER_MODEL)
def _dense_search(self, query: str, k: int) -> list[Chunk]:
vec = self.embedder.encode(query).tolist()
hits = self.qdrant.search(
collection_name=config.QDRANT_COLLECTION, query_vector=vec, limit=k
)
return [
Chunk(text=h.payload["text"], source=h.payload["source"], page=h.payload["page"], score=h.score)
for h in hits
]
def _sparse_search(self, query: str, k: int) -> list[Chunk]:
scores = self.bm25.get_scores(query.lower().split())
top_idx = sorted(range(len(scores)), key=lambda i: scores[i], reverse=True)[:k]
return [
Chunk(
text=self.bm25_chunks[i]["text"],
source=self.bm25_chunks[i]["source"],
page=self.bm25_chunks[i]["page"],
score=scores[i],
)
for i in top_idx
]
@staticmethod
def _dedupe(chunks: list[Chunk]) -> list[Chunk]:
seen, out = set(), []
for c in chunks:
key = (c.source, c.page, c.text[:50])
if key not in seen:
seen.add(key)
out.append(c)
return out
def retrieve(self, query: str) -> list[Chunk]:
dense = self._dense_search(query, config.RETRIEVAL_TOP_K)
sparse = self._sparse_search(query, config.RETRIEVAL_TOP_K)
candidates = self._dedupe(dense + sparse)
if self.reranker and candidates:
pairs = [(query, c.text) for c in candidates]
scores = self.reranker.predict(pairs)
for c, s in zip(candidates, scores):
c.score = float(s)
candidates.sort(key=lambda c: c.score, reverse=True)
return candidates[: config.FINAL_TOP_K]
# Singleton, loaded once when the API process starts.
_retriever: Retriever | None = None
def get_retriever() -> Retriever:
global _retriever
if _retriever is None:
_retriever = Retriever()
return _retriever