rag-document-qa / evaluation /eval_e2e.py
Amrita P
feat: implement advanced RAG pipeline (cross-encoder, contextual chunks, streaming, confidence gating)
4f25e4a
Raw
History Blame Contribute Delete
10.4 kB
"""End-to-end evaluation: faithfulness, answer relevance, and keyword coverage.
Usage (from the rag-qa/ directory):
python -m evaluation.eval_e2e path/to/doc.pdf [more.pdf ...]
"""
import logging
import sys
from dataclasses import dataclass, field
from pathlib import Path
from ingestion.embedder import Embedder
from ingestion.pipeline import IngestionPipeline
from retrieval.index import VectorIndex
from retrieval.searcher import search
from generation.generator import Generator
from evaluation.judge import LLMJudge
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Placeholder test cases — replace with real questions from your PDFs.
# ---------------------------------------------------------------------------
PLACEHOLDER_TEST_CASES: list[dict] = [
{
"question": "What two components does RAG combine to generate answers?",
"expected_source": "original_rag_paper.pdf",
"expected_page": 1,
"expected_answer_contains": ["retrieval", "generation"],
},
{
"question": "Which dataset is used to evaluate open-domain QA in the RAG paper?",
"expected_source": "original_rag_paper.pdf",
"expected_page": 6,
"expected_answer_contains": ["Natural Questions", "TriviaQA"],
},
{
"question": "What is the role of the retriever in the RAG architecture?",
"expected_source": "original_rag_paper.pdf",
"expected_page": 2,
"expected_answer_contains": ["relevant", "documents", "passages"],
},
]
# ---------------------------------------------------------------------------
# Dataclasses
# ---------------------------------------------------------------------------
@dataclass
class CaseResult:
question: str
retrieval_hit: bool
faithfulness_score: float
faithfulness_explanation: str
relevance_score: float
relevance_explanation: str
keyword_coverage: float
keywords_found: list[str] = field(default_factory=list)
keywords_missing: list[str] = field(default_factory=list)
confidence_level: str = "high"
generated_answer: str = "" # stored for debugging; not printed by default
@property
def composite_score(self) -> float:
return (self.faithfulness_score + self.relevance_score + self.keyword_coverage) / 3
@dataclass
class E2EMetrics:
n_queries: int
retrieval_hit_rate: float
avg_faithfulness: float
avg_relevance: float
avg_keyword_coverage: float
per_case: list[CaseResult] = field(default_factory=list)
@property
def avg_composite(self) -> float:
return (self.avg_faithfulness + self.avg_relevance + self.avg_keyword_coverage) / 3
def worst_cases(self, n: int = 3) -> list[CaseResult]:
return sorted(self.per_case, key=lambda c: c.composite_score)[:n]
def __str__(self) -> str:
return (
f"Hit={self.retrieval_hit_rate:.2f} "
f"Faith={self.avg_faithfulness:.2f} "
f"Relev={self.avg_relevance:.2f} "
f"Keywords={self.avg_keyword_coverage:.2f} "
f"(n={self.n_queries})"
)
# ---------------------------------------------------------------------------
# Core evaluation
# ---------------------------------------------------------------------------
def evaluate_e2e(
test_cases: list[dict],
embedder: Embedder,
index: VectorIndex,
generator: Generator,
rate_limit_delay: float = 1.0,
k: int = 5,
use_reranking: bool = True,
) -> E2EMetrics:
"""Run the full search→generate pipeline on each test case and score outputs.
Test case keys:
question (str, required)
expected_source (str, optional) — filename for retrieval-hit check
expected_page (int, optional) — 1-based page for retrieval-hit check
expected_answer_contains (list[str], optional) — keywords for coverage check
Args:
test_cases: List of test-case dicts.
embedder: Shared Embedder (same model used at ingest time).
index: Populated VectorIndex.
generator: Generator instance (Gemini model).
rate_limit_delay: Seconds to wait between Gemini judge calls.
k: Chunks to retrieve per query.
use_reranking: Passed through to search(); default True.
"""
if not test_cases:
raise ValueError("test_cases must be non-empty")
judge = LLMJudge(rate_limit_delay=rate_limit_delay)
case_results: list[CaseResult] = []
for i, case in enumerate(test_cases, start=1):
question = case["question"]
logger.info("[%d/%d] %s", i, len(test_cases), question[:70])
search_resp = search(question, embedder, index, k=k, use_reranking=use_reranking)
answer_obj = generator.generate_answer(
question, search_resp.chunks, max_score=search_resp.max_score
)
answer_text = answer_obj.answer
hit = _retrieval_hit(
search_resp.chunks,
case.get("expected_source", ""),
int(case.get("expected_page", -1)),
)
context_text = "\n\n".join(
f"[Source {j}] {r.metadata.get('source','?')}, p.{r.metadata.get('page_num','?')}\n"
f"{r.metadata.get('text','')}"
for j, r in enumerate(search_resp.chunks, start=1)
)
faith_score, faith_expl = judge.score_faithfulness(context_text, answer_text)
relev_score, relev_expl = judge.score_relevance(question, answer_text)
expected_kw: list[str] = case.get("expected_answer_contains", [])
answer_lower = answer_text.lower()
found = [kw for kw in expected_kw if kw.lower() in answer_lower]
missing = [kw for kw in expected_kw if kw.lower() not in answer_lower]
coverage = len(found) / len(expected_kw) if expected_kw else 1.0
case_results.append(CaseResult(
question=question,
retrieval_hit=hit,
faithfulness_score=faith_score,
faithfulness_explanation=faith_expl,
relevance_score=relev_score,
relevance_explanation=relev_expl,
keyword_coverage=coverage,
keywords_found=found,
keywords_missing=missing,
confidence_level=answer_obj.confidence_level,
generated_answer=answer_text,
))
n = len(case_results)
return E2EMetrics(
n_queries=n,
retrieval_hit_rate=sum(c.retrieval_hit for c in case_results) / n,
avg_faithfulness=sum(c.faithfulness_score for c in case_results) / n,
avg_relevance=sum(c.relevance_score for c in case_results) / n,
avg_keyword_coverage=sum(c.keyword_coverage for c in case_results) / n,
per_case=case_results,
)
# ---------------------------------------------------------------------------
# Printing
# ---------------------------------------------------------------------------
def print_e2e_report(metrics: E2EMetrics, show_explanations: bool = False) -> None:
"""Print a per-question table, aggregate averages, and worst performers."""
Q, H, F, R, K = 40, 7, 8, 8, 10
total = Q + H + F + R + K
div = "-" * total
print(f"\n{'End-to-End Evaluation':^{total}}")
print(div)
print("Question".ljust(Q) + "Hit".ljust(H) + "Faith.".ljust(F) + "Relev.".ljust(R) + "Keywords".ljust(K))
print(div)
for c in metrics.per_case:
q = (c.question[:Q - 2] + "…") if len(c.question) > Q - 1 else c.question
kw = f"{c.keyword_coverage:.2f}" + (f" (-{len(c.keywords_missing)})" if c.keywords_missing else "")
print(
q.ljust(Q)
+ ("✓" if c.retrieval_hit else "✗").ljust(H)
+ f"{c.faithfulness_score:.2f}".ljust(F)
+ f"{c.relevance_score:.2f}".ljust(R)
+ kw.ljust(K)
)
if show_explanations:
print(f" Faith: {c.faithfulness_explanation}")
print(f" Relev: {c.relevance_explanation}")
if c.keywords_missing:
print(f" Missing: {c.keywords_missing}")
print(div)
print(
"AVERAGE".ljust(Q)
+ f"{metrics.retrieval_hit_rate:.2f}".ljust(H)
+ f"{metrics.avg_faithfulness:.2f}".ljust(F)
+ f"{metrics.avg_relevance:.2f}".ljust(R)
+ f"{metrics.avg_keyword_coverage:.2f}".ljust(K)
)
print(div)
worst = metrics.worst_cases(n=min(3, metrics.n_queries))
if worst:
print("\nWorst-performing questions (by composite score):")
for c in worst:
q = c.question[:70] + ("…" if len(c.question) > 70 else "")
print(f" [{c.composite_score:.2f}] {q}")
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _retrieval_hit(chunks, expected_source: str, expected_page: int) -> bool:
if not expected_source or expected_page < 0:
return False
for result in chunks:
meta = result.metadata
if (Path(meta.get("source", "")).name == Path(expected_source).name
and meta.get("page_num") == expected_page):
return True
return False
# ---------------------------------------------------------------------------
# Entry point
# ---------------------------------------------------------------------------
if __name__ == "__main__":
logging.basicConfig(level=logging.INFO, format="%(levelname)s | %(name)s | %(message)s")
if len(sys.argv) < 2:
print("Usage: python -m evaluation.eval_e2e <path/to/doc.pdf> [more.pdf ...]")
sys.exit(0)
pdf_paths = [Path(p) for p in sys.argv[1:]]
print("Loading embedder and generator...")
embedder = Embedder()
index = VectorIndex(dimension=embedder.dimension)
pipeline = IngestionPipeline(embedder=embedder, index=index, strategy="recursive_character")
generator = Generator()
for p in pdf_paths:
r = pipeline.ingest_pdf(p)
print(f"Ingested {r.file}: {r.chunks} chunks" if not r.error else f"Error: {r.error}")
print("\nRunning end-to-end evaluation...")
metrics = evaluate_e2e(PLACEHOLDER_TEST_CASES, embedder, index, generator)
print_e2e_report(metrics, show_explanations=True)