-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathreranker.py
More file actions
107 lines (85 loc) · 3.42 KB
/
Copy pathreranker.py
File metadata and controls
107 lines (85 loc) · 3.42 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
98
99
100
101
102
103
104
105
106
107
"""
Reranker — Cross-encoder reranking for multi-stage retrieval.
Stage 1 (bi-encoder): fast cosine similarity retrieval → top-K candidates
Stage 2 (cross-encoder): precise relevance scoring → reranked results
Uses sentence-transformers CrossEncoder with ms-marco-MiniLM-L-6-v2 (22M params).
Falls back to pass-through if cross-encoder is unavailable.
"""
import logging
import os
from typing import List, Optional, Tuple
logger = logging.getLogger(__name__)
_cross_encoder = None
_cross_encoder_loaded = False
def _get_cross_encoder():
global _cross_encoder, _cross_encoder_loaded
if _cross_encoder_loaded:
return _cross_encoder
_cross_encoder_loaded = True
try:
from sentence_transformers import CrossEncoder
device = "cpu" if os.environ.get("PYTORCH_MPS_ENABLED") == "0" else None
model = CrossEncoder("cross-encoder/ms-marco-MiniLM-L-6-v2", max_length=512,
device=device)
_cross_encoder = model
logger.info("CrossEncoder loaded: cross-encoder/ms-marco-MiniLM-L-6-v2")
except Exception as e:
logger.warning("CrossEncoder unavailable (%s), reranking disabled", e)
_cross_encoder = None
return _cross_encoder
get_cross_encoder = _get_cross_encoder
def rerank(query: str, candidates: List[dict],
content_key: str = "content",
top_k: int = 0,
score_key: str = "rerank_score") -> List[dict]:
"""Rerank candidates using cross-encoder.
Args:
query: The search query
candidates: List of dicts with at least a content_key field
content_key: Key to extract text from each candidate
top_k: Max results to return (0 = return all, reranked)
score_key: Key to store the cross-encoder score in each result
Returns:
Reranked list of candidates with score_key added.
Falls back to original order if cross-encoder is unavailable.
"""
if not candidates:
return candidates
ce = _get_cross_encoder()
if ce is None:
return candidates[:top_k] if top_k > 0 else candidates
pairs = []
for c in candidates:
text = c.get(content_key, "")
if not text:
text = str(c)
pairs.append((query, text[:512]))
try:
scores = ce.predict(pairs)
scored = list(zip(scores, candidates))
scored.sort(key=lambda x: x[0], reverse=True)
results = []
for score, candidate in scored:
candidate = dict(candidate)
candidate[score_key] = round(float(score), 4)
results.append(candidate)
if top_k > 0:
results = results[:top_k]
return results
except Exception as e:
logger.warning("Cross-encoder rerank failed: %s", e)
return candidates[:top_k] if top_k > 0 else candidates
def rerank_pairs(query: str, texts: List[str]) -> List[Tuple[float, int]]:
"""Score query-text pairs and return sorted (score, original_index)."""
ce = _get_cross_encoder()
if ce is None:
return [(0.0, i) for i in range(len(texts))]
pairs = [(query, t[:512]) for t in texts]
try:
scores = ce.predict(pairs)
indexed = [(float(s), i) for i, s in enumerate(scores)]
indexed.sort(key=lambda x: x[0], reverse=True)
return indexed
except Exception as e:
logger.warning("Cross-encoder scoring failed: %s", e)
return [(0.0, i) for i in range(len(texts))]