McQbis/document-intelligence-rag
0
1from __future__ import annotations2 3import os4from dataclasses import dataclass, field5from typing import Dict, List, Optional6 7from tqdm import tqdm8 9 10@dataclass11class EvalResults:12 """BEIR evaluation summary for a single dataset/split."""13 14 dataset: str15 split: str16 top_k: int17 queries_evaluated: int18 recall_at_1: float19 recall_at_5: float20 recall_at_10: float21 ndcg_at_10: float = 0.022 extra: dict = field(default_factory=dict)23 24 def __str__(self) -> str:25 lines = [26 f"\n{'='*40}",27 f"BEIR Evaluation — {self.dataset} ({self.split})",28 f"{'='*40}",29 f"Queries evaluated : {self.queries_evaluated}",30 f"Recall@1 : {self.recall_at_1:.4f}",31 f"Recall@5 : {self.recall_at_5:.4f}",32 f"Recall@10 : {self.recall_at_10:.4f}",33 f"nDCG@10 : {self.ndcg_at_10:.4f}",34 f"{'='*40}",35 ]36 return "\n".join(lines)37 38 39class BEIREvaluator:40 BEIR_BASE_URL = (41 "https://public.ukp.informatik.tu-darmstadt.de/"42 "thakur/BEIR/datasets/{dataset}.zip"43 )44 45 def __init__(46 self,47 retriever, # HybridRetriever or QueryRouter48 data_dir: str = "beir-data",49 query_prefix: str = "Represent this sentence for searching relevant passages: ",50 ):51 self.retriever = retriever52 self.data_dir = data_dir53 self.query_prefix = query_prefix54 55 56 def run(57 self,58 dataset: str,59 split: str = "test",60 top_k: int = 10,61 max_queries: Optional[int] = None,62 ) -> EvalResults:63 """Download dataset (if needed), build index, evaluate."""64 from beir import util65 from beir.datasets.data_loader import GenericDataLoader66 67 print(f"[beir] Downloading {dataset}…")68 url = self.BEIR_BASE_URL.format(dataset=dataset)69 data_path = util.download_and_unzip(url, self.data_dir)70 71 print(f"[beir] Loading {dataset}/{split}")72 corpus, queries, qrels = GenericDataLoader(73 data_folder=os.path.join(data_path)74 ).load(split=split)75 print(f"[beir] corpus={len(corpus):,} queries={len(queries):,}")76 77 chunks, doc_id_mapping = self._corpus_to_chunks(corpus)78 chunk_to_doc: Dict[int, str] = {79 id(c): doc_id_mapping[i] for i, c in enumerate(chunks)80 }81 82 print("[beir] Building index…")83 self.retriever.build_index(chunks)84 print("[beir] Index ready.")85 86 hits_1 = hits_5 = hits_10 = 087 ndcg_sum = 0.088 evaluated = 089 90 query_items = list(queries.items())91 if max_queries:92 query_items = query_items[:max_queries]93 94 for query_id, query_text in tqdm(query_items, desc="eval"):95 relevant = set(qrels[query_id].keys())96 if not relevant:97 continue98 evaluated += 199 100 # instruction-style prefix improves embedding models (e.g., BGE)101 prefixed = self.query_prefix + query_text102 results = self._search(prefixed, top_k=top_k)103 retrieved = [chunk_to_doc[id(chunk)] for chunk, _ in results]104 105 if any(d in relevant for d in retrieved[:1]):106 hits_1 += 1107 if any(d in relevant for d in retrieved[:5]):108 hits_5 += 1109 if any(d in relevant for d in retrieved[:10]):110 hits_10 += 1111 112 ndcg_sum += self._ndcg_at_k(retrieved, relevant, k=10)113 114 return EvalResults(115 dataset=dataset,116 split=split,117 top_k=top_k,118 queries_evaluated=evaluated,119 recall_at_1=hits_1 / evaluated,120 recall_at_5=hits_5 / evaluated,121 recall_at_10=hits_10 / evaluated,122 ndcg_at_10=ndcg_sum / evaluated,123 )124 125 126 def _search(self, query: str, top_k: int):127 if hasattr(self.retriever, "search"):128 return self.retriever.search(query, top_k=top_k)129 raise TypeError(f"Unsupported retriever type: {type(self.retriever)}")130 131 @staticmethod132 def _corpus_to_chunks(corpus: dict):133 """Convert BEIR corpus into unified chunk format."""134 135 class _Chunk:136 def __init__(self, text: str):137 self.text = text138 self.page = 0139 self.source = ""140 self.file_type = "beir"141 self.chunk_index = 0142 self.metadata: dict = {}143 144 chunks = []145 doc_ids = []146 for doc_id, doc in corpus.items():147 parts = []148 if doc.get("title"):149 parts.append(doc["title"])150 if doc.get("text"):151 parts.append(doc["text"])152 chunks.append(_Chunk(text="\n".join(parts)))153 doc_ids.append(doc_id)154 155 return chunks, doc_ids156 157 @staticmethod158 def _ndcg_at_k(retrieved: List[str], relevant: set, k: int) -> float:159 import math160 161 dcg = sum(162 1.0 / math.log2(i + 2)163 for i, doc_id in enumerate(retrieved[:k])164 if doc_id in relevant165 )166 ideal_hits = min(len(relevant), k)167 idcg = sum(1.0 / math.log2(i + 2) for i in range(ideal_hits))168 return dcg / idcg if idcg > 0 else 0.0