Spaces:
Sleeping
Sleeping
Amrita P
feat: implement advanced RAG pipeline (cross-encoder, contextual chunks, streaming, confidence gating)
4f25e4a | """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 | |
| # --------------------------------------------------------------------------- | |
| 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 | |
| def composite_score(self) -> float: | |
| return (self.faithfulness_score + self.relevance_score + self.keyword_coverage) / 3 | |
| 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) | |
| 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) | |