Make benchmarks use pydantic evals

This commit is contained in:
Yiorgis Gozadinos 2025-09-26 09:46:41 +03:00
parent 53557b0629
commit 7cc4e13b57
No known key found for this signature in database
3 changed files with 111 additions and 33 deletions

View file

@ -57,6 +57,7 @@ dev = [
"logfire>=4.7.0", "logfire>=4.7.0",
"mkdocs>=1.6.1", "mkdocs>=1.6.1",
"mkdocs-material>=9.6.14", "mkdocs-material>=9.6.14",
"pydantic-evals>=1.0.8",
"pre-commit>=4.2.0", "pre-commit>=4.2.0",
"pyright>=1.1.405", "pyright>=1.1.405",
"pytest>=8.4.2", "pytest>=8.4.2",

View file

@ -3,12 +3,19 @@ from pathlib import Path
import logfire import logfire
from datasets import Dataset, load_dataset from datasets import Dataset, load_dataset
from llm_judge import LLMJudge from llm_judge import ANSWER_EQUIVALENCE_RUBRIC
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.console import Console
from rich.progress import Progress from rich.progress import Progress
from haiku.rag import logging # noqa from haiku.rag import logging # noqa
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.config import Config
from haiku.rag.logging import configure_cli_logging from haiku.rag.logging import configure_cli_logging
from haiku.rag.qa import get_qa_agent from haiku.rag.qa import get_qa_agent
@ -17,6 +24,7 @@ logfire.instrument_pydantic_ai()
configure_cli_logging() configure_cli_logging()
console = Console() console = Console()
QA_JUDGE_MODEL: str = "qwen3"
db_path = Path(__file__).parent / "data" / "benchmark.lancedb" db_path = Path(__file__).parent / "data" / "benchmark.lancedb"
@ -113,49 +121,116 @@ async def run_qa_benchmark(k: int | None = None):
if k is not None: if k is not None:
corpus = corpus.select(range(min(k, len(corpus)))) corpus = corpus.select(range(min(k, len(corpus))))
judge = LLMJudge() cases: list[Case[str, str, dict[str, str]]] = []
correct_answers = 0 for index, doc in enumerate(corpus, start=1):
total_questions = 0 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}"
with Progress() as progress: cases.append(
task = progress.add_task("[yellow]Running QA benchmark...", total=len(corpus)) Case(
name=case_name,
inputs=question,
expected_output=expected_answer,
metadata={
"document_id": str(doc_id),
"case_index": str(index),
},
)
)
async with HaikuRAG(db_path) as rag: judge_model = OpenAIChatModel(
qa = get_qa_agent(rag) model_name=QA_JUDGE_MODEL,
for doc in corpus: provider=OllamaProvider(base_url=f"{Config.OLLAMA_BASE_URL}/v1"),
question = doc["question"] # type: ignore )
expected_answer = doc["answer"] # type: ignore
# Really small models might fail, let's account for that in try/except evaluation_dataset = EvalDataset[str, str, dict[str, str]](
try: cases=cases,
generated_answer = await qa.answer(question) evaluators=[
is_equivalent = await judge.judge_answers( IsInstance(type_name="str"),
question, generated_answer, expected_answer LLMJudge(
) rubric=ANSWER_EQUIVALENCE_RUBRIC,
console.print(f"Question: {question}") include_input=True,
console.print(f"Expected: {expected_answer}") include_expected_output=True,
console.print(f"Generated: {generated_answer}") model=judge_model,
console.print(f"Equivalent: {is_equivalent}\n") assertion={
"evaluation_name": "answer_equivalent",
"include_reason": True,
},
),
],
)
if is_equivalent: console.print("[yellow]Running QA benchmark...[/yellow]")
correct_answers += 1
except Exception as e: total_processed = 0
console.print(f"[red]Error processing question: {question}[/red]") passing_cases = 0
console.print(f"[red]{e}[/red]") failures: list[ReportCaseFailure[str, str, dict[str, str]]] = []
finally:
total_questions += 1 async with HaikuRAG(db_path) as rag:
qa = get_qa_agent(rag)
async def answer_question(question: str) -> str:
return await qa.answer(question)
for case in evaluation_dataset.cases:
console.print(f"\n[bold]Evaluating case:[/bold] {case.name}")
single_case_dataset = EvalDataset[str, str, dict[str, str]](
cases=[case],
evaluators=evaluation_dataset.evaluators,
)
report = await single_case_dataset.evaluate(
answer_question,
name="qa_answer",
max_concurrency=1,
progress=False,
)
total_processed += 1
if report.cases:
result_case = report.cases[0]
equivalence = result_case.assertions.get("answer_equivalent")
console.print(f"Question: {result_case.inputs}")
console.print(f"Expected: {result_case.expected_output}")
console.print(f"Generated: {result_case.output}")
if equivalence is not None:
console.print( console.print(
"Current score:", correct_answers, "/", total_questions f"Equivalent: {equivalence.value}"
+ (f"{equivalence.reason}" if equivalence.reason else "")
) )
progress.advance(task) if equivalence.value:
passing_cases += 1
console.print("")
accuracy = correct_answers / total_questions if total_questions > 0 else 0 if report.failures:
failures.extend(report.failures)
failure = report.failures[0]
console.print("[red]Failure encountered during case evaluation:[/red]")
console.print(f"Question: {failure.inputs}")
console.print(f"Error: {failure.error_message}")
console.print("")
total_cases = total_processed
accuracy = passing_cases / total_cases if total_cases > 0 else 0
console.print("\n=== QA Benchmark Results ===", style="bold cyan") console.print("\n=== QA Benchmark Results ===", style="bold cyan")
console.print(f"Total questions: {total_questions}") console.print(f"Total questions: {total_cases}")
console.print(f"Correct answers: {correct_answers}") console.print(f"Correct answers: {passing_cases}")
console.print(f"QA Accuracy: {accuracy:.4f} ({accuracy * 100:.2f}%)") console.print(f"QA Accuracy: {accuracy:.4f} ({accuracy * 100:.2f}%)")
if failures:
console.print("[red]\nSummary of failures:[/red]")
for failure in failures:
console.print(f"Case: {failure.name}")
console.print(f"Question: {failure.inputs}")
console.print(f"Error: {failure.error_message}")
console.print("")
async def main(): async def main():
await populate_db() await populate_db()

