Metafazer/finrag-backend
0
1"""Retrieval evaluation harness for measuring search quality.2 3Provides metrics to evaluate retrieval independently from generation,4so we can catch retrieval failures before they pollute LLM answers.5 6Supported metrics:7- Precision@k: fraction of retrieved docs that are relevant8- Recall@k: fraction of relevant docs that were retrieved9- MRR (Mean Reciprocal Rank): average 1/rank of first relevant doc10- Hit Rate@k: fraction of queries where at least one relevant doc is in top-k11- NDCG@k: Normalized Discounted Cumulative Gain (position-aware relevance)12 13Design decisions:14- Evaluation dataset is a list of (query, relevant_chunk_ids) pairs.15 This is the minimal unit for retrieval evaluation.16- Metrics are computed per-query and then averaged (macro averaging).17- JSON-serializable evaluation dataset format for versioning in git.18- Retriever-agnostic: works with any callable that returns ranked dicts.19 20Debt: DAY-6-002 โ Golden evaluation dataset is hand-crafted and small.21 Day 13 will build a proper 50+ Q/A pair dataset with RAGAS.22"""23 24import json25import math26from collections.abc import Callable27from dataclasses import dataclass, field28from pathlib import Path29 30import structlog31 32logger = structlog.get_logger(__name__)33 34 35# --------------------------------------------------------------------------- #36# Data Structures37# --------------------------------------------------------------------------- #38 39 40@dataclass41class EvalQuery:42 """A single evaluation query with known relevant chunk IDs.43 44 Args:45 query: The natural language query string.46 relevant_chunk_ids: Set of chunk IDs that are relevant answers.47 metadata: Optional metadata for categorizing eval results48 (e.g., query type, difficulty).49 """50 51 query: str52 relevant_chunk_ids: set[str]53 metadata: dict = field(default_factory=dict)54 55 56@dataclass57class QueryResult:58 """Evaluation result for a single query.59 60 Args:61 query: The query string.62 precision_at_k: Precision@k score.63 recall_at_k: Recall@k score.64 reciprocal_rank: 1/rank of first relevant result (0 if none).65 hit: Whether any relevant doc was in top-k.66 ndcg_at_k: NDCG@k score.67 retrieved_ids: List of retrieved chunk IDs in order.68 relevant_ids: Set of relevant chunk IDs.69 """70 71 query: str72 precision_at_k: float73 recall_at_k: float74 reciprocal_rank: float75 hit: bool76 ndcg_at_k: float77 retrieved_ids: list[str]78 relevant_ids: set[str]79 80 81@dataclass82class EvalReport:83 """Aggregated evaluation report across all queries.84 85 Args:86 k: The k value used for evaluation.87 num_queries: Total number of queries evaluated.88 mean_precision: Mean Precision@k across all queries.89 mean_recall: Mean Recall@k across all queries.90 mrr: Mean Reciprocal Rank across all queries.91 hit_rate: Fraction of queries with at least one hit in top-k.92 mean_ndcg: Mean NDCG@k across all queries.93 per_query: Individual QueryResult for each query.94 """95 96 k: int97 num_queries: int98 mean_precision: float99 mean_recall: float100 mrr: float101 hit_rate: float102 mean_ndcg: float103 per_query: list[QueryResult]104 105 106# --------------------------------------------------------------------------- #107# Metric Functions108# --------------------------------------------------------------------------- #109 110 111def precision_at_k(retrieved_ids: list[str], relevant_ids: set[str], k: int) -> float:112 """Compute Precision@k.113 114 Fraction of retrieved documents (top-k) that are relevant.115 116 Args:117 retrieved_ids: Ordered list of retrieved chunk IDs.118 relevant_ids: Set of known relevant chunk IDs.119 k: Cutoff rank.120 121 Returns:122 Precision@k in [0, 1].123 """124 if k == 0:125 return 0.0126 top_k = retrieved_ids[:k]127 relevant_in_top_k = sum(1 for cid in top_k if cid in relevant_ids)128 return relevant_in_top_k / k129 130 131def recall_at_k(retrieved_ids: list[str], relevant_ids: set[str], k: int) -> float:132 """Compute Recall@k.133 134 Fraction of all relevant documents that appear in top-k.135 136 Args:137 retrieved_ids: Ordered list of retrieved chunk IDs.138 relevant_ids: Set of known relevant chunk IDs.139 k: Cutoff rank.140 141 Returns:142 Recall@k in [0, 1].143 """144 if not relevant_ids:145 return 0.0146 top_k = retrieved_ids[:k]147 relevant_in_top_k = sum(1 for cid in top_k if cid in relevant_ids)148 return relevant_in_top_k / len(relevant_ids)149 150 151def reciprocal_rank(retrieved_ids: list[str], relevant_ids: set[str]) -> float:152 """Compute Reciprocal Rank.153 154 1 / rank of the first relevant result. Returns 0 if no155 relevant result is found.156 157 Args:158 retrieved_ids: Ordered list of retrieved chunk IDs.159 relevant_ids: Set of known relevant chunk IDs.160 161 Returns:162 Reciprocal rank in (0, 1] or 0 if no hit.163 """164 for rank, cid in enumerate(retrieved_ids, start=1):165 if cid in relevant_ids:166 return 1.0 / rank167 return 0.0168 169 170def ndcg_at_k(retrieved_ids: list[str], relevant_ids: set[str], k: int) -> float:171 """Compute Normalized Discounted Cumulative Gain at k.172 173 Uses binary relevance (1 if relevant, 0 otherwise).174 Rewards placing relevant documents at higher ranks.175 176 Args:177 retrieved_ids: Ordered list of retrieved chunk IDs.178 relevant_ids: Set of known relevant chunk IDs.179 k: Cutoff rank.180 181 Returns:182 NDCG@k in [0, 1].183 """184 if not relevant_ids or k == 0:185 return 0.0186 187 top_k = retrieved_ids[:k]188 189 # DCG: sum of relevance / log2(rank + 1)190 dcg = 0.0191 for rank, cid in enumerate(top_k, start=1):192 if cid in relevant_ids:193 dcg += 1.0 / math.log2(rank + 1)194 195 # Ideal DCG: all relevant docs at top ranks196 ideal_k = min(len(relevant_ids), k)197 idcg = sum(1.0 / math.log2(rank + 1) for rank in range(1, ideal_k + 1))198 199 if idcg == 0:200 return 0.0201 202 return dcg / idcg203 204 205# --------------------------------------------------------------------------- #206# RetrievalEvaluator207# --------------------------------------------------------------------------- #208 209 210class RetrievalEvaluator:211 """Evaluator for retrieval quality.212 213 Runs a set of evaluation queries against a retriever function214 and computes standard IR metrics.215 216 Args:217 retriever_fn: A callable that takes (query, n_results) and218 returns a list of result dicts with 'chunk_id' keys.219 k: The cutoff rank for evaluation metrics.220 """221 222 def __init__(223 self,224 retriever_fn: Callable[[str, int], list[dict]],225 k: int = 5,226 ) -> None:227 """Initialize the evaluator.228 229 Args:230 retriever_fn: Retriever callable: (query, n_results) -> results.231 k: Cutoff rank for @k metrics (default 5).232 """233 self._retriever_fn = retriever_fn234 self._k = k235 236 def evaluate(self, eval_queries: list[EvalQuery]) -> EvalReport:237 """Run evaluation across all queries and compute metrics.238 239 Args:240 eval_queries: List of EvalQuery with known relevant docs.241 242 Returns:243 EvalReport with aggregated and per-query metrics.244 """245 if not eval_queries:246 return EvalReport(247 k=self._k,248 num_queries=0,249 mean_precision=0.0,250 mean_recall=0.0,251 mrr=0.0,252 hit_rate=0.0,253 mean_ndcg=0.0,254 per_query=[],255 )256 257 per_query_results: list[QueryResult] = []258 259 for eq in eval_queries:260 # Run retrieval261 results = self._retriever_fn(eq.query, self._k)262 retrieved_ids = [r["chunk_id"] for r in results]263 264 # Compute metrics265 p_at_k = precision_at_k(retrieved_ids, eq.relevant_chunk_ids, self._k)266 r_at_k = recall_at_k(retrieved_ids, eq.relevant_chunk_ids, self._k)267 rr = reciprocal_rank(retrieved_ids, eq.relevant_chunk_ids)268 hit = rr > 0269 ndcg = ndcg_at_k(retrieved_ids, eq.relevant_chunk_ids, self._k)270 271 qr = QueryResult(272 query=eq.query,273 precision_at_k=p_at_k,274 recall_at_k=r_at_k,275 reciprocal_rank=rr,276 hit=hit,277 ndcg_at_k=ndcg,278 retrieved_ids=retrieved_ids,279 relevant_ids=eq.relevant_chunk_ids,280 )281 per_query_results.append(qr)282 283 logger.debug(284 "eval_query_complete",285 query=eq.query[:60],286 precision=f"{p_at_k:.3f}",287 recall=f"{r_at_k:.3f}",288 rr=f"{rr:.3f}",289 hit=hit,290 )291 292 # Aggregate293 n = len(per_query_results)294 report = EvalReport(295 k=self._k,296 num_queries=n,297 mean_precision=sum(q.precision_at_k for q in per_query_results) / n,298 mean_recall=sum(q.recall_at_k for q in per_query_results) / n,299 mrr=sum(q.reciprocal_rank for q in per_query_results) / n,300 hit_rate=sum(1 for q in per_query_results if q.hit) / n,301 mean_ndcg=sum(q.ndcg_at_k for q in per_query_results) / n,302 per_query=per_query_results,303 )304 305 logger.info(306 "eval_complete",307 k=self._k,308 num_queries=n,309 mean_precision=f"{report.mean_precision:.3f}",310 mean_recall=f"{report.mean_recall:.3f}",311 mrr=f"{report.mrr:.3f}",312 hit_rate=f"{report.hit_rate:.3f}",313 mean_ndcg=f"{report.mean_ndcg:.3f}",314 )315 316 return report317 318 319# --------------------------------------------------------------------------- #320# Dataset I/O321# --------------------------------------------------------------------------- #322 323 324def load_eval_dataset(path: Path) -> list[EvalQuery]:325 """Load evaluation queries from a JSON file.326 327 Expected format:328 [329 {330 "query": "What was Apple's revenue in 2024?",331 "relevant_chunk_ids": ["aapl_rev_001", "aapl_rev_002"],332 "metadata": {"category": "numerical_extraction"}333 },334 ...335 ]336 337 Args:338 path: Path to the JSON eval dataset.339 340 Returns:341 List of EvalQuery objects.342 343 Raises:344 FileNotFoundError: If the dataset file doesn't exist.345 """346 if not path.exists():347 msg = f"Eval dataset not found: {path}"348 raise FileNotFoundError(msg)349 350 with open(path) as f:351 data = json.load(f)352 353 queries = [354 EvalQuery(355 query=item["query"],356 relevant_chunk_ids=set(item["relevant_chunk_ids"]),357 metadata=item.get("metadata", {}),358 )359 for item in data360 ]361 362 logger.info("eval_dataset_loaded", path=str(path), num_queries=len(queries))363 return queries364 365 366def save_eval_report(report: EvalReport, path: Path) -> None:367 """Save an evaluation report to JSON.368 369 Args:370 report: The EvalReport to save.371 path: Output file path.372 """373 path.parent.mkdir(parents=True, exist_ok=True)374 375 data = {376 "k": report.k,377 "num_queries": report.num_queries,378 "mean_precision": report.mean_precision,379 "mean_recall": report.mean_recall,380 "mrr": report.mrr,381 "hit_rate": report.hit_rate,382 "mean_ndcg": report.mean_ndcg,383 "per_query": [384 {385 "query": qr.query,386 "precision_at_k": qr.precision_at_k,387 "recall_at_k": qr.recall_at_k,388 "reciprocal_rank": qr.reciprocal_rank,389 "hit": qr.hit,390 "ndcg_at_k": qr.ndcg_at_k,391 "retrieved_ids": qr.retrieved_ids,392 "relevant_ids": sorted(qr.relevant_ids),393 }394 for qr in report.per_query395 ],396 }397 398 with open(path, "w") as f:399 json.dump(data, f, indent=2)400 401 logger.info("eval_report_saved", path=str(path))402 