Spaces:
Runtime error
Runtime error
| 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}, | |
| } | |
| 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() | |