View file

@ -1143,6 +1143,7 @@ dev = [
{ name = "mkdocs" }, { name = "mkdocs" },
{ name = "mkdocs-material" }, { name = "mkdocs-material" },
{ name = "pre-commit" }, { name = "pre-commit" },
{ name = "pydantic-evals" },
{ name = "pyright" }, { name = "pyright" },
{ name = "pytest" }, { name = "pytest" },
{ name = "pytest-asyncio" }, { name = "pytest-asyncio" },
@ -1176,6 +1177,7 @@ dev = [
{ name = "mkdocs", specifier = ">=1.6.1" }, { name = "mkdocs", specifier = ">=1.6.1" },
{ name = "mkdocs-material", specifier = ">=9.6.14" }, { name = "mkdocs-material", specifier = ">=9.6.14" },
{ name = "pre-commit", specifier = ">=4.2.0" }, { name = "pre-commit", specifier = ">=4.2.0" },
{ name = "pydantic-evals", specifier = ">=1.0.8" },
{ name = "pyright", specifier = ">=1.1.405" }, { name = "pyright", specifier = ">=1.1.405" },
{ name = "pytest", specifier = ">=8.4.2" }, { name = "pytest", specifier = ">=8.4.2" },
{ name = "pytest-asyncio", specifier = ">=1.2.0" }, { name = "pytest-asyncio", specifier = ">=1.2.0" },