Merge pull request #86 from ggozad/feat/pydantic-evaluations
Use pydantic evals for running our own evaluations
This commit is contained in:
commit
73209b203c
4 changed files with 130 additions and 36 deletions
|
|
@ -3,7 +3,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`.
|
||||
`tests/generate_benchmark_db.py`. 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.
|
||||
|
||||
## Recall
|
||||
|
||||
|
|
@ -24,7 +26,9 @@ The recall obtained is ~0.79 for matching in the top result, raising to ~0.91 fo
|
|||
|
||||
## Question/Answer evaluation
|
||||
|
||||
Again using the same dataset, we use a QA agent to answer the question. In addition we use an LLM judge (using the Ollama `qwen3`) to evaluate whether the answer is correct or not. The obtained accuracy is as follows:
|
||||
Again using the same dataset, we use a QA agent to answer the question.
|
||||
`pydantic-evals` runs each case and coordinates an LLM judge (Ollama `qwen3`) to
|
||||
determine whether the answer is correct. The obtained accuracy is as follows:
|
||||
|
||||
| Embedding Model | QA Model | Accuracy | Reranker |
|
||||
|------------------------------------|-----------------------------------|-----------|------------------------|
|
||||
|
|
|
|||
|
|
@ -57,6 +57,7 @@ dev = [
|
|||
"logfire>=4.7.0",
|
||||
"mkdocs>=1.6.1",
|
||||
"mkdocs-material>=9.6.14",
|
||||
"pydantic-evals>=1.0.8",
|
||||
"pre-commit>=4.2.0",
|
||||
"pyright>=1.1.405",
|
||||
"pytest>=8.4.2",
|
||||
|
|
|
|||
|
|
@ -3,20 +3,28 @@ from pathlib import Path
|
|||
|
||||
import logfire
|
||||
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.progress import Progress
|
||||
|
||||
from haiku.rag import logging # noqa
|
||||
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
|
||||
|
||||
logfire.configure(send_to_logfire="if-token-present")
|
||||
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"
|
||||
|
||||
|
||||
|
|
@ -113,49 +121,128 @@ async def run_qa_benchmark(k: int | None = None):
|
|||
if k is not None:
|
||||
corpus = corpus.select(range(min(k, len(corpus))))
|
||||
|
||||
judge = LLMJudge()
|
||||
correct_answers = 0
|
||||
total_questions = 0
|
||||
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}"
|
||||
|
||||
with Progress() as progress:
|
||||
task = progress.add_task("[yellow]Running QA benchmark...", total=len(corpus))
|
||||
cases.append(
|
||||
Case(
|
||||
name=case_name,
|
||||
inputs=question,
|
||||
expected_output=expected_answer,
|
||||
metadata={
|
||||
"document_id": str(doc_id),
|
||||
"case_index": str(index),
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
judge_model = OpenAIChatModel(
|
||||
model_name=QA_JUDGE_MODEL,
|
||||
provider=OllamaProvider(base_url=f"{Config.OLLAMA_BASE_URL}/v1"),
|
||||
)
|
||||
|
||||
evaluation_dataset = EvalDataset[str, str, dict[str, str]](
|
||||
cases=cases,
|
||||
evaluators=[
|
||||
IsInstance(type_name="str"),
|
||||
LLMJudge(
|
||||
rubric=ANSWER_EQUIVALENCE_RUBRIC,
|
||||
include_input=True,
|
||||
include_expected_output=True,
|
||||
model=judge_model,
|
||||
assertion={
|
||||
"evaluation_name": "answer_equivalent",
|
||||
"include_reason": True,
|
||||
},
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
total_processed = 0
|
||||
passing_cases = 0
|
||||
failures: list[ReportCaseFailure[str, str, dict[str, str]]] = []
|
||||
|
||||
with Progress(console=console) as progress:
|
||||
qa_task = progress.add_task(
|
||||
"[yellow]Evaluating QA cases...",
|
||||
total=len(evaluation_dataset.cases),
|
||||
)
|
||||
|
||||
async with HaikuRAG(db_path) as rag:
|
||||
qa = get_qa_agent(rag)
|
||||
for doc in corpus:
|
||||
question = doc["question"] # type: ignore
|
||||
expected_answer = doc["answer"] # type: ignore
|
||||
|
||||
# Really small models might fail, let's account for that in try/except
|
||||
try:
|
||||
generated_answer = await qa.answer(question)
|
||||
is_equivalent = await judge.judge_answers(
|
||||
question, generated_answer, expected_answer
|
||||
async def answer_question(question: str) -> str:
|
||||
return await qa.answer(question)
|
||||
|
||||
for case in evaluation_dataset.cases:
|
||||
progress.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")
|
||||
progress.console.print(f"Question: {result_case.inputs}")
|
||||
progress.console.print(f"Expected: {result_case.expected_output}")
|
||||
progress.console.print(f"Generated: {result_case.output}")
|
||||
if equivalence is not None:
|
||||
progress.console.print(
|
||||
f"Equivalent: {equivalence.value}"
|
||||
+ (f" — {equivalence.reason}" if equivalence.reason else "")
|
||||
)
|
||||
if equivalence.value:
|
||||
passing_cases += 1
|
||||
|
||||
progress.console.print("")
|
||||
|
||||
if report.failures:
|
||||
failures.extend(report.failures)
|
||||
failure = report.failures[0]
|
||||
progress.console.print(
|
||||
"[red]Failure encountered during case evaluation:[/red]"
|
||||
)
|
||||
console.print(f"Question: {question}")
|
||||
console.print(f"Expected: {expected_answer}")
|
||||
console.print(f"Generated: {generated_answer}")
|
||||
console.print(f"Equivalent: {is_equivalent}\n")
|
||||
progress.console.print(f"Question: {failure.inputs}")
|
||||
progress.console.print(f"Error: {failure.error_message}")
|
||||
progress.console.print("")
|
||||
|
||||
if is_equivalent:
|
||||
correct_answers += 1
|
||||
except Exception as e:
|
||||
console.print(f"[red]Error processing question: {question}[/red]")
|
||||
console.print(f"[red]{e}[/red]")
|
||||
finally:
|
||||
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
|
||||
progress.console.print(
|
||||
f"[green]Accuracy: {(passing_cases / total_processed):.4f} "
|
||||
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
|
||||
|
||||
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"Total questions: {total_cases}")
|
||||
console.print(f"Correct answers: {passing_cases}")
|
||||
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():
|
||||
await populate_db()
|
||||
|
|
|
|||
2
uv.lock
2
uv.lock
|
|
@ -1143,6 +1143,7 @@ dev = [
|
|||
{ name = "mkdocs" },
|
||||
{ name = "mkdocs-material" },
|
||||
{ name = "pre-commit" },
|
||||
{ name = "pydantic-evals" },
|
||||
{ name = "pyright" },
|
||||
{ name = "pytest" },
|
||||
{ name = "pytest-asyncio" },
|
||||
|
|
@ -1176,6 +1177,7 @@ dev = [
|
|||
{ name = "mkdocs", specifier = ">=1.6.1" },
|
||||
{ name = "mkdocs-material", specifier = ">=9.6.14" },
|
||||
{ name = "pre-commit", specifier = ">=4.2.0" },
|
||||
{ name = "pydantic-evals", specifier = ">=1.0.8" },
|
||||
{ name = "pyright", specifier = ">=1.1.405" },
|
||||
{ name = "pytest", specifier = ">=8.4.2" },
|
||||
{ name = "pytest-asyncio", specifier = ">=1.2.0" },
|
||||
|
|
|
|||
Loading…
Reference in a new issue