From 57e86d00fd54aa69495706088899f0c2c778f1d8 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Fri, 26 Sep 2025 13:49:05 +0300 Subject: [PATCH] Refactor evaluations so that we can perform with multiple datasets.. Introduce Wix dataset. --- .gitignore | 2 +- docs/benchmarks.md | 5 +- evaluations/__init__.py | 0 .../benchmark.py | 200 +++++++++++------- evaluations/config.py | 46 ++++ evaluations/datasets/__init__.py | 8 + evaluations/datasets/repliqa.py | 58 +++++ evaluations/datasets/wix.py | 81 +++++++ {tests => evaluations}/llm_judge.py | 0 tests/test_qa.py | 3 +- 10 files changed, 325 insertions(+), 78 deletions(-) create mode 100644 evaluations/__init__.py rename tests/generate_benchmark_db.py => evaluations/benchmark.py (57%) create mode 100644 evaluations/config.py create mode 100644 evaluations/datasets/__init__.py create mode 100644 evaluations/datasets/repliqa.py create mode 100644 evaluations/datasets/wix.py rename {tests => evaluations}/llm_judge.py (100%) diff --git a/.gitignore b/.gitignore index 1da3f5ad..2aacff62 100644 --- a/.gitignore +++ b/.gitignore @@ -11,7 +11,7 @@ wheels/ # tests .coverage* -tests/data/ +evaluations/data/ .pytest_cache/ .ruff_cache/ diff --git a/docs/benchmarks.md b/docs/benchmarks.md index 9705e677..2d828fa4 100644 --- a/docs/benchmarks.md +++ b/docs/benchmarks.md @@ -2,8 +2,9 @@ We use the [repliqa](https://huggingface.co/datasets/ServiceNow/repliqa) dataset for the evaluation of `haiku.rag`. -You can perform your own evaluations using as example the script found at -`tests/generate_benchmark_db.py`. The evaluation flow is orchestrated with +You can perform your own evaluations with the Typer CLI in +`evaluations/benchmark.py`, for example `python -m evaluations.benchmark repliqa`. +The evaluation flow is orchestrated with [`pydantic-evals`](https://github.com/pydantic/pydantic-ai/tree/main/libs/pydantic-evals), which we leverage for dataset management, scoring, and report generation. diff --git a/evaluations/__init__.py b/evaluations/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/generate_benchmark_db.py b/evaluations/benchmark.py similarity index 57% rename from tests/generate_benchmark_db.py rename to evaluations/benchmark.py index f1209ac5..1e42640f 100644 --- a/tests/generate_benchmark_db.py +++ b/evaluations/benchmark.py @@ -1,59 +1,75 @@ import asyncio -from pathlib import Path +from collections.abc import Mapping +from typing import Any, cast import logfire -from datasets import Dataset, load_dataset -from llm_judge import ANSWER_EQUIVALENCE_RUBRIC +import typer from pydantic_ai.models.openai import OpenAIChatModel from pydantic_ai.providers.ollama import OllamaProvider -from pydantic_evals import Case from pydantic_evals import Dataset as EvalDataset from pydantic_evals.evaluators import IsInstance, LLMJudge from pydantic_evals.reporting import ReportCaseFailure from rich.console import Console from rich.progress import Progress -from haiku.rag import logging # noqa +from evaluations.config import DatasetSpec, RetrievalSample +from evaluations.datasets import DATASETS +from evaluations.llm_judge import ANSWER_EQUIVALENCE_RUBRIC +from haiku.rag import logging # noqa: F401 from haiku.rag.client import HaikuRAG from haiku.rag.config import Config from haiku.rag.logging import configure_cli_logging from haiku.rag.qa import get_qa_agent +QA_JUDGE_MODEL = "qwen3" + logfire.configure(send_to_logfire="if-token-present", service_name="evals") logfire.instrument_pydantic_ai() configure_cli_logging() console = Console() -QA_JUDGE_MODEL: str = "qwen3" -db_path = Path(__file__).parent / "data" / "benchmark.lancedb" - -async def populate_db(): - ds: Dataset = load_dataset("ServiceNow/repliqa")["repliqa_3"] # type: ignore - corpus = ds.filter(lambda doc: doc["document_topic"] == "News Stories") +async def populate_db(spec: DatasetSpec) -> None: + spec.db_path.parent.mkdir(parents=True, exist_ok=True) + corpus = spec.document_loader() + if spec.document_limit is not None: + corpus = corpus.select(range(min(spec.document_limit, len(corpus)))) with Progress() as progress: task = progress.add_task("[green]Populating database...", total=len(corpus)) - - async with HaikuRAG(db_path) as rag: + async with HaikuRAG(spec.db_path) as rag: for doc in corpus: - uri = doc["document_id"] # type: ignore - existing_doc = await rag.get_document_by_uri(uri) - if existing_doc is not None: + doc_mapping = cast(Mapping[str, Any], doc) + payload = spec.document_mapper(doc_mapping) + if payload is None: + progress.advance(task) + continue + + existing = await rag.get_document_by_uri(payload.uri) + if existing is not None: progress.advance(task) continue await rag.create_document( - content=doc["document_extracted"], # type: ignore - uri=uri, + content=payload.content, + uri=payload.uri, + title=payload.title, + metadata=payload.metadata, ) progress.advance(task) rag.store.vacuum() -async def run_match_benchmark(): - ds: Dataset = load_dataset("ServiceNow/repliqa")["repliqa_3"] # type: ignore - corpus = ds.filter(lambda doc: doc["document_topic"] == "News Stories") +def _is_relevant_match(retrieved_uri: str | None, sample: RetrievalSample) -> bool: + return retrieved_uri is not None and retrieved_uri in sample.expected_uris + + +async def run_retrieval_benchmark(spec: DatasetSpec) -> dict[str, float] | None: + if spec.retrieval_loader is None or spec.retrieval_mapper is None: + console.print("Skipping retrieval benchmark; no retrieval config.") + return None + + corpus = spec.retrieval_loader() correct_at_1 = 0 correct_at_2 = 0 @@ -64,42 +80,45 @@ async def run_match_benchmark(): task = progress.add_task( "[blue]Running retrieval benchmark...", total=len(corpus) ) - - async with HaikuRAG(db_path) as rag: + async with HaikuRAG(spec.db_path) as rag: for doc in corpus: - doc_id = doc["document_id"] # type: ignore - expected_answer = doc["answer"] # type: ignore - if expected_answer == "The answer is not found in the document.": + doc_mapping = cast(Mapping[str, Any], doc) + sample = spec.retrieval_mapper(doc_mapping) + if sample is None or sample.skip: + progress.advance(task) + continue + + matches = await rag.search(query=sample.question, limit=3) + if not matches: progress.advance(task) continue - matches = await rag.search( - query=doc["question"], # type: ignore - limit=3, - ) total_queries += 1 - # Check position of correct document in results for position, (chunk, _) in enumerate(matches): - assert chunk.document_id is not None, ( - "Chunk document_id should not be None" + retrieved = ( + await rag.get_document_by_id(chunk.document_id) + if chunk.document_id is not None + else None ) - retrieved = await rag.get_document_by_id(chunk.document_id) - if retrieved and retrieved.uri == doc_id: - if position == 0: # First position + if retrieved and _is_relevant_match(retrieved.uri, sample): + if position == 0: correct_at_1 += 1 correct_at_2 += 1 correct_at_3 += 1 - elif position == 1: # Second position + elif position == 1: correct_at_2 += 1 correct_at_3 += 1 - elif position == 2: # Third position + elif position == 2: correct_at_3 += 1 break progress.advance(task) - # Calculate recall metrics + if total_queries == 0: + console.print("No retrieval cases to evaluate.") + return None + recall_at_1 = correct_at_1 / total_queries recall_at_2 = correct_at_2 / total_queries recall_at_3 = correct_at_3 / total_queries @@ -110,35 +129,24 @@ async def run_match_benchmark(): console.print(f"Recall@2: {recall_at_2:.4f}") console.print(f"Recall@3: {recall_at_3:.4f}") - return {"recall@1": recall_at_1, "recall@2": recall_at_2, "recall@3": recall_at_3} + return { + "recall@1": recall_at_1, + "recall@2": recall_at_2, + "recall@3": recall_at_3, + } -async def run_qa_benchmark(k: int | None = None): - """Run QA benchmarking on the corpus.""" - ds: Dataset = load_dataset("ServiceNow/repliqa")["repliqa_3"] # type: ignore - corpus = ds.filter(lambda doc: doc["document_topic"] == "News Stories") +async def run_qa_benchmark( + spec: DatasetSpec, qa_limit: int | None = None +) -> ReportCaseFailure[str, str, dict[str, str]] | None: + corpus = spec.qa_loader() + if qa_limit is not None: + corpus = corpus.select(range(min(qa_limit, len(corpus)))) - if k is not None: - corpus = corpus.select(range(min(k, len(corpus)))) - - cases: list[Case[str, str, dict[str, str]]] = [] - for index, doc in enumerate(corpus, start=1): - question = doc["question"] # type: ignore[index] - expected_answer = doc["answer"] # type: ignore[index] - doc_id = doc["document_id"] # type: ignore[index] - case_name = f"{index}_{doc_id}" if doc_id is not None else f"case_{index}" - - cases.append( - Case( - name=case_name, - inputs=question, - expected_output=expected_answer, - metadata={ - "document_id": str(doc_id), - "case_index": str(index), - }, - ) - ) + cases = [ + spec.qa_case_builder(index, cast(Mapping[str, Any], doc)) + for index, doc in enumerate(corpus, start=1) + ] judge_model = OpenAIChatModel( model_name=QA_JUDGE_MODEL, @@ -172,7 +180,7 @@ async def run_qa_benchmark(k: int | None = None): total=len(evaluation_dataset.cases), ) - async with HaikuRAG(db_path) as rag: + async with HaikuRAG(spec.db_path) as rag: qa = get_qa_agent(rag) async def answer_question(question: str) -> str: @@ -227,6 +235,7 @@ async def run_qa_benchmark(k: int | None = None): f"{passing_cases}/{total_processed}[/green]" ) progress.advance(qa_task) + total_cases = total_processed accuracy = passing_cases / total_cases if total_cases > 0 else 0 @@ -243,16 +252,61 @@ async def run_qa_benchmark(k: int | None = None): console.print(f"Error: {failure.error_message}") console.print("") + return failures[0] if failures else None -async def main(): - await populate_db() - console.print("Running retrieval benchmarks...", style="bold blue") - await run_match_benchmark() +async def evaluate_dataset( + spec: DatasetSpec, + skip_retrieval: bool, + skip_qa: bool, + qa_limit: int | None, +) -> None: + console.print(f"Using dataset: {spec.key}", style="bold magenta") + await populate_db(spec) - console.print("\nRunning QA benchmarks...", style="bold yellow") - await run_qa_benchmark() + if not skip_retrieval: + console.print("Running retrieval benchmarks...", style="bold blue") + await run_retrieval_benchmark(spec) + else: + console.print("Skipping retrieval benchmark by request.") + + if not skip_qa: + console.print("\nRunning QA benchmarks...", style="bold yellow") + await run_qa_benchmark(spec, qa_limit=qa_limit) + else: + console.print("Skipping QA benchmark by request.") + + +app = typer.Typer(help="Run retrieval and QA benchmarks for configured datasets.") + + +@app.command() +def run( + dataset: str = typer.Argument(..., help="Dataset key to evaluate."), + skip_retrieval: bool = typer.Option( + False, "--skip-retrieval", help="Skip retrieval benchmark." + ), + skip_qa: bool = typer.Option(False, "--skip-qa", help="Skip QA benchmark."), + qa_limit: int | None = typer.Option( + None, "--qa-limit", help="Limit number of QA cases." + ), +) -> None: + spec = DATASETS.get(dataset.lower()) + if spec is None: + valid_datasets = ", ".join(sorted(DATASETS)) + raise typer.BadParameter( + f"Unknown dataset '{dataset}'. Choose from: {valid_datasets}" + ) + + asyncio.run( + evaluate_dataset( + spec=spec, + skip_retrieval=skip_retrieval, + skip_qa=skip_qa, + qa_limit=qa_limit, + ) + ) if __name__ == "__main__": - asyncio.run(main()) + app() diff --git a/evaluations/config.py b/evaluations/config.py new file mode 100644 index 00000000..8e230dda --- /dev/null +++ b/evaluations/config.py @@ -0,0 +1,46 @@ +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +from datasets import Dataset +from pydantic_evals import Case + + +@dataclass +class DocumentPayload: + uri: str + content: str + title: str | None = None + metadata: dict[str, Any] | None = None + + +@dataclass +class RetrievalSample: + question: str + expected_uris: tuple[str, ...] + skip: bool = False + + +DocumentLoader = Callable[[], Dataset] +DocumentMapper = Callable[[Mapping[str, Any]], DocumentPayload | None] +RetrievalLoader = Callable[[], Dataset] +RetrievalMapper = Callable[[Mapping[str, Any]], RetrievalSample | None] +CaseBuilder = Callable[[int, Mapping[str, Any]], Case[str, str, dict[str, str]]] + + +@dataclass +class DatasetSpec: + key: str + db_filename: str + document_loader: DocumentLoader + document_mapper: DocumentMapper + qa_loader: DocumentLoader + qa_case_builder: CaseBuilder + retrieval_loader: RetrievalLoader | None = None + retrieval_mapper: RetrievalMapper | None = None + document_limit: int | None = None + + @property + def db_path(self) -> Path: + return Path(__file__).parent / "data" / self.db_filename diff --git a/evaluations/datasets/__init__.py b/evaluations/datasets/__init__.py new file mode 100644 index 00000000..94a00c34 --- /dev/null +++ b/evaluations/datasets/__init__.py @@ -0,0 +1,8 @@ +from evaluations.config import DatasetSpec + +from .repliqa import REPLIQ_SPEC +from .wix import WIX_SPEC + +DATASETS: dict[str, DatasetSpec] = {spec.key: spec for spec in (REPLIQ_SPEC, WIX_SPEC)} + +__all__ = ["DATASETS"] diff --git a/evaluations/datasets/repliqa.py b/evaluations/datasets/repliqa.py new file mode 100644 index 00000000..2fc21d0f --- /dev/null +++ b/evaluations/datasets/repliqa.py @@ -0,0 +1,58 @@ +from collections.abc import Mapping +from typing import Any, cast + +from datasets import Dataset, DatasetDict, load_dataset +from pydantic_evals import Case + +from evaluations.config import DatasetSpec, DocumentPayload, RetrievalSample + + +def load_repliqa_corpus() -> Dataset: + dataset_dict = cast(DatasetDict, load_dataset("ServiceNow/repliqa")) + dataset = cast(Dataset, dataset_dict["repliqa_3"]) + return dataset.filter(lambda doc: doc["document_topic"] == "News Stories") + + +def map_repliqa_document(doc: Mapping[str, Any]) -> DocumentPayload: + return DocumentPayload( + uri=str(doc["document_id"]), + content=doc["document_extracted"], + ) + + +def map_repliqa_retrieval(doc: Mapping[str, Any]) -> RetrievalSample | None: + expected_answer = doc["answer"] + if expected_answer == "The answer is not found in the document.": + return None + return RetrievalSample( + question=doc["question"], + expected_uris=(str(doc["document_id"]),), + ) + + +def build_repliqa_case( + index: int, doc: Mapping[str, Any] +) -> Case[str, str, dict[str, str]]: + document_id = doc["document_id"] + case_name = f"{index}_{document_id}" if document_id is not None else f"case_{index}" + return Case( + name=case_name, + inputs=doc["question"], + expected_output=doc["answer"], + metadata={ + "document_id": str(document_id), + "case_index": str(index), + }, + ) + + +REPLIQ_SPEC = DatasetSpec( + key="repliqa", + db_filename="repliqa.lancedb", + document_loader=load_repliqa_corpus, + document_mapper=map_repliqa_document, + qa_loader=load_repliqa_corpus, + qa_case_builder=build_repliqa_case, + retrieval_loader=load_repliqa_corpus, + retrieval_mapper=map_repliqa_retrieval, +) diff --git a/evaluations/datasets/wix.py b/evaluations/datasets/wix.py new file mode 100644 index 00000000..112ba171 --- /dev/null +++ b/evaluations/datasets/wix.py @@ -0,0 +1,81 @@ +import json +from collections.abc import Iterable, Mapping +from typing import Any, cast + +from datasets import Dataset, DatasetDict, load_dataset +from pydantic_evals import Case + +from evaluations.config import DatasetSpec, DocumentPayload, RetrievalSample + + +def load_wix_corpus() -> Dataset: + dataset_dict = cast(DatasetDict, load_dataset("Wix/WixQA", "wix_kb_corpus")) + return cast(Dataset, dataset_dict["train"]) + + +def map_wix_document(doc: Mapping[str, Any]) -> DocumentPayload: + article_id = doc.get("id") + url = doc.get("url") + uri = str(article_id) if article_id is not None else str(url) + + metadata: dict[str, str] = {} + if article_id is not None: + metadata["article_id"] = str(article_id) + if url: + metadata["url"] = str(url) + + return DocumentPayload( + uri=uri, + content=doc["contents"], + title=doc.get("title"), + metadata=metadata or None, + ) + + +def load_wix_qa() -> Dataset: + dataset_dict = cast(DatasetDict, load_dataset("Wix/WixQA", "wixqa_expertwritten")) + return cast(Dataset, dataset_dict["train"]) + + +def map_wix_retrieval(doc: Mapping[str, Any]) -> RetrievalSample | None: + article_ids: Iterable[int | str] | None = doc.get("article_ids") + if not article_ids: + return None + + expected_uris = tuple(str(article_id) for article_id in article_ids) + return RetrievalSample( + question=doc["question"], + expected_uris=expected_uris, + ) + + +def build_wix_case( + index: int, doc: Mapping[str, Any] +) -> Case[str, str, dict[str, str]]: + article_ids = tuple(str(article_id) for article_id in doc.get("article_ids") or []) + joined_ids = "-".join(article_ids) + case_name = f"{index}_{joined_ids}" if joined_ids else f"case_{index}" + + metadata = { + "case_index": str(index), + "document_ids": json.dumps(article_ids), + } + + return Case( + name=case_name, + inputs=doc["question"], + expected_output=doc["answer"], + metadata=metadata, + ) + + +WIX_SPEC = DatasetSpec( + key="wix", + db_filename="wix.lancedb", + document_loader=load_wix_corpus, + document_mapper=map_wix_document, + qa_loader=load_wix_qa, + qa_case_builder=build_wix_case, + retrieval_loader=load_wix_qa, + retrieval_mapper=map_wix_retrieval, +) diff --git a/tests/llm_judge.py b/evaluations/llm_judge.py similarity index 100% rename from tests/llm_judge.py rename to evaluations/llm_judge.py diff --git a/tests/test_qa.py b/tests/test_qa.py index 2fa2d417..1f2f4022 100644 --- a/tests/test_qa.py +++ b/tests/test_qa.py @@ -1,12 +1,11 @@ import pytest from datasets import Dataset +from evaluations.llm_judge import LLMJudge from haiku.rag.client import HaikuRAG from haiku.rag.config import Config from haiku.rag.qa.agent import QuestionAnswerAgent -from .llm_judge import LLMJudge - OPENAI_AVAILABLE = bool(Config.OPENAI_API_KEY) ANTHROPIC_AVAILABLE = bool(Config.ANTHROPIC_API_KEY) VLLM_QA_AVAILABLE = bool(Config.VLLM_QA_BASE_URL)