CoolFace
Apppublic

Metafazer/finrag-backend

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
eval_harness.py402 linesDownload Raw Back to retrieval
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