citation retrieval scoring

This commit is contained in:
Yiorgis Gozadinos 2026-04-28 12:46:00 +03:00
parent d00befd0c4
commit 7d288c525e
No known key found for this signature in database
5 changed files with 351 additions and 37 deletions

View file

@ -1,6 +1,6 @@
import asyncio import asyncio
import shutil import shutil
from collections.abc import Mapping from collections.abc import Awaitable, Callable, Mapping
from pathlib import Path from pathlib import Path
from typing import Any, Literal, cast from typing import Any, Literal, cast
@ -8,15 +8,21 @@ import logfire
import typer import typer
from dotenv import find_dotenv, load_dotenv from dotenv import find_dotenv, load_dotenv
from huggingface_hub import HfApi, snapshot_download from huggingface_hub import HfApi, snapshot_download
from pydantic_evals import Case, Dataset as EvalDataset from pydantic_evals import Case, Dataset as EvalDataset, set_eval_attribute
from pydantic_evals.evaluators import LLMJudge from pydantic_evals.evaluators import Evaluator, LLMJudge
from pydantic_evals.reporting import ReportCaseFailure 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 evaluations.config import DatasetSpec from evaluations.config import DatasetSpec
from evaluations.datasets import DATASETS from evaluations.datasets import DATASETS
from evaluations.evaluators import ANSWER_EQUIVALENCE_RUBRIC from evaluations.evaluators import (
ANSWER_EQUIVALENCE_RUBRIC,
CitationMAPEvaluator,
CitationMRREvaluator,
MAPEvaluator,
MRREvaluator,
)
from evaluations.skill_runner import SkillFactory, run_skill_question from evaluations.skill_runner import SkillFactory, run_skill_question
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.config import AppConfig, find_config_file, load_yaml_config from haiku.rag.config import AppConfig, find_config_file, load_yaml_config
@ -25,6 +31,11 @@ from haiku.rag.logging import configure_cli_logging
from haiku.rag.agents.qa import get_qa_agent from haiku.rag.agents.qa import get_qa_agent
from haiku.rag.utils import get_model, parse_model_option from haiku.rag.utils import get_model, parse_model_option
_CITATION_EVALUATORS: dict[type[Evaluator], type[Evaluator]] = {
MRREvaluator: CitationMRREvaluator,
MAPEvaluator: CitationMAPEvaluator,
}
Target = Literal["qa", "rag-skill", "analysis-skill"] Target = Literal["qa", "rag-skill", "analysis-skill"]
TARGETS: tuple[Target, ...] = ("qa", "rag-skill", "analysis-skill") TARGETS: tuple[Target, ...] = ("qa", "rag-skill", "analysis-skill")
@ -289,6 +300,44 @@ def _skill_factory_for_target(target: Target) -> SkillFactory:
raise ValueError(f"target {target!r} is not a skill target") raise ValueError(f"target {target!r} is not a skill target")
def _citation_evaluator_for(retrieval_evaluator: Evaluator | None) -> Evaluator | None:
"""Return the citation-scoring twin of the dataset's retrieval evaluator."""
if retrieval_evaluator is None:
return None
twin = _CITATION_EVALUATORS.get(type(retrieval_evaluator))
return twin() if twin is not None else None
def _attach_relevant_uris(
cases: list[Case[str, str, dict[str, Any]]],
spec: DatasetSpec,
limit: int | None,
) -> None:
"""Augment QA cases with `relevant_uris` joined from retrieval samples.
Mutates each case's metadata in place. Cases with no matching retrieval
sample (by question) are left untouched.
"""
if spec.retrieval_loader is None or spec.retrieval_mapper is None:
return
corpus = spec.retrieval_loader()
if limit is not None:
corpus = corpus.select(range(min(limit, len(corpus))))
expected_by_question: dict[str, tuple[str, ...]] = {}
for raw in corpus:
sample = spec.retrieval_mapper(cast(Mapping[str, Any], raw))
if sample is None or sample.skip:
continue
expected_by_question[sample.question] = sample.expected_uris
for case in cases:
uris = expected_by_question.get(case.inputs)
if uris is None:
continue
metadata = case.metadata if case.metadata is not None else {}
metadata["relevant_uris"] = list(uris)
case.metadata = metadata
async def run_qa_benchmark( async def run_qa_benchmark(
spec: DatasetSpec, spec: DatasetSpec,
config: AppConfig, config: AppConfig,
@ -309,28 +358,33 @@ async def run_qa_benchmark(
] ]
judge_config = judge_model or config.qa.model judge_config = judge_model or config.qa.model
judge = get_model(judge_config, config) skill_config = (skill_model or config.qa.model) if target != "qa" else None
db = spec.db_path(db_path)
citation_evaluator: Evaluator | None = None
if target != "qa":
_attach_relevant_uris(cases, spec, limit)
citation_evaluator = _citation_evaluator_for(spec.retrieval_evaluator)
evaluators: list[Evaluator] = [
LLMJudge(
rubric=ANSWER_EQUIVALENCE_RUBRIC,
include_input=True,
include_expected_output=True,
model=get_model(judge_config, config),
assertion={
"evaluation_name": "answer_equivalent",
"include_reason": True,
},
),
]
if citation_evaluator is not None:
evaluators.append(citation_evaluator)
evaluation_dataset = EvalDataset[str, str, dict[str, str]]( evaluation_dataset = EvalDataset[str, str, dict[str, str]](
name=spec.key, name=spec.key, cases=cases, evaluators=evaluators
cases=cases,
evaluators=[
LLMJudge(
rubric=ANSWER_EQUIVALENCE_RUBRIC,
include_input=True,
include_expected_output=True,
model=judge,
assertion={
"evaluation_name": "answer_equivalent",
"include_reason": True,
},
),
],
) )
db = spec.db_path(db_path)
skill_config = skill_model or config.qa.model if target != "qa" else None
eval_name = name if name is not None else f"{spec.key}_qa_evaluation" eval_name = name if name is not None else f"{spec.key}_qa_evaluation"
experiment_metadata = build_experiment_metadata( experiment_metadata = build_experiment_metadata(
dataset_key=spec.key, dataset_key=spec.key,
@ -341,6 +395,15 @@ async def run_qa_benchmark(
skill_config=skill_config, skill_config=skill_config,
) )
async def _evaluate(answer_fn: Callable[[str], Awaitable[str]]):
return await evaluation_dataset.evaluate(
answer_fn,
name=eval_name,
max_concurrency=1,
progress=True,
metadata=experiment_metadata,
)
if target == "qa": if target == "qa":
async with HaikuRAG(db, config=config) as rag: async with HaikuRAG(db, config=config) as rag:
qa = get_qa_agent( qa = get_qa_agent(
@ -351,13 +414,7 @@ async def run_qa_benchmark(
answer, _ = await qa.answer(question) answer, _ = await qa.answer(question)
return answer return answer
report = await evaluation_dataset.evaluate( report = await _evaluate(answer_question)
answer_question,
name=eval_name,
max_concurrency=1,
progress=True,
metadata=experiment_metadata,
)
else: else:
skill_factory = _skill_factory_for_target(target) skill_factory = _skill_factory_for_target(target)
assert skill_config is not None assert skill_config is not None
@ -371,15 +428,10 @@ async def run_qa_benchmark(
question=question, question=question,
skill_model=resolved_skill_model, skill_model=resolved_skill_model,
) )
set_eval_attribute("cited_uris", result.cited_uris)
return result.answer return result.answer
report = await evaluation_dataset.evaluate( report = await _evaluate(answer_question)
answer_question,
name=eval_name,
max_concurrency=1,
progress=True,
metadata=experiment_metadata,
)
passing_cases = sum( passing_cases = sum(
1 1
@ -389,7 +441,6 @@ async def run_qa_benchmark(
) )
total_processed = len(report.cases) total_processed = len(report.cases)
failures = report.failures failures = report.failures
accuracy = passing_cases / total_processed if total_processed > 0 else 0 accuracy = passing_cases / total_processed if total_processed > 0 else 0
console.print("\n=== QA Benchmark Results ===", style="bold cyan") console.print("\n=== QA Benchmark Results ===", style="bold cyan")
@ -397,6 +448,30 @@ async def run_qa_benchmark(
console.print(f"Correct answers: {passing_cases}") 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 citation_evaluator is not None:
score_key = citation_evaluator.get_default_evaluation_name()
scores = [
case.scores[score_key].value
for case in report.cases
if score_key in case.scores
]
if scores:
cited_count = sum(
1 for case in report.cases if case.attributes.get("cited_uris")
)
mean_citations = sum(
len(case.attributes.get("cited_uris") or []) for case in report.cases
) / len(report.cases)
mean_score = sum(scores) / len(scores)
console.print(
f"\n=== Citation Retrieval ({score_key}) ===", style="bold cyan"
)
console.print(f"Mean {score_key}: {mean_score:.4f}")
console.print(
f"Cite rate (≥1 citation): {cited_count / len(report.cases):.2%}"
)
console.print(f"Mean citations per case: {mean_citations:.2f}")
if failures: if failures:
console.print("[red]\nSummary of failures:[/red]") console.print("[red]\nSummary of failures:[/red]")
for failure in failures: for failure in failures:

View file

@ -1,3 +1,7 @@
from evaluations.evaluators.citation import (
CitationMAPEvaluator,
CitationMRREvaluator,
)
from evaluations.evaluators.judge import ( from evaluations.evaluators.judge import (
ANSWER_EQUIVALENCE_RUBRIC, ANSWER_EQUIVALENCE_RUBRIC,
LLMJudge, LLMJudge,
@ -8,6 +12,8 @@ from evaluations.evaluators.mrr import MRREvaluator
__all__ = [ __all__ = [
"ANSWER_EQUIVALENCE_RUBRIC", "ANSWER_EQUIVALENCE_RUBRIC",
"CitationMAPEvaluator",
"CitationMRREvaluator",
"LLMJudge", "LLMJudge",
"LLMJudgeResponseSchema", "LLMJudgeResponseSchema",
"MAPEvaluator", "MAPEvaluator",

View file

@ -0,0 +1,60 @@
from dataclasses import dataclass
from pydantic_evals.evaluators import Evaluator, EvaluatorContext
def _cited_uris(ctx: EvaluatorContext) -> list[str]:
return list(ctx.attributes.get("cited_uris") or [])
def _relevant_uris(ctx: EvaluatorContext) -> set[str]:
if ctx.metadata is None:
return set()
return set(ctx.metadata.get("relevant_uris", []))
@dataclass
class CitationMRREvaluator(Evaluator):
"""Reciprocal rank over the URIs the skill cited via the `cite` tool.
Reads ``cited_uris`` from ``ctx.attributes`` (recorded during the task run
via :func:`pydantic_evals.set_eval_attribute`) and ``relevant_uris`` from
``ctx.metadata``. Returns ``1.0/rank`` of the first cited URI that is in
the relevant set, or ``0.0`` if none match.
Use for single-document datasets, mirroring :class:`MRREvaluator`.
"""
evaluation_name: str = "cited_mrr"
def evaluate(self, ctx: EvaluatorContext) -> float:
relevant = _relevant_uris(ctx)
for rank, uri in enumerate(_cited_uris(ctx), start=1):
if uri in relevant:
return 1.0 / rank
return 0.0
@dataclass
class CitationMAPEvaluator(Evaluator):
"""Average precision over the URIs the skill cited via the `cite` tool.
Same input shape as :class:`CitationMRREvaluator`; use for multi-document
datasets, mirroring :class:`MAPEvaluator`.
"""
evaluation_name: str = "cited_map"
def evaluate(self, ctx: EvaluatorContext) -> float:
relevant = _relevant_uris(ctx)
if not relevant:
return 0.0
precisions: list[float] = []
found = 0
for rank, uri in enumerate(_cited_uris(ctx), start=1):
if uri in relevant:
found += 1
precisions.append(found / rank)
if not precisions:
return 0.0
return sum(precisions) / len(relevant)

View file

@ -322,3 +322,94 @@ class TestRunQaBenchmarkSkillTarget:
assert _skill_factory_for_target("analysis-skill") is analysis_factory assert _skill_factory_for_target("analysis-skill") is analysis_factory
with pytest.raises(ValueError, match="not a skill target"): with pytest.raises(ValueError, match="not a skill target"):
_skill_factory_for_target("qa") # type: ignore[arg-type] _skill_factory_for_target("qa") # type: ignore[arg-type]
class TestCitationEvaluatorWiring:
def test_returns_mrr_twin_for_mrr_evaluator(self) -> None:
from evaluations.benchmark import _citation_evaluator_for
from evaluations.evaluators import CitationMRREvaluator, MRREvaluator
result = _citation_evaluator_for(MRREvaluator())
assert isinstance(result, CitationMRREvaluator)
def test_returns_map_twin_for_map_evaluator(self) -> None:
from evaluations.benchmark import _citation_evaluator_for
from evaluations.evaluators import CitationMAPEvaluator, MAPEvaluator
result = _citation_evaluator_for(MAPEvaluator())
assert isinstance(result, CitationMAPEvaluator)
def test_returns_none_for_no_evaluator(self) -> None:
from evaluations.benchmark import _citation_evaluator_for
assert _citation_evaluator_for(None) is None
class TestAttachRelevantUris:
def test_joins_by_question(self) -> None:
from pydantic_evals import Case
from evaluations.benchmark import _attach_relevant_uris
from evaluations.config import RetrievalSample
from evaluations.evaluators import MRREvaluator
cases: list[Case[str, str, dict]] = [
Case(name="c1", inputs="What is X?", expected_output="X is a thing"),
Case(
name="c2",
inputs="What is Y?",
expected_output="Y is another",
metadata={"existing": "value"},
),
Case(
name="c3",
inputs="What is Z?",
expected_output="not in retrieval set",
),
]
spec = DatasetSpec(
key="test",
db_filename="test.lancedb",
document_loader=lambda: None, # type: ignore[arg-type] # ty: ignore[invalid-argument-type]
document_mapper=lambda doc: None,
qa_loader=lambda: [], # type: ignore[arg-type] # ty: ignore[invalid-argument-type]
qa_case_builder=lambda idx, doc: None, # type: ignore[arg-type] # ty: ignore[invalid-argument-type]
retrieval_loader=lambda: [ # type: ignore[arg-type] # ty: ignore[invalid-argument-type]
{"q": "What is X?", "uris": ("uri-x",)},
{"q": "What is Y?", "uris": ("uri-y1", "uri-y2")},
],
retrieval_mapper=lambda d: RetrievalSample(
question=d["q"], expected_uris=d["uris"]
),
retrieval_evaluator=MRREvaluator(),
)
_attach_relevant_uris(cases, spec, limit=None)
assert cases[0].metadata == {"relevant_uris": ["uri-x"]}
assert cases[1].metadata == {
"existing": "value",
"relevant_uris": ["uri-y1", "uri-y2"],
}
# case c3 has no matching retrieval sample — metadata untouched
assert cases[2].metadata is None
def test_no_op_without_retrieval_loader(self) -> None:
from pydantic_evals import Case
from evaluations.benchmark import _attach_relevant_uris
cases: list[Case[str, str, dict]] = [
Case(name="c1", inputs="q", expected_output="a"),
]
spec = DatasetSpec(
key="test",
db_filename="test.lancedb",
document_loader=lambda: None, # type: ignore[arg-type] # ty: ignore[invalid-argument-type]
document_mapper=lambda doc: None,
qa_loader=lambda: [], # type: ignore[arg-type] # ty: ignore[invalid-argument-type]
qa_case_builder=lambda idx, doc: None, # type: ignore[arg-type] # ty: ignore[invalid-argument-type]
)
_attach_relevant_uris(cases, spec, limit=None)
assert cases[0].metadata is None

View file

@ -0,0 +1,82 @@
from unittest.mock import MagicMock
from evaluations.evaluators.citation import (
CitationMAPEvaluator,
CitationMRREvaluator,
)
def _ctx(cited: list[str], relevant: list[str]) -> MagicMock:
ctx = MagicMock()
ctx.metadata = {"relevant_uris": relevant}
ctx.attributes = {"cited_uris": cited}
return ctx
class TestCitationMRREvaluator:
def setup_method(self) -> None:
self.evaluator = CitationMRREvaluator()
def test_first_citation_is_relevant(self) -> None:
assert self.evaluator.evaluate(_ctx(["a", "b"], ["a"])) == 1.0
def test_second_citation_is_relevant(self) -> None:
assert self.evaluator.evaluate(_ctx(["a", "b"], ["b"])) == 0.5
def test_no_citations(self) -> None:
assert self.evaluator.evaluate(_ctx([], ["a"])) == 0.0
def test_no_relevant(self) -> None:
assert self.evaluator.evaluate(_ctx(["a"], [])) == 0.0
def test_no_matches(self) -> None:
assert self.evaluator.evaluate(_ctx(["a", "b"], ["c"])) == 0.0
def test_metadata_none(self) -> None:
ctx = MagicMock()
ctx.metadata = None
ctx.attributes = {"cited_uris": ["a"]}
assert self.evaluator.evaluate(ctx) == 0.0
def test_attribute_missing(self) -> None:
ctx = MagicMock()
ctx.metadata = {"relevant_uris": ["a"]}
ctx.attributes = {}
assert self.evaluator.evaluate(ctx) == 0.0
def test_evaluation_name(self) -> None:
assert self.evaluator.evaluation_name == "cited_mrr"
class TestCitationMAPEvaluator:
def setup_method(self) -> None:
self.evaluator = CitationMAPEvaluator()
def test_all_relevant_first(self) -> None:
# Both relevant docs cited at ranks 1 and 2: AP = (1/1 + 2/2) / 2 = 1.0
assert self.evaluator.evaluate(_ctx(["a", "b"], ["a", "b"])) == 1.0
def test_partial_match(self) -> None:
# Cited a, x, b. relevant a, b. P@1 = 1/1, P@3 = 2/3. AP = (1 + 2/3)/2
assert (
self.evaluator.evaluate(_ctx(["a", "x", "b"], ["a", "b"]))
== (1.0 + 2 / 3) / 2
)
def test_no_matches(self) -> None:
assert self.evaluator.evaluate(_ctx(["x", "y"], ["a", "b"])) == 0.0
def test_no_relevant(self) -> None:
assert self.evaluator.evaluate(_ctx(["a"], [])) == 0.0
def test_no_citations(self) -> None:
assert self.evaluator.evaluate(_ctx([], ["a"])) == 0.0
def test_metadata_none(self) -> None:
ctx = MagicMock()
ctx.metadata = None
ctx.attributes = {"cited_uris": ["a"]}
assert self.evaluator.evaluate(ctx) == 0.0
def test_evaluation_name(self) -> None:
assert self.evaluator.evaluation_name == "cited_map"