Coverage for agentos/rag/reranker.py: 0%
141 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-09 07:12 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-09 07:12 +0800
1"""Re-ranking for RAG pipeline.
3Cross-encoder and LLM-based reranking to refine retrieval results.
4Supports: cross-encoder (sentence-transformers), LLM reranking,
5and simple heuristic reranking (diversity, freshness).
6"""
8from __future__ import annotations
10from dataclasses import dataclass
11from typing import Any
14@dataclass
15class RerankConfig:
16 """Configuration for reranking."""
18 method: str = "cross_encoder" # cross_encoder | llm | diversity | mmr
19 model: str = "cross-encoder/ms-marco-MiniLM-L-6-v2"
20 top_n: int = 5 # number of results after reranking
21 diversity_lambda: float = 0.5 # MMR diversity weight
22 llm_prompt_template: str = "" # custom prompt for LLM reranker
23 batch_size: int = 8
26class Reranker:
27 """Re-rank retrieval results for improved relevance.
29 Methods:
30 - cross_encoder: Uses sentence-transformers cross-encoder for precision.
31 - mmr: Maximal Marginal Relevance for diversity.
32 - llm: Uses an LLM to score relevance of each passage.
33 """
35 def __init__(self, config: RerankConfig | None = None):
36 self.config = config or RerankConfig()
37 self._cross_encoder = None
38 self._embed_fn = None # for MMR diversity
40 async def rerank(
41 self,
42 query: str,
43 passages: list[dict[str, Any]],
44 ) -> list[dict[str, Any]]:
45 """Re-rank passages by relevance to query.
47 Args:
48 query: Original search query.
49 passages: List of dicts with 'text' and 'score' keys.
51 Returns:
52 Re-ranked list with updated 'rerank_score' key.
53 """
54 if not passages:
55 return []
57 if self.config.method == "cross_encoder":
58 return await self._cross_encode_rerank(query, passages)
59 elif self.config.method == "mmr":
60 return self._mmr_rerank(query, passages)
61 elif self.config.method == "llm":
62 return await self._llm_rerank(query, passages)
63 else:
64 # diversity: sort by text length variability as proxy
65 return self._diversity_rerank(passages)
67 async def _cross_encode_rerank(
68 self,
69 query: str,
70 passages: list[dict[str, Any]],
71 ) -> list[dict[str, Any]]:
72 """Use cross-encoder model for relevance scoring."""
73 try:
74 from sentence_transformers import CrossEncoder
75 except ImportError:
76 # Fallback: use simple heuristic based on term overlap
77 return self._fallback_rerank(query, passages)
79 if self._cross_encoder is None:
80 self._cross_encoder = CrossEncoder(self.config.model)
82 pairs = [(query, p["text"]) for p in passages]
83 scores = self._cross_encoder.predict(pairs, batch_size=self.config.batch_size)
85 for p, s in zip(passages, scores):
86 p["rerank_score"] = float(s)
87 p["rerank_method"] = "cross_encoder"
89 passages.sort(key=lambda x: x.get("rerank_score", 0), reverse=True)
90 return passages[: self.config.top_n]
92 def _mmr_rerank(
93 self,
94 query: str,
95 passages: list[dict[str, Any]],
96 ) -> list[dict[str, Any]]:
97 """Maximal Marginal Relevance: balance relevance with diversity.
99 Without embeddings, uses Jaccard similarity on token sets as proxy.
100 """
101 if not passages:
102 return []
104 texts = [p["text"] for p in passages]
106 # Tokenize for diversity computation
107 token_sets = []
108 for t in texts:
109 import re
111 tokens = set(re.findall(r"\w+", t.lower()))
112 token_sets.append(tokens)
114 query_tokens = set(re.findall(r"\w+", query.lower())) if query else set()
116 def _jaccard_sim(a: set, b: set) -> float:
117 if not a or not b:
118 return 0.0
119 return len(a & b) / len(a | b)
121 # Initial relevance scores (original scores or query overlap)
122 relevance = []
123 for i, (p, ts) in enumerate(zip(passages, token_sets)):
124 if query_tokens:
125 rel = _jaccard_sim(query_tokens, ts)
126 else:
127 rel = p.get("score", 0.0)
128 relevance.append(rel)
130 selected = []
131 remaining = list(range(len(passages)))
133 while remaining and len(selected) < self.config.top_n:
134 best_idx = None
135 best_score = -float("inf")
137 for idx in remaining:
138 diversity = (
139 min((1.0 - _jaccard_sim(token_sets[idx], token_sets[s])) for s in selected)
140 if selected
141 else 1.0
142 )
144 score = (
145 self.config.diversity_lambda * relevance[idx]
146 + (1 - self.config.diversity_lambda) * diversity
147 )
149 if score > best_score:
150 best_score = score
151 best_idx = idx
153 if best_idx is not None:
154 selected.append(best_idx)
155 remaining.remove(best_idx)
156 else:
157 break
159 result = [passages[i] for i in selected]
160 for i, p in enumerate(result):
161 p["rerank_score"] = relevance[selected[i]]
162 p["rerank_method"] = "mmr"
164 return result
166 async def _llm_rerank(
167 self,
168 query: str,
169 passages: list[dict[str, Any]],
170 ) -> list[dict[str, Any]]:
171 """LLM-based reranking: ask an LLM to score passage relevance.
173 Falls back to fallback heuristic if no LLM is configured.
174 """
175 # This is a framework hook — the actual LLM call is done by the caller
176 # by injecting an llm_call function or using the default heuristic
177 return self._fallback_rerank(query, passages)
179 def _diversity_rerank(
180 self,
181 passages: list[dict[str, Any]],
182 ) -> list[dict[str, Any]]:
183 """Simple diversity reranking: penalize similar-length passages."""
184 import re
186 texts = [p["text"] for p in passages]
187 token_sets = []
188 for t in texts:
189 token_sets.append(set(re.findall(r"\w+", t.lower())))
191 # Score: original score * diversity bonus (penalize similarity to higher-ranked)
192 for i, p in enumerate(passages):
193 diversity_penalty = 0.0
194 for j in range(i):
195 if token_sets[i] and token_sets[j]:
196 overlap = len(token_sets[i] & token_sets[j]) / len(
197 token_sets[i] | token_sets[j]
198 )
199 diversity_penalty += overlap * 0.1
200 p["rerank_score"] = p.get("score", 0.5) * (1.0 - min(diversity_penalty, 0.5))
201 p["rerank_method"] = "diversity"
203 passages.sort(key=lambda x: x.get("rerank_score", 0), reverse=True)
204 return passages[: self.config.top_n]
206 def _fallback_rerank(
207 self,
208 query: str,
209 passages: list[dict[str, Any]],
210 ) -> list[dict[str, Any]]:
211 """Fallback: simple term-overlap heuristic rerank."""
212 import re
214 query_tokens = set(re.findall(r"\w+", query.lower())) if query else set()
216 for p in passages:
217 text_tokens = set(re.findall(r"\w+", p["text"].lower()))
218 if query_tokens and text_tokens:
219 overlap = len(query_tokens & text_tokens) / max(len(query_tokens), 1)
220 p["rerank_score"] = p.get("score", 0.0) * 0.5 + overlap * 0.5
221 else:
222 p["rerank_score"] = p.get("score", 0.0)
223 p["rerank_method"] = "fallback_heuristic"
225 passages.sort(key=lambda x: x.get("rerank_score", 0), reverse=True)
226 return passages[: self.config.top_n]
229class DiversityRanker:
230 """Diversity-focused reranker for varied search results."""
232 def __init__(self, lambda_param: float = 0.6):
233 self.lambda_param = lambda_param
235 def rerank(
236 self,
237 passages: list[dict[str, Any]],
238 top_n: int = 5,
239 ) -> list[dict[str, Any]]:
240 """Maximize result diversity while keeping relevance high."""
241 if not passages:
242 return []
244 texts = [p["text"] for p in passages]
245 import re
247 token_sets = []
248 for t in texts:
249 token_sets.append(set(re.findall(r"\w+", t.lower())))
251 def sim(a: set, b: set) -> float:
252 if not a or not b:
253 return 0.0
254 return len(a & b) / len(a | b)
256 selected = [0]
257 remaining = set(range(1, len(passages)))
259 while len(selected) < min(top_n, len(passages)):
260 best_idx = -1
261 best_score = -float("inf")
262 for idx in remaining:
263 max_sim = max(sim(token_sets[idx], token_sets[s]) for s in selected)
264 score = (
265 self.lambda_param * passages[idx].get("score", 0.5)
266 - (1 - self.lambda_param) * max_sim
267 )
268 if score > best_score:
269 best_score = score
270 best_idx = idx
271 if best_idx < 0:
272 break
273 selected.append(best_idx)
274 remaining.remove(best_idx)
276 return [passages[i] for i in selected]