| """ |
| Evaluate retrieval performance (recall@K, hit-rate) for the RAG engine. |
| |
| Expected ground-truth file format (JSONL or JSON list): |
| - Each item is an object with keys: "query" (str) and "relevant_ids" (list[int] or list[str]) |
| - `relevant_ids` should match the index positions in `data/chunks.pkl` or a stable id present in the chunk dicts. |
| |
| Usage: |
| python evaluate_retrieval.py --gt ground_truth.jsonl --k 10 |
| |
| If your chunk metadata uses custom IDs (e.g., a `chunk_id` field), pass `--id-field chunk_id`. |
| """ |
| from __future__ import annotations |
| import argparse |
| import json |
| import os |
| import pickle |
| import logging |
| from typing import List, Any |
| import math |
| from collections import defaultdict |
|
|
| from rag_engine import RAGEngine |
| from config import CHUNKS_PATH |
|
|
| logging.basicConfig(level=logging.INFO) |
| logger = logging.getLogger(__name__) |
|
|
|
|
| def load_ground_truth(path: str) -> List[dict]: |
| if not os.path.exists(path): |
| raise FileNotFoundError(path) |
| with open(path, "r", encoding="utf-8") as f: |
| txt = f.read().strip() |
| if not txt: |
| return [] |
| try: |
| data = json.loads(txt) |
| except json.JSONDecodeError: |
| |
| data = [json.loads(l) for l in txt.splitlines() if l.strip()] |
| return data |
|
|
|
|
| def map_chunk_ids(chunks: List[dict], id_field: str = None) -> dict: |
| """Return a mapping from id -> index in chunks list. |
| |
| If `id_field` is None, uses the integer index as the id. |
| """ |
| mapping = {} |
| if id_field: |
| for i, c in enumerate(chunks): |
| if id_field in c: |
| mapping[c[id_field]] = i |
| else: |
| for i, c in enumerate(chunks): |
| mapping[i] = i |
| return mapping |
|
|
|
|
| def recall_at_k(rag: RAGEngine, gt: List[dict], k: int = 10, id_field: str | None = None) -> dict: |
| """Compute Recall@K, Precision@K, MRR, and nDCG@K with optional per-language breakdown. |
| |
| Ground-truth entries may include a `lang` field (e.g. 'en' or 'hinglish'). |
| Returns a stats dict with overall metrics and a `per_language` mapping. |
| """ |
| |
| rag.initialize() |
|
|
| |
| with open(CHUNKS_PATH, "rb") as f: |
| chunks = pickle.load(f) |
|
|
| id2idx = map_chunk_ids(chunks, id_field) |
|
|
| total = 0 |
| hits = 0 |
| sum_precision = 0.0 |
| sum_rr = 0.0 |
| sum_ndcg = 0.0 |
|
|
| per_language = defaultdict(lambda: {"total": 0, "hits": 0, "sum_precision": 0.0, "sum_rr": 0.0, "sum_ndcg": 0.0}) |
|
|
| for item in gt: |
| query = item.get("query") or item.get("q") |
| relevant = item.get("relevant_ids") or item.get("relevant") |
| lang = item.get("lang", "en") |
| if not query or not relevant: |
| continue |
| total += 1 |
| per_language[lang]["total"] += 1 |
|
|
| |
| relevant_idxs = set() |
| for r in relevant: |
| if id_field and r in id2idx: |
| relevant_idxs.add(id2idx[r]) |
| else: |
| try: |
| relevant_idxs.add(int(r)) |
| except Exception: |
| pass |
|
|
| |
| results = rag.retrieve(query, top_k=k) |
| retrieved_order = [] |
| for r in results: |
| idx = None |
| |
| if id_field and isinstance(r, dict) and id_field in r and r[id_field] in id2idx: |
| idx = id2idx[r[id_field]] |
|
|
| |
| if idx is None and isinstance(r, dict) and "text" in r: |
| r_text = r.get("text") |
| for i, c in enumerate(chunks): |
| try: |
| c_text = c.get("text") if isinstance(c, dict) else str(c) |
| except Exception: |
| c_text = str(c) |
| if c_text == r_text: |
| idx = i |
| break |
|
|
| |
| if idx is None: |
| try: |
| idx = chunks.index(r) |
| except Exception: |
| idx = None |
|
|
| if idx is not None and idx not in retrieved_order: |
| retrieved_order.append(idx) |
|
|
| |
| top_k_list = retrieved_order[:k] |
| hits_k = sum(1 for idx in top_k_list if idx in relevant_idxs) |
| precision_k = hits_k / k if k > 0 else 0.0 |
| hit_flag = hits_k > 0 |
| if hit_flag: |
| hits += 1 |
| per_language[lang]["hits"] += 1 |
|
|
| |
| rr = 0.0 |
| for pos, idx in enumerate(retrieved_order, start=1): |
| if idx in relevant_idxs: |
| rr = 1.0 / pos |
| break |
|
|
| |
| dcg = 0.0 |
| for pos, idx in enumerate(top_k_list, start=1): |
| if idx in relevant_idxs: |
| dcg += 1.0 / math.log2(pos + 1) |
|
|
| ideal_rel = min(len(relevant_idxs), k) |
| idcg = sum(1.0 / math.log2(i + 1) for i in range(1, ideal_rel + 1)) if ideal_rel > 0 else 0.0 |
| ndcg = dcg / idcg if idcg > 0 else 0.0 |
|
|
| sum_precision += precision_k |
| sum_rr += rr |
| sum_ndcg += ndcg |
|
|
| per_language[lang]["sum_precision"] += precision_k |
| per_language[lang]["sum_rr"] += rr |
| per_language[lang]["sum_ndcg"] += ndcg |
|
|
| |
| recall = hits / total if total else 0.0 |
| precision_at_k = sum_precision / total if total else 0.0 |
| mrr = sum_rr / total if total else 0.0 |
| ndcg = sum_ndcg / total if total else 0.0 |
|
|
| |
| per_lang_stats = {} |
| for lang, vals in per_language.items(): |
| t = vals["total"] |
| per_lang_stats[lang] = { |
| "total": t, |
| "recall_at_k": (vals["hits"] / t) if t else 0.0, |
| "precision_at_k": (vals["sum_precision"] / t) if t else 0.0, |
| "mrr": (vals["sum_rr"] / t) if t else 0.0, |
| "ndcg": (vals["sum_ndcg"] / t) if t else 0.0, |
| } |
|
|
| return { |
| "recall_at_k": recall, |
| "precision_at_k": precision_at_k, |
| "mrr": mrr, |
| "ndcg": ndcg, |
| "total_queries": total, |
| "hits": hits, |
| "per_language": per_lang_stats, |
| } |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--gt", required=True, help="Ground truth JSON/JSONL file") |
| parser.add_argument("--k", type=int, default=10, help="Top-K to evaluate") |
| parser.add_argument("--id-field", help="Field name in chunk dicts used as stable id") |
| args = parser.parse_args() |
|
|
| gt = load_ground_truth(args.gt) |
| rag = RAGEngine() |
| stats = recall_at_k(rag, gt, args.k, args.id_field) |
|
|
| logger.info(f"Recall@{args.k}: {stats['recall_at_k']:.4f} ({stats['hits']}/{stats['total_queries']})") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|