improve output
This commit is contained in:
parent
87f28ed96b
commit
19ad1b9d7d
1 changed files with 53 additions and 45 deletions
|
|
@ -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,63 +162,71 @@ 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]]] = []
|
||||||
|
|
||||||
async with HaikuRAG(db_path) as rag:
|
with Progress(console=console) as progress:
|
||||||
qa = get_qa_agent(rag)
|
qa_task = progress.add_task(
|
||||||
|
"[yellow]Evaluating QA cases...",
|
||||||
|
total=len(evaluation_dataset.cases),
|
||||||
|
)
|
||||||
|
|
||||||
async def answer_question(question: str) -> str:
|
async with HaikuRAG(db_path) as rag:
|
||||||
return await qa.answer(question)
|
qa = get_qa_agent(rag)
|
||||||
|
|
||||||
for case in evaluation_dataset.cases:
|
async def answer_question(question: str) -> str:
|
||||||
console.print(f"\n[bold]Evaluating case:[/bold] {case.name}")
|
return await qa.answer(question)
|
||||||
|
|
||||||
single_case_dataset = EvalDataset[str, str, dict[str, str]](
|
for case in evaluation_dataset.cases:
|
||||||
cases=[case],
|
progress.console.print(f"\n[bold]Evaluating case:[/bold] {case.name}")
|
||||||
evaluators=evaluation_dataset.evaluators,
|
|
||||||
)
|
|
||||||
|
|
||||||
report = await single_case_dataset.evaluate(
|
single_case_dataset = EvalDataset[str, str, dict[str, str]](
|
||||||
answer_question,
|
cases=[case],
|
||||||
name="qa_answer",
|
evaluators=evaluation_dataset.evaluators,
|
||||||
max_concurrency=1,
|
)
|
||||||
progress=False,
|
|
||||||
)
|
|
||||||
|
|
||||||
total_processed += 1
|
report = await single_case_dataset.evaluate(
|
||||||
|
answer_question,
|
||||||
|
name="qa_answer",
|
||||||
|
max_concurrency=1,
|
||||||
|
progress=False,
|
||||||
|
)
|
||||||
|
|
||||||
if report.cases:
|
total_processed += 1
|
||||||
result_case = report.cases[0]
|
|
||||||
|
|
||||||
equivalence = result_case.assertions.get("answer_equivalent")
|
if report.cases:
|
||||||
console.print(f"Question: {result_case.inputs}")
|
result_case = report.cases[0]
|
||||||
console.print(f"Expected: {result_case.expected_output}")
|
|
||||||
console.print(f"Generated: {result_case.output}")
|
equivalence = result_case.assertions.get("answer_equivalent")
|
||||||
if equivalence is not None:
|
progress.console.print(f"Question: {result_case.inputs}")
|
||||||
console.print(
|
progress.console.print(f"Expected: {result_case.expected_output}")
|
||||||
f"Equivalent: {equivalence.value}"
|
progress.console.print(f"Generated: {result_case.output}")
|
||||||
+ (f" — {equivalence.reason}" if equivalence.reason else "")
|
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]"
|
||||||
)
|
)
|
||||||
if equivalence.value:
|
progress.console.print(f"Question: {failure.inputs}")
|
||||||
passing_cases += 1
|
progress.console.print(f"Error: {failure.error_message}")
|
||||||
|
progress.console.print("")
|
||||||
|
|
||||||
console.print("")
|
progress.console.print(
|
||||||
|
f"[green]Accuracy: {(passing_cases / total_processed):.4f} "
|
||||||
if report.failures:
|
f"{passing_cases}/{total_processed}[/green]"
|
||||||
failures.extend(report.failures)
|
)
|
||||||
failure = report.failures[0]
|
progress.advance(qa_task)
|
||||||
console.print("[red]Failure encountered during case evaluation:[/red]")
|
|
||||||
console.print(f"Question: {failure.inputs}")
|
|
||||||
console.print(f"Error: {failure.error_message}")
|
|
||||||
console.print("")
|
|
||||||
|
|
||||||
console.print(
|
|
||||||
f"[green]Accuracy: {(passing_cases / total_processed):.4f}[/green]"
|
|
||||||
)
|
|
||||||
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
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue