From 42e62f92a1b44237f6f4d0c5e715bcd92c2fa35d Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Wed, 9 Jul 2025 10:27:32 +0300 Subject: [PATCH] Better formatting of benchmark script --- docs/benchmarks.md | 2 +- tests/generate_benchmark_db.py | 143 +++++++++++++++++++-------------- 2 files changed, 83 insertions(+), 62 deletions(-) diff --git a/docs/benchmarks.md b/docs/benchmarks.md index e1fc2bf7..948cba93 100644 --- a/docs/benchmarks.md +++ b/docs/benchmarks.md @@ -15,7 +15,7 @@ The recall obtained is ~0.73 for matching in the top result, raising to ~0.75 fo | Model | Document in top 1 | Document in top 3 | |---------------------------------------|-------------------|-------------------| | Ollama / `mxbai-embed-large` | 0.73 | 0.75 | -| OpenAI / `text-embeddings-3-small` | | | +| OpenAI / `text-embeddings-3-small` | 0.75 | 0.88 | ## Question/Answer evaluation diff --git a/tests/generate_benchmark_db.py b/tests/generate_benchmark_db.py index 20bb01b2..70b9c468 100644 --- a/tests/generate_benchmark_db.py +++ b/tests/generate_benchmark_db.py @@ -3,28 +3,37 @@ from pathlib import Path from datasets import Dataset, load_dataset from llm_judge import LLMJudge -from tqdm import tqdm +from rich.console import Console +from rich.progress import Progress from haiku.rag.client import HaikuRAG from haiku.rag.qa import get_qa_agent +console = Console() + db_path = Path(__file__).parent / "data" / "benchmark.sqlite" async def populate_db(): - if (db_path).exists(): - print("Benchmark database already exists. Skipping creation.") - return - ds: Dataset = load_dataset("ServiceNow/repliqa")["repliqa_3"] # type: ignore corpus = ds.filter(lambda doc: doc["document_topic"] == "News Stories") - async with HaikuRAG(db_path) as rag: - for i, doc in enumerate(tqdm(corpus)): - await rag.create_document( - content=doc["document_extracted"], # type: ignore - uri=doc["document_id"], # type: ignore - ) + with Progress() as progress: + task = progress.add_task("[green]Populating database...", total=len(corpus)) + + async with HaikuRAG(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: + progress.advance(task) + continue + + await rag.create_document( + content=doc["document_extracted"], # type: ignore + uri=uri, + ) + progress.advance(task) async def run_match_benchmark(): @@ -36,41 +45,48 @@ async def run_match_benchmark(): correct_at_3 = 0 total_queries = 0 - async with HaikuRAG(db_path) as rag: - for i, doc in enumerate(tqdm(corpus)): - doc_id = doc["document_id"] # type: ignore - matches = await rag.search( - query=doc["question"], # type: ignore - limit=3, - ) + with Progress() as progress: + task = progress.add_task( + "[blue]Running retrieval benchmark...", total=len(corpus) + ) - total_queries += 1 + async with HaikuRAG(db_path) as rag: + for doc in corpus: + doc_id = doc["document_id"] # type: ignore + matches = await rag.search( + query=doc["question"], # type: ignore + limit=3, + ) - # Check position of correct document in results - for position, (chunk, _) in enumerate(matches): - retrieved = await rag.get_document_by_id(chunk.document_id) - if retrieved and retrieved.uri == doc_id: - if position == 0: # First position - correct_at_1 += 1 - correct_at_2 += 1 - correct_at_3 += 1 - elif position == 1: # Second position - correct_at_2 += 1 - correct_at_3 += 1 - elif position == 2: # Third position - correct_at_3 += 1 - break + total_queries += 1 + + # Check position of correct document in results + for position, (chunk, _) in enumerate(matches): + retrieved = await rag.get_document_by_id(chunk.document_id) + if retrieved and retrieved.uri == doc_id: + if position == 0: # First position + correct_at_1 += 1 + correct_at_2 += 1 + correct_at_3 += 1 + elif position == 1: # Second position + correct_at_2 += 1 + correct_at_3 += 1 + elif position == 2: # Third position + correct_at_3 += 1 + break + + progress.advance(task) # Calculate recall metrics recall_at_1 = correct_at_1 / total_queries recall_at_2 = correct_at_2 / total_queries recall_at_3 = correct_at_3 / total_queries - print("\n=== Retrieval Benchmark Results ===") - print(f"Total queries: {total_queries}") - print(f"Recall@1: {recall_at_1:.4f}") - print(f"Recall@2: {recall_at_2:.4f}") - print(f"Recall@3: {recall_at_3:.4f}") + console.print("\n=== Retrieval Benchmark Results ===", style="bold cyan") + console.print(f"Total queries: {total_queries}") + console.print(f"Recall@1: {recall_at_1:.4f}") + 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} @@ -87,42 +103,47 @@ async def run_qa_benchmark(k: int | None = None): correct_answers = 0 total_questions = 0 - async with HaikuRAG(db_path) as rag: - qa = get_qa_agent(rag) + with Progress() as progress: + task = progress.add_task("[yellow]Running QA benchmark...", total=len(corpus)) - for i, doc in enumerate(tqdm(corpus, desc="QA Benchmarking")): - question = doc["question"] # type: ignore - expected_answer = doc["answer"] # type: ignore + async with HaikuRAG(db_path) as rag: + qa = get_qa_agent(rag) - generated_answer = await qa.answer(question) - is_equivalent = await judge.judge_answers( - question, generated_answer, expected_answer - ) - print(f"Question: {question}") - print(f"Expected: {expected_answer}") - print(f"Generated: {generated_answer}") - print(f"Equivalent: {is_equivalent}\n") + for doc in corpus: + question = doc["question"] # type: ignore + expected_answer = doc["answer"] # type: ignore - if is_equivalent: - correct_answers += 1 - total_questions += 1 - print("Current score:", correct_answers, "/", total_questions) + generated_answer = await qa.answer(question) + is_equivalent = await judge.judge_answers( + question, generated_answer, expected_answer + ) + console.print(f"Question: {question}") + console.print(f"Expected: {expected_answer}") + console.print(f"Generated: {generated_answer}") + console.print(f"Equivalent: {is_equivalent}\n") + + if is_equivalent: + correct_answers += 1 + total_questions += 1 + console.print("Current score:", correct_answers, "/", total_questions) + + progress.advance(task) accuracy = correct_answers / total_questions if total_questions > 0 else 0 - print("\n=== QA Benchmark Results ===") - print(f"Total questions: {total_questions}") - print(f"Correct answers: {correct_answers}") - print(f"QA Accuracy: {accuracy:.4f} ({accuracy * 100:.2f}%)") + console.print("\n=== QA Benchmark Results ===", style="bold cyan") + console.print(f"Total questions: {total_questions}") + console.print(f"Correct answers: {correct_answers}") + console.print(f"QA Accuracy: {accuracy:.4f} ({accuracy * 100:.2f}%)") async def main(): await populate_db() - print("Running retrieval benchmarks...") + console.print("Running retrieval benchmarks...", style="bold blue") await run_match_benchmark() - print("\nRunning QA benchmarks...") + console.print("\nRunning QA benchmarks...", style="bold yellow") await run_qa_benchmark()