mathidot111/First_agent_template
0
1from __future__ import annotations2 3import argparse4import json5import traceback6from dataclasses import dataclass7from datetime import datetime8from pathlib import Path9from types import SimpleNamespace10from typing import Any11 12from eval.rag_eval import (13 REPORT_DIR,14 build_index,15 ensure_dirs,16 evaluate_retrieval,17 load_eval_corpus,18 write_reports,19)20 21 22DEFAULT_DATASETS = ["beir/scifact", "beir/fiqa", "open-ragbench", "local-options"]23SMOKE_DEFAULTS = {24 "beir/scifact": {"max_corpus_docs": 200, "max_queries": 10},25 "beir/fiqa": {"max_corpus_docs": 500, "max_queries": 10},26 "open-ragbench": {"max_corpus_docs": 20, "max_queries": 5},27 "t2-ragbench": {"max_corpus_docs": 20, "max_queries": 5},28 "local-options": {"max_corpus_docs": None, "max_queries": 3},29}30 31 32@dataclass33class DatasetRun:34 dataset: str35 status: str36 metrics: dict[str, Any] | None37 json_report: str | None38 markdown_report: str | None39 error: str | None = None40 41 42def parse_dataset_list(value: str) -> list[str]:43 datasets = [item.strip() for item in value.split(",") if item.strip()]44 return datasets or DEFAULT_DATASETS45 46 47def build_dataset_args(args: argparse.Namespace, dataset: str) -> SimpleNamespace:48 defaults = SMOKE_DEFAULTS.get(dataset, {"max_corpus_docs": None, "max_queries": None})49 return SimpleNamespace(50 dataset=dataset,51 split=args.split,52 top_k=args.top_k,53 chunk_size=args.chunk_size,54 chunk_overlap=args.chunk_overlap,55 max_corpus_docs=args.max_corpus_docs56 if args.max_corpus_docs is not None57 else defaults["max_corpus_docs"],58 max_queries=args.max_queries if args.max_queries is not None else defaults["max_queries"],59 rebuild=args.rebuild,60 use_hybrid=args.use_hybrid,61 use_reranker=args.use_reranker,62 reranker_model=args.reranker_model,63 reranker_candidates=args.reranker_candidates,64 )65 66 67def run_one(dataset: str, args: argparse.Namespace) -> DatasetRun:68 dataset_args = build_dataset_args(args, dataset)69 print(70 f"\n=== Running {dataset} "71 f"(top_k={dataset_args.top_k}, max_corpus_docs={dataset_args.max_corpus_docs}, "72 f"max_queries={dataset_args.max_queries}, rebuild={dataset_args.rebuild}, "73 f"use_hybrid={dataset_args.use_hybrid}, "74 f"use_reranker={dataset_args.use_reranker}) ==="75 )76 77 corpus = load_eval_corpus(dataset_args)78 index = build_index(79 corpus,80 chunk_size=dataset_args.chunk_size,81 chunk_overlap=dataset_args.chunk_overlap,82 rebuild=dataset_args.rebuild,83 )84 report = evaluate_retrieval(85 corpus,86 index,87 dataset_args.top_k,88 use_hybrid=dataset_args.use_hybrid,89 chunk_size=dataset_args.chunk_size,90 chunk_overlap=dataset_args.chunk_overlap,91 use_reranker=dataset_args.use_reranker,92 reranker_model_name=dataset_args.reranker_model,93 reranker_candidates=dataset_args.reranker_candidates,94 )95 json_path, md_path = write_reports(report)96 print(json.dumps(report["metrics"], ensure_ascii=False, indent=2))97 98 return DatasetRun(99 dataset=dataset,100 status="passed",101 metrics=report["metrics"],102 json_report=str(json_path),103 markdown_report=str(md_path),104 )105 106 107def write_suite_report(runs: list[DatasetRun], output_name: str | None) -> tuple[Path, Path]:108 ensure_dirs()109 timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")110 stem = output_name or f"rag_eval_suite_{timestamp}"111 json_path = REPORT_DIR / f"{stem}.json"112 md_path = REPORT_DIR / f"{stem}.md"113 114 payload = {115 "created_at": datetime.now().isoformat(timespec="seconds"),116 "runs": [run.__dict__ for run in runs],117 }118 json_path.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8")119 120 lines = ["# RAG Eval Suite", ""]121 for run in runs:122 lines.append(f"## {run.dataset}")123 lines.append("")124 lines.append(f"- status: `{run.status}`")125 if run.error:126 lines.append(f"- error: `{run.error}`")127 if run.metrics:128 for key, value in run.metrics.items():129 lines.append(f"- `{key}`: {value:.4f}" if isinstance(value, float) else f"- `{key}`: {value}")130 if run.markdown_report:131 lines.append(f"- report: `{run.markdown_report}`")132 lines.append("")133 md_path.write_text("\n".join(lines), encoding="utf-8")134 return json_path, md_path135 136 137def parse_args() -> argparse.Namespace:138 parser = argparse.ArgumentParser(description="Run a RAG retrieval eval suite.")139 parser.add_argument(140 "--datasets",141 default=",".join(DEFAULT_DATASETS),142 help="Comma-separated datasets: beir/scifact, beir/fiqa, open-ragbench, t2-ragbench, local-options",143 )144 parser.add_argument("--split", default="test")145 parser.add_argument("--top-k", type=int, default=5)146 parser.add_argument("--chunk-size", type=int, default=512)147 parser.add_argument("--chunk-overlap", type=int, default=64)148 parser.add_argument("--max-corpus-docs", type=int, default=None)149 parser.add_argument("--max-queries", type=int, default=None)150 parser.add_argument("--rebuild", action="store_true")151 parser.add_argument("--use-hybrid", action="store_true")152 parser.add_argument("--use-reranker", action="store_true")153 parser.add_argument("--reranker-model", default="cross-encoder/ms-marco-MiniLM-L-6-v2")154 parser.add_argument("--reranker-candidates", type=int, default=25)155 parser.add_argument("--fail-fast", action="store_true")156 parser.add_argument("--output-name", default=None, help="Suite report filename stem under eval/reports.")157 return parser.parse_args()158 159 160def main() -> None:161 args = parse_args()162 runs: list[DatasetRun] = []163 164 for dataset in parse_dataset_list(args.datasets):165 try:166 runs.append(run_one(dataset, args))167 except Exception as exc:168 error = f"{type(exc).__name__}: {exc}"169 print(f"\n*** {dataset} failed: {error}")170 if args.fail_fast:171 raise172 traceback.print_exc()173 runs.append(174 DatasetRun(175 dataset=dataset,176 status="failed",177 metrics=None,178 json_report=None,179 markdown_report=None,180 error=error,181 )182 )183 184 json_path, md_path = write_suite_report(runs, args.output_name)185 print(f"\nSuite JSON report: {json_path}")186 print(f"Suite Markdown report: {md_path}")187 188 if any(run.status == "failed" for run in runs):189 raise SystemExit(1)190 191 192if __name__ == "__main__":193 main()194 