CoolFace
Apppublic

McQbis/document-intelligence-rag

sourceHugging Faceupdated 10d agoView on Hugging Face
0likes
beir_eval.py168 linesDownload Raw Back to evaluation
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