from __future__ import annotations import argparse import json import traceback from dataclasses import dataclass from datetime import datetime from pathlib import Path from types import SimpleNamespace from typing import Any from eval.rag_eval import ( REPORT_DIR, build_index, ensure_dirs, evaluate_retrieval, load_eval_corpus, write_reports, ) DEFAULT_DATASETS = ["beir/scifact", "beir/fiqa", "open-ragbench", "local-options"] SMOKE_DEFAULTS = { "beir/scifact": {"max_corpus_docs": 200, "max_queries": 10}, "beir/fiqa": {"max_corpus_docs": 500, "max_queries": 10}, "open-ragbench": {"max_corpus_docs": 20, "max_queries": 5}, "t2-ragbench": {"max_corpus_docs": 20, "max_queries": 5}, "local-options": {"max_corpus_docs": None, "max_queries": 3}, } @dataclass class DatasetRun: dataset: str status: str metrics: dict[str, Any] | None json_report: str | None markdown_report: str | None error: str | None = None def parse_dataset_list(value: str) -> list[str]: datasets = [item.strip() for item in value.split(",") if item.strip()] return datasets or DEFAULT_DATASETS def build_dataset_args(args: argparse.Namespace, dataset: str) -> SimpleNamespace: defaults = SMOKE_DEFAULTS.get(dataset, {"max_corpus_docs": None, "max_queries": None}) return SimpleNamespace( dataset=dataset, split=args.split, top_k=args.top_k, chunk_size=args.chunk_size, chunk_overlap=args.chunk_overlap, max_corpus_docs=args.max_corpus_docs if args.max_corpus_docs is not None else defaults["max_corpus_docs"], max_queries=args.max_queries if args.max_queries is not None else defaults["max_queries"], rebuild=args.rebuild, use_hybrid=args.use_hybrid, use_reranker=args.use_reranker, reranker_model=args.reranker_model, reranker_candidates=args.reranker_candidates, ) def run_one(dataset: str, args: argparse.Namespace) -> DatasetRun: dataset_args = build_dataset_args(args, dataset) print( f"\n=== Running {dataset} " f"(top_k={dataset_args.top_k}, max_corpus_docs={dataset_args.max_corpus_docs}, " f"max_queries={dataset_args.max_queries}, rebuild={dataset_args.rebuild}, " f"use_hybrid={dataset_args.use_hybrid}, " f"use_reranker={dataset_args.use_reranker}) ===" ) corpus = load_eval_corpus(dataset_args) index = build_index( corpus, chunk_size=dataset_args.chunk_size, chunk_overlap=dataset_args.chunk_overlap, rebuild=dataset_args.rebuild, ) report = evaluate_retrieval( corpus, index, dataset_args.top_k, use_hybrid=dataset_args.use_hybrid, chunk_size=dataset_args.chunk_size, chunk_overlap=dataset_args.chunk_overlap, use_reranker=dataset_args.use_reranker, reranker_model_name=dataset_args.reranker_model, reranker_candidates=dataset_args.reranker_candidates, ) json_path, md_path = write_reports(report) print(json.dumps(report["metrics"], ensure_ascii=False, indent=2)) return DatasetRun( dataset=dataset, status="passed", metrics=report["metrics"], json_report=str(json_path), markdown_report=str(md_path), ) def write_suite_report(runs: list[DatasetRun], output_name: str | None) -> tuple[Path, Path]: ensure_dirs() timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") stem = output_name or f"rag_eval_suite_{timestamp}" json_path = REPORT_DIR / f"{stem}.json" md_path = REPORT_DIR / f"{stem}.md" payload = { "created_at": datetime.now().isoformat(timespec="seconds"), "runs": [run.__dict__ for run in runs], } json_path.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8") lines = ["# RAG Eval Suite", ""] for run in runs: lines.append(f"## {run.dataset}") lines.append("") lines.append(f"- status: `{run.status}`") if run.error: lines.append(f"- error: `{run.error}`") if run.metrics: for key, value in run.metrics.items(): lines.append(f"- `{key}`: {value:.4f}" if isinstance(value, float) else f"- `{key}`: {value}") if run.markdown_report: lines.append(f"- report: `{run.markdown_report}`") lines.append("") md_path.write_text("\n".join(lines), encoding="utf-8") return json_path, md_path def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description="Run a RAG retrieval eval suite.") parser.add_argument( "--datasets", default=",".join(DEFAULT_DATASETS), help="Comma-separated datasets: beir/scifact, beir/fiqa, open-ragbench, t2-ragbench, local-options", ) parser.add_argument("--split", default="test") parser.add_argument("--top-k", type=int, default=5) parser.add_argument("--chunk-size", type=int, default=512) parser.add_argument("--chunk-overlap", type=int, default=64) parser.add_argument("--max-corpus-docs", type=int, default=None) parser.add_argument("--max-queries", type=int, default=None) parser.add_argument("--rebuild", action="store_true") parser.add_argument("--use-hybrid", action="store_true") parser.add_argument("--use-reranker", action="store_true") parser.add_argument("--reranker-model", default="cross-encoder/ms-marco-MiniLM-L-6-v2") parser.add_argument("--reranker-candidates", type=int, default=25) parser.add_argument("--fail-fast", action="store_true") parser.add_argument("--output-name", default=None, help="Suite report filename stem under eval/reports.") return parser.parse_args() def main() -> None: args = parse_args() runs: list[DatasetRun] = [] for dataset in parse_dataset_list(args.datasets): try: runs.append(run_one(dataset, args)) except Exception as exc: error = f"{type(exc).__name__}: {exc}" print(f"\n*** {dataset} failed: {error}") if args.fail_fast: raise traceback.print_exc() runs.append( DatasetRun( dataset=dataset, status="failed", metrics=None, json_report=None, markdown_report=None, error=error, ) ) json_path, md_path = write_suite_report(runs, args.output_name) print(f"\nSuite JSON report: {json_path}") print(f"Suite Markdown report: {md_path}") if any(run.status == "failed" for run in runs): raise SystemExit(1) if __name__ == "__main__": main()