File size: 2,247 Bytes
6ced533
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""CLI entry point for the embedding-model retrieval benchmark.

Usage:
    python run_benchmark.py                       # run the full model set from benchmark/config.py
    python run_benchmark.py --models kazalbrur/bangla-embed-e5-small-banglish
    python run_benchmark.py --query-modes raw      # skip the normalized-query experiment
"""

import argparse

from benchmark import config, report
from benchmark.runner import run_benchmark


def parse_args():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--models",
        nargs="+",
        default=None,
        help="HF model names to benchmark (default: full set in benchmark/config.py)",
    )
    parser.add_argument(
        "--query-modes",
        nargs="+",
        choices=["raw", "normalized"],
        default=None,
        help="Which query variants to evaluate (default: both)",
    )
    parser.add_argument(
        "--retrieval-modes",
        nargs="+",
        choices=["dense", "hybrid"],
        default=None,
        help="dense (embedding-only) and/or hybrid (BM25 + dense RRF) (default: both)",
    )
    parser.add_argument("--top-k", type=int, default=None)
    parser.add_argument("--results-dir", default=config.RESULTS_DIR)
    return parser.parse_args()


def main():
    args = parse_args()

    models = None
    if args.models:
        by_name = {m["name"]: m for m in config.MODELS}
        models = []
        for name in args.models:
            if name not in by_name:
                raise SystemExit(
                    f"Unknown model '{name}'. Add it to benchmark/config.py MODELS first "
                    f"(with its query/passage prefix)."
                )
            models.append(by_name[name])

    results = run_benchmark(
        models=models,
        query_modes=args.query_modes,
        retrieval_modes=args.retrieval_modes,
        top_k=args.top_k,
    )

    paths = report.save_results(results, args.results_dir)

    print("\n" + "=" * 100)
    print("SUMMARY")
    print("=" * 100)
    report.print_summary_table(results)

    print("\nSaved:")
    for label, path in paths.items():
        print(f"  {label}: {path}")


if __name__ == "__main__":
    main()