From 9a859c6ee5862cd4ac525ffff1c673e431ae1709 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Mon, 29 Sep 2025 21:13:57 +0300 Subject: [PATCH] Add option to skip db in evals --- evaluations/benchmark.py | 14 ++++++++------ src/haiku/rag/reranking/__init__.py | 3 +++ 2 files changed, 11 insertions(+), 6 deletions(-) diff --git a/evaluations/benchmark.py b/evaluations/benchmark.py index e548326f..f309bd7e 100644 --- a/evaluations/benchmark.py +++ b/evaluations/benchmark.py @@ -263,24 +263,22 @@ async def run_qa_benchmark( async def evaluate_dataset( spec: DatasetSpec, + skip_db: bool, 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) + if not skip_db: + console.print(f"Using dataset: {spec.key}", style="bold magenta") + await populate_db(spec) 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.") @@ -289,6 +287,9 @@ 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_db: bool = typer.Option( + False, "--skip-db", help="Skip updateing the evaluation db." + ), skip_retrieval: bool = typer.Option( False, "--skip-retrieval", help="Skip retrieval benchmark." ), @@ -307,6 +308,7 @@ def run( asyncio.run( evaluate_dataset( spec=spec, + skip_db=skip_db, skip_retrieval=skip_retrieval, skip_qa=skip_qa, qa_limit=qa_limit, diff --git a/src/haiku/rag/reranking/__init__.py b/src/haiku/rag/reranking/__init__.py index f63453c6..f4753d50 100644 --- a/src/haiku/rag/reranking/__init__.py +++ b/src/haiku/rag/reranking/__init__.py @@ -1,3 +1,5 @@ +import os + from haiku.rag.config import Config from haiku.rag.reranking.base import RerankerBase @@ -17,6 +19,7 @@ def get_reranker() -> RerankerBase | None: try: from haiku.rag.reranking.mxbai import MxBAIReranker + os.environ["TOKENIZERS_PARALLELISM"] = "true" _reranker = MxBAIReranker() return _reranker except ImportError: