Evaluation

Retrieval, answer, citation, support, and uncertainty metrics.

src/rag_evolution/evaluation.py · 200 lines · sha256 bcbb14119b04…

"""Retrieval, answer, citation, support, and uncertainty metrics."""

import math
import random
import re
import statistics
from collections import Counter, defaultdict
from typing import Any, Dict, Iterable, List, Mapping, Sequence, Tuple

from .models import Answer, QuestionExample, SearchResult
from .text import content_terms, split_sentences, tokenize


def precision_at_k(ranked_ids: Sequence[str], relevant_ids: Iterable[str], k: int) -> float:
    relevant = set(relevant_ids)
    if k <= 0:
        return 0.0
    return sum(item in relevant for item in ranked_ids[:k]) / k


def recall_at_k(ranked_ids: Sequence[str], relevant_ids: Iterable[str], k: int) -> float:
    relevant = set(relevant_ids)
    if not relevant:
        return 1.0
    return len(set(ranked_ids[:k]) & relevant) / len(relevant)


def reciprocal_rank(ranked_ids: Sequence[str], relevant_ids: Iterable[str]) -> float:
    relevant = set(relevant_ids)
    for rank, item in enumerate(ranked_ids, start=1):
        if item in relevant:
            return 1.0 / rank
    return 0.0


def ndcg_at_k(ranked_ids: Sequence[str], relevant_ids: Iterable[str], k: int) -> float:
    relevant = set(relevant_ids)
    gains = [1.0 if item in relevant else 0.0 for item in ranked_ids[:k]]
    dcg = sum(gain / math.log2(index + 2) for index, gain in enumerate(gains))
    ideal_length = min(k, len(relevant))
    ideal = sum(1.0 / math.log2(index + 2) for index in range(ideal_length))
    return dcg / ideal if ideal else 1.0


def _normalize_answer(text: str) -> str:
    text = re.sub(r"\[[^\]]+\]", " ", text.lower())
    text = re.sub(r"\b(a|an|the)\b", " ", text)
    return " ".join(tokenize(text))


def exact_match(prediction: str, reference: str) -> float:
    return float(_normalize_answer(prediction) == _normalize_answer(reference))


def token_f1(prediction: str, reference: str) -> float:
    predicted = Counter(_normalize_answer(prediction).split())
    expected = Counter(_normalize_answer(reference).split())
    overlap = sum((predicted & expected).values())
    if not predicted or not expected:
        return float(predicted == expected)
    precision = overlap / sum(predicted.values())
    recall = overlap / sum(expected.values())
    return 2 * precision * recall / (precision + recall) if precision + recall else 0.0


def citation_precision(answer: Answer, relevant_document_ids: Iterable[str]) -> float:
    if not answer.citations:
        return 0.0
    relevant = set(relevant_document_ids)
    return sum(citation.document_id in relevant for citation in answer.citations) / len(
        answer.citations
    )


def citation_recall(answer: Answer, relevant_document_ids: Iterable[str]) -> float:
    relevant = set(relevant_document_ids)
    if not relevant:
        return 1.0
    cited = {citation.document_id for citation in answer.citations}
    return len(cited & relevant) / len(relevant)


def lexical_faithfulness(answer: Answer) -> float:
    """Cheap diagnostic: fraction of answer sentences substantially supported by a quote.

    This is intentionally not called semantic entailment.  A production study
    should add human labels or a calibrated NLI/LLM judge and report judge error.
    """

    if answer.abstained:
        return 1.0
    quotes = [set(content_terms(citation.quote)) for citation in answer.citations]
    sentences = split_sentences(re.sub(r"\[[^\]]+\]", "", answer.text))
    if not sentences:
        return 0.0
    supported = 0
    for sentence in sentences:
        terms = set(content_terms(sentence))
        if not terms:
            continue
        best = max((len(terms & quote) / len(terms) for quote in quotes), default=0.0)
        supported += best >= 0.65
    return supported / len(sentences)


def _document_ranking(results: Sequence[SearchResult]) -> List[str]:
    seen = set()
    documents = []
    for result in results:
        document_id = result.chunk.document_id
        if document_id not in seen:
            documents.append(document_id)
            seen.add(document_id)
    return documents


def evaluate_retriever(
    retriever: Any, examples: Sequence[QuestionExample], k: int = 5
) -> List[Mapping[str, Any]]:
    """Return per-example metrics so slices and confidence intervals remain possible."""

    rows: List[Mapping[str, Any]] = []
    for example in examples:
        results = retriever.search(example.question, k)
        ranked = _document_ranking(results)
        rows.append(
            {
                "id": example.id,
                "tags": example.tags,
                f"precision@{k}": precision_at_k(ranked, example.relevant_document_ids, k),
                f"recall@{k}": recall_at_k(ranked, example.relevant_document_ids, k),
                "mrr": reciprocal_rank(ranked, example.relevant_document_ids),
                f"ndcg@{k}": ndcg_at_k(ranked, example.relevant_document_ids, k),
            }
        )
    return rows


def evaluate_pipeline(pipeline: Any, examples: Sequence[QuestionExample]) -> List[Mapping[str, Any]]:
    rows: List[Mapping[str, Any]] = []
    for example in examples:
        answer = pipeline.ask(example.question)
        ranked = _document_ranking(answer.contexts)
        rows.append(
            {
                "id": example.id,
                "tags": example.tags,
                "recall@5": recall_at_k(ranked, example.relevant_document_ids, 5),
                "mrr": reciprocal_rank(ranked, example.relevant_document_ids),
                "answer_em": exact_match(answer.text, example.reference_answer),
                "answer_f1": token_f1(answer.text, example.reference_answer),
                "citation_precision": citation_precision(answer, example.relevant_document_ids),
                "citation_recall": citation_recall(answer, example.relevant_document_ids),
                "lexical_faithfulness": lexical_faithfulness(answer),
                "abstained": float(answer.abstained),
                "confidence": answer.confidence,
            }
        )
    return rows


def aggregate_metrics(rows: Sequence[Mapping[str, Any]]) -> Dict[str, float]:
    numeric: Dict[str, List[float]] = defaultdict(list)
    for row in rows:
        for key, value in row.items():
            if isinstance(value, (int, float)) and not isinstance(value, bool):
                numeric[key].append(float(value))
    return {key: statistics.mean(values) for key, values in numeric.items() if values}


def metrics_by_tag(rows: Sequence[Mapping[str, Any]]) -> Mapping[str, Mapping[str, float]]:
    grouped: Dict[str, List[Mapping[str, Any]]] = defaultdict(list)
    for row in rows:
        for tag in row.get("tags", ()):  # type: ignore[union-attr]
            grouped[str(tag)].append(row)
    return {tag: aggregate_metrics(tag_rows) for tag, tag_rows in sorted(grouped.items())}


def paired_bootstrap_delta(
    left: Sequence[float],
    right: Sequence[float],
    iterations: int = 2000,
    seed: int = 7,
) -> Tuple[float, float, float]:
    """Paired bootstrap mean difference and a percentile 95% interval."""

    if len(left) != len(right) or not left:
        raise ValueError("paired samples must be non-empty and equally sized")
    randomizer = random.Random(seed)
    size = len(left)
    samples = []
    for _ in range(iterations):
        indices = [randomizer.randrange(size) for _ in range(size)]
        samples.append(statistics.mean(left[index] - right[index] for index in indices))
    samples.sort()
    lower = samples[int(0.025 * (iterations - 1))]
    upper = samples[int(0.975 * (iterations - 1))]
    observed = statistics.mean(a - b for a, b in zip(left, right))
    return observed, lower, upper