Team Ai
Apppublic

mathidot111/First_agent_template

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
run_eval_suite.py194 linesDownload Raw Back to eval
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