From a4839c0ad2b02d51aa746de4951e8f727c3d8133 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Thu, 6 Aug 2026 11:02:21 +0300 Subject: [PATCH] Score retrieval evals from search results --- evaluations/evaluations/benchmark.py | 11 +--- evaluations/tests/test_benchmark.py | 87 ++++++++++++++++++++++++++++ 2 files changed, 89 insertions(+), 9 deletions(-) diff --git a/evaluations/evaluations/benchmark.py b/evaluations/evaluations/benchmark.py index 2541ca05..1249660c 100644 --- a/evaluations/evaluations/benchmark.py +++ b/evaluations/evaluations/benchmark.py @@ -230,20 +230,13 @@ async def run_retrieval_benchmark( async with HaikuRAG(db, config=config, read_only=True) as rag: async def retrieval_target(question: str) -> list[str]: - chunks = await rag.search(query=question, limit=5) + chunks = await rag.search(query=question, limit=5, include_images=False) seen = set() identifiers = [] for result in chunks: - if result.document_id is None: - continue - doc = await rag.get_document_by_id(result.document_id) - if doc is None: - continue # Use arxiv_id from metadata if present, otherwise use URI - doc_id = doc.metadata.get("arxiv_id") if doc.metadata else None - if doc_id is None: - doc_id = doc.uri + doc_id = result.document_meta.get("arxiv_id") or result.document_uri if doc_id and doc_id not in seen: identifiers.append(doc_id) seen.add(doc_id) diff --git a/evaluations/tests/test_benchmark.py b/evaluations/tests/test_benchmark.py index dcdc27da..782a9c5e 100644 --- a/evaluations/tests/test_benchmark.py +++ b/evaluations/tests/test_benchmark.py @@ -455,6 +455,93 @@ class TestLoadCaseIds: assert _load_case_ids(None) is None +class TestRetrievalTarget: + def _spec(self) -> DatasetSpec: + from evaluations.config import RetrievalSample + from evaluations.evaluators import MAPEvaluator + + return DatasetSpec( + key="test", + db_filename="test.lancedb", + document_loader=lambda: None, # type: ignore[arg-type] # ty: ignore[invalid-argument-type] + document_mapper=lambda doc: None, + qa_loader=lambda: [], # type: ignore[arg-type] # ty: ignore[invalid-argument-type] + qa_case_builder=lambda idx, doc: None, # type: ignore[arg-type] # ty: ignore[invalid-argument-type] + retrieval_loader=lambda: [ # type: ignore[arg-type] # ty: ignore[invalid-argument-type] + {"q": "What is X?", "uris": ("uri-x",)}, + ], + retrieval_mapper=lambda d: RetrievalSample( + question=d["q"], expected_uris=d["uris"] + ), + retrieval_evaluator=MAPEvaluator(), + ) + + @pytest.mark.asyncio + async def test_scores_from_search_results_without_reading_documents( + self, tmp_path: Path + ) -> None: + from haiku.rag.store.models.chunk import SearchResult + + from evaluations.benchmark import run_retrieval_benchmark + + searches: list[dict] = [] + + class FakeRag: + async def search(self, **kwargs) -> list[SearchResult]: + searches.append(kwargs) + return [ + SearchResult( + content="x", + score=1.0, + document_id="doc-1", + document_uri="uri-x", + ) + ] + + async def get_document_by_id(self, document_id: str) -> None: + raise AssertionError( + "retrieval scoring must not read whole document rows" + ) + + fake = FakeRag() + with patch("evaluations.benchmark.HaikuRAG") as mock_haiku: + mock_haiku.return_value.__aenter__.return_value = fake + result = await run_retrieval_benchmark( + self._spec(), AppConfig(), db_path=tmp_path / "test.lancedb" + ) + + assert result is not None + assert result["map"] == 1.0 + assert searches[0]["include_images"] is False + + @pytest.mark.asyncio + async def test_prefers_metadata_identifier_over_uri(self, tmp_path: Path) -> None: + from haiku.rag.store.models.chunk import SearchResult + + from evaluations.benchmark import run_retrieval_benchmark + + class FakeRag: + async def search(self, **kwargs) -> list[SearchResult]: + return [ + SearchResult( + content="x", + score=1.0, + document_id="doc-1", + document_uri="uri-other", + document_meta={"arxiv_id": "uri-x"}, + ) + ] + + with patch("evaluations.benchmark.HaikuRAG") as mock_haiku: + mock_haiku.return_value.__aenter__.return_value = FakeRag() + result = await run_retrieval_benchmark( + self._spec(), AppConfig(), db_path=tmp_path / "test.lancedb" + ) + + assert result is not None + assert result["map"] == 1.0 + + class TestEvaluateDatasetCaseIds: def _spec(self) -> DatasetSpec: return DatasetSpec(