improve output

This commit is contained in:
Yiorgis Gozadinos 2025-09-26 10:46:15 +03:00
parent 87f28ed96b
commit 19ad1b9d7d
No known key found for this signature in database

View file

@ -19,7 +19,7 @@ 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
logfire.configure(send_to_logfire="if-token-present") logfire.configure(send_to_logfire="if-token-present", service_name="evals")
logfire.instrument_pydantic_ai() logfire.instrument_pydantic_ai()
configure_cli_logging() configure_cli_logging()
console = Console() console = Console()
@ -162,12 +162,16 @@ async def run_qa_benchmark(k: int | None = None):
], ],
) )
console.print("[yellow]Running QA benchmark...[/yellow]")
total_processed = 0 total_processed = 0
passing_cases = 0 passing_cases = 0
failures: list[ReportCaseFailure[str, str, dict[str, str]]] = [] 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: async with HaikuRAG(db_path) as rag:
qa = get_qa_agent(rag) qa = get_qa_agent(rag)
@ -175,7 +179,7 @@ async def run_qa_benchmark(k: int | None = None):
return await qa.answer(question) return await qa.answer(question)
for case in evaluation_dataset.cases: for case in evaluation_dataset.cases:
console.print(f"\n[bold]Evaluating case:[/bold] {case.name}") progress.console.print(f"\n[bold]Evaluating case:[/bold] {case.name}")
single_case_dataset = EvalDataset[str, str, dict[str, str]]( single_case_dataset = EvalDataset[str, str, dict[str, str]](
cases=[case], cases=[case],
@ -195,30 +199,34 @@ async def run_qa_benchmark(k: int | None = None):
result_case = report.cases[0] result_case = report.cases[0]
equivalence = result_case.assertions.get("answer_equivalent") equivalence = result_case.assertions.get("answer_equivalent")
console.print(f"Question: {result_case.inputs}") progress.console.print(f"Question: {result_case.inputs}")
console.print(f"Expected: {result_case.expected_output}") progress.console.print(f"Expected: {result_case.expected_output}")
console.print(f"Generated: {result_case.output}") progress.console.print(f"Generated: {result_case.output}")
if equivalence is not None: if equivalence is not None:
console.print( progress.console.print(
f"Equivalent: {equivalence.value}" f"Equivalent: {equivalence.value}"
+ (f"{equivalence.reason}" if equivalence.reason else "") + (f"{equivalence.reason}" if equivalence.reason else "")
) )
if equivalence.value: if equivalence.value:
passing_cases += 1 passing_cases += 1
console.print("") progress.console.print("")
if report.failures: if report.failures:
failures.extend(report.failures) failures.extend(report.failures)
failure = report.failures[0] failure = report.failures[0]
console.print("[red]Failure encountered during case evaluation:[/red]") progress.console.print(
console.print(f"Question: {failure.inputs}") "[red]Failure encountered during case evaluation:[/red]"
console.print(f"Error: {failure.error_message}")
console.print("")
console.print(
f"[green]Accuracy: {(passing_cases / total_processed):.4f}[/green]"
) )
progress.console.print(f"Question: {failure.inputs}")
progress.console.print(f"Error: {failure.error_message}")
progress.console.print("")
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 total_cases = total_processed
accuracy = passing_cases / total_cases if total_cases > 0 else 0 accuracy = passing_cases / total_cases if total_cases > 0 else 0