Coverage for agentos/rag/reranker.py: 0%

141 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-06 21:19 +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 

12 

13 

14@dataclass 

15class RerankConfig: 

16 """Configuration for reranking.""" 

17 

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 

24 

25 

26class Reranker: 

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

28 

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

34 

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 

39 

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. 

46 

47 Args: 

48 query: Original search query. 

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

50 

51 Returns: 

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

53 """ 

54 if not passages: 

55 return [] 

56 

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) 

66 

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) 

78 

79 if self._cross_encoder is None: 

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

81 

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

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

84 

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

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

87 p["rerank_method"] = "cross_encoder" 

88 

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

90 return passages[: self.config.top_n] 

91 

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. 

98 

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

100 """ 

101 if not passages: 

102 return [] 

103 

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

105 

106 # Tokenize for diversity computation 

107 token_sets = [] 

108 for t in texts: 

109 import re 

110 

111 tokens = set(re.findall(r"\w+", t.lower())) 

112 token_sets.append(tokens) 

113 

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

115 

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) 

120 

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) 

129 

130 selected = [] 

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

132 

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

134 best_idx = None 

135 best_score = -float("inf") 

136 

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 ) 

143 

144 score = ( 

145 self.config.diversity_lambda * relevance[idx] 

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

147 ) 

148 

149 if score > best_score: 

150 best_score = score 

151 best_idx = idx 

152 

153 if best_idx is not None: 

154 selected.append(best_idx) 

155 remaining.remove(best_idx) 

156 else: 

157 break 

158 

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" 

163 

164 return result 

165 

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. 

172 

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) 

178 

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 

185 

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

190 

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" 

202 

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

204 return passages[: self.config.top_n] 

205 

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 

213 

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

215 

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" 

224 

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

226 return passages[: self.config.top_n] 

227 

228 

229class DiversityRanker: 

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

231 

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

233 self.lambda_param = lambda_param 

234 

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

243 

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

245 import re 

246 

247 token_sets = [] 

248 for t in texts: 

249 token_sets.append(set(re.findall(r"\w+", t.lower()))) 

250 

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) 

255 

256 selected = [0] 

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

258 

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) 

275 

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