File size: 5,614 Bytes
6dfa658
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
from __future__ import annotations

import argparse
import json
import subprocess
import sys
from pathlib import Path


ROOT = Path(__file__).resolve().parents[1]


def run_command(command: list[str]) -> None:
    print("\n$", " ".join(command))
    subprocess.run(command, cwd=ROOT, check=True)


def read_summary(path: Path) -> dict[str, float]:
    with path.open("r", encoding="utf-8") as f:
        data = json.load(f)
    return data["summary"]


def main() -> None:
    parser = argparse.ArgumentParser(
        description="Run Base RAG and Fine-tuned RAG on the same corpus, benchmark, and answer generator"
    )
    parser.add_argument("--data-dir", type=Path, default=Path("data"))
    parser.add_argument("--output-dir", type=Path, default=Path("outputs/submission_eval"))
    parser.add_argument("--limit", type=int, default=None)
    parser.add_argument("--top-k", type=int, default=5)
    parser.add_argument("--generation-mode", choices=["extractive", "local_hf", "ollama"], default="extractive")
    parser.add_argument("--generation-model", default=None)
    parser.add_argument("--max-new-tokens", type=int, default=192)
    parser.add_argument("--temperature", type=float, default=0.0)
    parser.add_argument("--base-retriever", choices=["bm25", "dense", "hybrid"], default="bm25")
    parser.add_argument("--base-embedding-model", default="intfloat/multilingual-e5-base")
    parser.add_argument("--finetuned-retriever", choices=["bm25", "dense", "hybrid"], default="bm25")
    parser.add_argument("--finetuned-embedding-model", default=None)
    parser.add_argument("--finetuned-reranker-model", default=None)
    parser.add_argument("--rerank-top-n", type=int, default=15)
    args = parser.parse_args()

    has_finetuned_component = bool(args.finetuned_embedding_model or args.finetuned_reranker_model)
    if not has_finetuned_component:
        print(
            "WARNING: No fine-tuned embedding or reranker checkpoint was provided. "
            "This run is a smoke check with the same retrieval setup for both systems. "
            "For a real fine-tuned RAG comparison, pass --finetuned-reranker-model or "
            "--finetuned-embedding-model."
        )

    args.output_dir.mkdir(parents=True, exist_ok=True)
    base_output = args.output_dir / "base_rag_qa.json"
    finetuned_output = args.output_dir / "finetuned_rag_qa.json"
    comparison_output = args.output_dir / "base_vs_finetuned_summary.json"

    common = [
        sys.executable,
        "scripts/evaluate_qa.py",
        "--data-dir",
        str(args.data_dir),
        "--top-k",
        str(args.top_k),
        "--generation-mode",
        args.generation_mode,
        "--max-new-tokens",
        str(args.max_new_tokens),
        "--temperature",
        str(args.temperature),
    ]
    if args.generation_model:
        common += ["--generation-model", args.generation_model]
    if args.limit:
        common += ["--limit", str(args.limit)]

    base_cmd = common + [
        "--retriever",
        args.base_retriever,
        "--embedding-model",
        args.base_embedding_model,
        "--output",
        str(base_output),
    ]

    finetuned_cmd = common + [
        "--retriever",
        args.finetuned_retriever,
        "--output",
        str(finetuned_output),
    ]
    if args.finetuned_embedding_model:
        finetuned_cmd += ["--embedding-model", args.finetuned_embedding_model]
    else:
        finetuned_cmd += ["--embedding-model", args.base_embedding_model]
    if args.finetuned_reranker_model:
        finetuned_cmd += [
            "--reranker-model",
            args.finetuned_reranker_model,
            "--rerank-top-n",
            str(args.rerank_top_n),
        ]

    run_command(base_cmd)
    run_command(finetuned_cmd)

    base_summary = read_summary(base_output)
    finetuned_summary = read_summary(finetuned_output)
    deltas = {
        key: finetuned_summary.get(key, 0.0) - base_summary.get(key, 0.0)
        for key in sorted(set(base_summary) | set(finetuned_summary))
    }
    comparison = {
        "note": "Both systems were evaluated on the same corpus, benchmark, and answer generator.",
        "warning": None
        if has_finetuned_component
        else "No fine-tuned embedding or reranker checkpoint was provided; this is a smoke check, not a real fine-tuned RAG comparison.",
        "base_config": {
            "retriever": args.base_retriever,
            "embedding_model": args.base_embedding_model if args.base_retriever != "bm25" else None,
            "generation_mode": args.generation_mode,
            "generation_model": args.generation_model,
        },
        "finetuned_config": {
            "retriever": args.finetuned_retriever,
            "embedding_model": (args.finetuned_embedding_model or args.base_embedding_model)
            if args.finetuned_retriever != "bm25"
            else None,
            "reranker_model": args.finetuned_reranker_model,
            "generation_mode": args.generation_mode,
            "generation_model": args.generation_model,
        },
        "base_summary": base_summary,
        "finetuned_summary": finetuned_summary,
        "delta_finetuned_minus_base": deltas,
    }
    with comparison_output.open("w", encoding="utf-8") as f:
        json.dump(comparison, f, ensure_ascii=False, indent=2)

    print("\nComparison summary")
    for key in sorted(deltas):
        print(f"{key}: base={base_summary.get(key, 0.0):.4f} fine={finetuned_summary.get(key, 0.0):.4f} delta={deltas[key]:+.4f}")
    print(f"Wrote {comparison_output}")


if __name__ == "__main__":
    main()