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

1"""Re-ranking for RAG pipeline. 

2 

3Cross-encoder and LLM-based reranking to refine retrieval results. 

4Supports: cross-encoder (sentence-transformers), LLM reranking, 

5and simple heuristic reranking (diversity, freshness). 

6""" 

7 

8from __future__ import annotations 

9 

10from dataclasses import dataclass 

11from typing import Any, Dict, List, Optional 

12 

13 

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 

23 

24 

25class Reranker: 

26 """Re-rank retrieval results for improved relevance. 

27 

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 """ 

33 

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 

38 

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. 

45 

46 Args: 

47 query: Original search query. 

48 passages: List of dicts with 'text' and 'score' keys. 

49 

50 Returns: 

51 Re-ranked list with updated 'rerank_score' key. 

52 """ 

53 if not passages: 

54 return [] 

55 

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) 

65 

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) 

77 

78 if self._cross_encoder is None: 

79 self._cross_encoder = CrossEncoder(self.config.model) 

80 

81 pairs = [(query, p["text"]) for p in passages] 

82 scores = self._cross_encoder.predict(pairs, batch_size=self.config.batch_size) 

83 

84 for p, s in zip(passages, scores): 

85 p["rerank_score"] = float(s) 

86 p["rerank_method"] = "cross_encoder" 

87 

88 passages.sort(key=lambda x: x.get("rerank_score", 0), reverse=True) 

89 return passages[: self.config.top_n] 

90 

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. 

97 

98 Without embeddings, uses Jaccard similarity on token sets as proxy. 

99 """ 

100 if not passages: 

101 return [] 

102 

103 texts = [p["text"] for p in passages] 

104 

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) 

111 

112 query_tokens = set(re.findall(r'\w+', query.lower())) if query else set() 

113 

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) 

118 

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) 

127 

128 selected = [] 

129 remaining = list(range(len(passages))) 

130 

131 while remaining and len(selected) < self.config.top_n: 

132 best_idx = None 

133 best_score = -float("inf") 

134 

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 

140 

141 score = ( 

142 self.config.diversity_lambda * relevance[idx] 

143 + (1 - self.config.diversity_lambda) * diversity 

144 ) 

145 

146 if score > best_score: 

147 best_score = score 

148 best_idx = idx 

149 

150 if best_idx is not None: 

151 selected.append(best_idx) 

152 remaining.remove(best_idx) 

153 else: 

154 break 

155 

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" 

160 

161 return result 

162 

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. 

169 

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) 

175 

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 

182 

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()))) 

187 

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" 

198 

199 passages.sort(key=lambda x: x.get("rerank_score", 0), reverse=True) 

200 return passages[: self.config.top_n] 

201 

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 

209 

210 query_tokens = set(re.findall(r'\w+', query.lower())) if query else set() 

211 

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" 

220 

221 passages.sort(key=lambda x: x.get("rerank_score", 0), reverse=True) 

222 return passages[: self.config.top_n] 

223 

224 

225class DiversityRanker: 

226 """Diversity-focused reranker for varied search results.""" 

227 

228 def __init__(self, lambda_param: float = 0.6): 

229 self.lambda_param = lambda_param 

230 

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 [] 

239 

240 texts = [p["text"] for p in passages] 

241 import re 

242 

243 token_sets = [] 

244 for t in texts: 

245 token_sets.append(set(re.findall(r'\w+', t.lower()))) 

246 

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) 

251 

252 selected = [0] 

253 remaining = set(range(1, len(passages))) 

254 

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) 

271 

272 return [passages[i] for i in selected]