Coverage for agentos/rag/reranker.py: 0%
142 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +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, Dict, List, Optional
14@dataclass
15class RerankConfig:
16 """Configuration for reranking."""
17 method: str = "cross_encoder" # cross_encoder | llm | diversity | mmr
18 model: str = "cross-encoder/ms-marco-MiniLM-L-6-v2"
19 top_n: int = 5 # number of results after reranking
20 diversity_lambda: float = 0.5 # MMR diversity weight
21 llm_prompt_template: str = "" # custom prompt for LLM reranker
22 batch_size: int = 8
25class Reranker:
26 """Re-rank retrieval results for improved relevance.
28 Methods:
29 - cross_encoder: Uses sentence-transformers cross-encoder for precision.
30 - mmr: Maximal Marginal Relevance for diversity.
31 - llm: Uses an LLM to score relevance of each passage.
32 """
34 def __init__(self, config: Optional[RerankConfig] = None):
35 self.config = config or RerankConfig()
36 self._cross_encoder = None
37 self._embed_fn = None # for MMR diversity
39 async def rerank(
40 self,
41 query: str,
42 passages: List[Dict[str, Any]],
43 ) -> List[Dict[str, Any]]:
44 """Re-rank passages by relevance to query.
46 Args:
47 query: Original search query.
48 passages: List of dicts with 'text' and 'score' keys.
50 Returns:
51 Re-ranked list with updated 'rerank_score' key.
52 """
53 if not passages:
54 return []
56 if self.config.method == "cross_encoder":
57 return await self._cross_encode_rerank(query, passages)
58 elif self.config.method == "mmr":
59 return self._mmr_rerank(query, passages)
60 elif self.config.method == "llm":
61 return await self._llm_rerank(query, passages)
62 else:
63 # diversity: sort by text length variability as proxy
64 return self._diversity_rerank(passages)
66 async def _cross_encode_rerank(
67 self,
68 query: str,
69 passages: List[Dict[str, Any]],
70 ) -> List[Dict[str, Any]]:
71 """Use cross-encoder model for relevance scoring."""
72 try:
73 from sentence_transformers import CrossEncoder
74 except ImportError:
75 # Fallback: use simple heuristic based on term overlap
76 return self._fallback_rerank(query, passages)
78 if self._cross_encoder is None:
79 self._cross_encoder = CrossEncoder(self.config.model)
81 pairs = [(query, p["text"]) for p in passages]
82 scores = self._cross_encoder.predict(pairs, batch_size=self.config.batch_size)
84 for p, s in zip(passages, scores):
85 p["rerank_score"] = float(s)
86 p["rerank_method"] = "cross_encoder"
88 passages.sort(key=lambda x: x.get("rerank_score", 0), reverse=True)
89 return passages[: self.config.top_n]
91 def _mmr_rerank(
92 self,
93 query: str,
94 passages: List[Dict[str, Any]],
95 ) -> List[Dict[str, Any]]:
96 """Maximal Marginal Relevance: balance relevance with diversity.
98 Without embeddings, uses Jaccard similarity on token sets as proxy.
99 """
100 if not passages:
101 return []
103 texts = [p["text"] for p in passages]
105 # Tokenize for diversity computation
106 token_sets = []
107 for t in texts:
108 import re
109 tokens = set(re.findall(r'\w+', t.lower()))
110 token_sets.append(tokens)
112 query_tokens = set(re.findall(r'\w+', query.lower())) if query else set()
114 def _jaccard_sim(a: set, b: set) -> float:
115 if not a or not b:
116 return 0.0
117 return len(a & b) / len(a | b)
119 # Initial relevance scores (original scores or query overlap)
120 relevance = []
121 for i, (p, ts) in enumerate(zip(passages, token_sets)):
122 if query_tokens:
123 rel = _jaccard_sim(query_tokens, ts)
124 else:
125 rel = p.get("score", 0.0)
126 relevance.append(rel)
128 selected = []
129 remaining = list(range(len(passages)))
131 while remaining and len(selected) < self.config.top_n:
132 best_idx = None
133 best_score = -float("inf")
135 for idx in remaining:
136 diversity = min(
137 (1.0 - _jaccard_sim(token_sets[idx], token_sets[s]))
138 for s in selected
139 ) if selected else 1.0
141 score = (
142 self.config.diversity_lambda * relevance[idx]
143 + (1 - self.config.diversity_lambda) * diversity
144 )
146 if score > best_score:
147 best_score = score
148 best_idx = idx
150 if best_idx is not None:
151 selected.append(best_idx)
152 remaining.remove(best_idx)
153 else:
154 break
156 result = [passages[i] for i in selected]
157 for i, p in enumerate(result):
158 p["rerank_score"] = relevance[selected[i]]
159 p["rerank_method"] = "mmr"
161 return result
163 async def _llm_rerank(
164 self,
165 query: str,
166 passages: List[Dict[str, Any]],
167 ) -> List[Dict[str, Any]]:
168 """LLM-based reranking: ask an LLM to score passage relevance.
170 Falls back to fallback heuristic if no LLM is configured.
171 """
172 # This is a framework hook — the actual LLM call is done by the caller
173 # by injecting an llm_call function or using the default heuristic
174 return self._fallback_rerank(query, passages)
176 def _diversity_rerank(
177 self,
178 passages: List[Dict[str, Any]],
179 ) -> List[Dict[str, Any]]:
180 """Simple diversity reranking: penalize similar-length passages."""
181 import re
183 texts = [p["text"] for p in passages]
184 token_sets = []
185 for t in texts:
186 token_sets.append(set(re.findall(r'\w+', t.lower())))
188 # Score: original score * diversity bonus (penalize similarity to higher-ranked)
189 scored = []
190 for i, p in enumerate(passages):
191 diversity_penalty = 0.0
192 for j in range(i):
193 if token_sets[i] and token_sets[j]:
194 overlap = len(token_sets[i] & token_sets[j]) / len(token_sets[i] | token_sets[j])
195 diversity_penalty += overlap * 0.1
196 p["rerank_score"] = p.get("score", 0.5) * (1.0 - min(diversity_penalty, 0.5))
197 p["rerank_method"] = "diversity"
199 passages.sort(key=lambda x: x.get("rerank_score", 0), reverse=True)
200 return passages[: self.config.top_n]
202 def _fallback_rerank(
203 self,
204 query: str,
205 passages: List[Dict[str, Any]],
206 ) -> List[Dict[str, Any]]:
207 """Fallback: simple term-overlap heuristic rerank."""
208 import re
210 query_tokens = set(re.findall(r'\w+', query.lower())) if query else set()
212 for p in passages:
213 text_tokens = set(re.findall(r'\w+', p["text"].lower()))
214 if query_tokens and text_tokens:
215 overlap = len(query_tokens & text_tokens) / max(len(query_tokens), 1)
216 p["rerank_score"] = p.get("score", 0.0) * 0.5 + overlap * 0.5
217 else:
218 p["rerank_score"] = p.get("score", 0.0)
219 p["rerank_method"] = "fallback_heuristic"
221 passages.sort(key=lambda x: x.get("rerank_score", 0), reverse=True)
222 return passages[: self.config.top_n]
225class DiversityRanker:
226 """Diversity-focused reranker for varied search results."""
228 def __init__(self, lambda_param: float = 0.6):
229 self.lambda_param = lambda_param
231 def rerank(
232 self,
233 passages: List[Dict[str, Any]],
234 top_n: int = 5,
235 ) -> List[Dict[str, Any]]:
236 """Maximize result diversity while keeping relevance high."""
237 if not passages:
238 return []
240 texts = [p["text"] for p in passages]
241 import re
243 token_sets = []
244 for t in texts:
245 token_sets.append(set(re.findall(r'\w+', t.lower())))
247 def sim(a: set, b: set) -> float:
248 if not a or not b:
249 return 0.0
250 return len(a & b) / len(a | b)
252 selected = [0]
253 remaining = set(range(1, len(passages)))
255 while len(selected) < min(top_n, len(passages)):
256 best_idx = -1
257 best_score = -float("inf")
258 for idx in remaining:
259 max_sim = max(sim(token_sets[idx], token_sets[s]) for s in selected)
260 score = (
261 self.lambda_param * passages[idx].get("score", 0.5)
262 - (1 - self.lambda_param) * max_sim
263 )
264 if score > best_score:
265 best_score = score
266 best_idx = idx
267 if best_idx < 0:
268 break
269 selected.append(best_idx)
270 remaining.remove(best_idx)
272 return [passages[i] for i in selected]