File size: 3,754 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
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
"""Flattens raw benchmark results into JSON + CSV artifacts."""

import csv
import json
import os

MODEL_FIELDS = [
    "model",
    "embedding_dim",
    "model_size_mb",
    "param_count",
    "load_time_sec",
    "doc_embed_time_sec",
    "doc_throughput_docs_per_sec",
    "raw_vector_storage_mb",
    "qdrant_collection_disk_mb",
    "peak_query_encode_memory_mb",
]

MODE_FIELDS = [
    "mean_embedding_latency_ms",
    "mean_retrieval_latency_ms",
    "queries_per_sec",
    "recall@10",
    "recall@50",
    "mrr",
    "ndcg@10",
]

SUMMARY_FIELDS = MODEL_FIELDS[:1] + ["retrieval_mode", "query_mode"] + MODEL_FIELDS[1:] + MODE_FIELDS

CATEGORY_FIELDS = [
    "model",
    "retrieval_mode",
    "query_mode",
    "category",
    "n",
    "recall@10",
    "recall@50",
    "mrr",
    "ndcg@10",
]


def _summary_rows(results):
    rows = []
    for model_result in results:
        for retrieval_mode, query_mode_results in model_result["retrieval_modes"].items():
            for query_mode, mode_result in query_mode_results.items():
                row = {k: model_result.get(k) for k in MODEL_FIELDS}
                row["retrieval_mode"] = retrieval_mode
                row["query_mode"] = query_mode
                for metric in MODE_FIELDS:
                    row[metric] = mode_result.get(metric)
                rows.append(row)
    return rows


def _category_rows(results):
    rows = []
    for model_result in results:
        for retrieval_mode, query_mode_results in model_result["retrieval_modes"].items():
            for query_mode, mode_result in query_mode_results.items():
                for category, stats in mode_result["per_category"].items():
                    rows.append(
                        {
                            "model": model_result["model"],
                            "retrieval_mode": retrieval_mode,
                            "query_mode": query_mode,
                            "category": category,
                            **stats,
                        }
                    )
    return rows


def _write_csv(path, fieldnames, rows):
    with open(path, "w", newline="", encoding="utf-8") as f:
        writer = csv.DictWriter(f, fieldnames=fieldnames)
        writer.writeheader()
        for row in rows:
            writer.writerow(row)


def save_results(results, results_dir):
    os.makedirs(results_dir, exist_ok=True)

    raw_path = os.path.join(results_dir, "raw_results.json")
    with open(raw_path, "w", encoding="utf-8") as f:
        json.dump(results, f, ensure_ascii=False, indent=2)

    summary_rows = _summary_rows(results)
    summary_path = os.path.join(results_dir, "summary.csv")
    _write_csv(summary_path, SUMMARY_FIELDS, summary_rows)

    category_rows = _category_rows(results)
    category_path = os.path.join(results_dir, "summary_by_category.csv")
    _write_csv(category_path, CATEGORY_FIELDS, category_rows)

    return {
        "raw_results": raw_path,
        "summary": summary_path,
        "summary_by_category": category_path,
    }


def print_summary_table(results):
    header = (
        f"{'model':<55} {'retrieval':<8} {'mode':<11} {'dim':>5} {'size(MB)':>9} "
        f"{'R@10':>6} {'R@50':>6} {'MRR':>6} {'nDCG@10':>8} {'lat(ms)':>8} {'q/s':>7}"
    )
    print(header)
    print("-" * len(header))
    for row in _summary_rows(results):
        print(
            f"{row['model']:<55} {row['retrieval_mode']:<8} {row['query_mode']:<11} "
            f"{row['embedding_dim']:>5} {row['model_size_mb']:>9.1f} "
            f"{row['recall@10']:>6.3f} {row['recall@50']:>6.3f} {row['mrr']:>6.3f} "
            f"{row['ndcg@10']:>8.3f} {row['mean_retrieval_latency_ms']:>8.1f} "
            f"{row['queries_per_sec']:>7.1f}"
        )