File size: 12,459 Bytes
93a281a
 
5a40cb1
c65eea1
 
a6ca468
c65eea1
 
 
 
a6ca468
 
 
85944a2
c65eea1
 
 
 
 
 
 
 
 
85944a2
c65eea1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
85944a2
 
 
 
c65eea1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
85944a2
 
 
a6ca468
 
022010c
 
 
 
 
 
 
abc6ea2
022010c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
abc6ea2
022010c
 
abc6ea2
022010c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c65eea1
a6ca468
 
5a40cb1
 
022010c
c65eea1
85944a2
a6ca468
85944a2
c65eea1
a6ca468
85944a2
a6ca468
85944a2
c65eea1
85944a2
 
 
 
 
c65eea1
 
85944a2
c65eea1
 
85944a2
 
c65eea1
022010c
 
 
c65eea1
85944a2
 
c65eea1
a6ca468
c65eea1
a6ca468
c65eea1
 
 
 
 
5a40cb1
022010c
5a40cb1
c65eea1
 
 
 
 
 
 
 
85944a2
 
c65eea1
 
 
 
85944a2
c65eea1
85944a2
c65eea1
 
85944a2
5a40cb1
022010c
 
5a40cb1
85944a2
c65eea1
 
 
 
 
 
 
 
85944a2
 
 
 
 
c65eea1
85944a2
c65eea1
 
 
85944a2
 
 
 
c65eea1
 
 
 
 
85944a2
 
 
 
 
 
 
c65eea1
 
85944a2
c65eea1
 
85944a2
 
 
 
 
 
 
a6ca468
 
 
 
93a281a
5a40cb1
 
 
 
 
 
 
 
 
 
022010c
 
 
 
 
85944a2
 
5a40cb1
 
 
 
85944a2
 
 
 
 
 
 
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
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
from __future__ import annotations

import asyncio
import json
import logging
import sys
import time
from pathlib import Path
from statistics import mean, median

sys.path.insert(0, str(Path(__file__).resolve().parent.parent))

from src.core.config import DATA_DIR
from src.core.constants import GROUND_TRUTH_PATH, QUERIES_PATH
from src.search.bm25_search import BM25Search
from src.search.hybrid import HybridSearch
from src.search.vector_search import VectorSearch

logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)


def precision_at_k(retrieved: list[str], relevant: set[str], k: int) -> float:
    if k <= 0 or not retrieved:
        return 0.0
    top_k = retrieved[:k]
    if not top_k:
        return 0.0
    return len([doc for doc in top_k if doc in relevant]) / len(top_k)


def recall_at_k(retrieved: list[str], relevant: set[str], k: int) -> float:
    if not relevant:
        return 0.0
    top_k = retrieved[:k]
    hits = len([doc for doc in top_k if doc in relevant])
    return hits / len(relevant)


def mean_reciprocal_rank(retrieved: list[str], relevant: set[str]) -> float:
    for i, doc in enumerate(retrieved, start=1):
        if doc in relevant:
            return 1.0 / i
    return 0.0


def ndcg_at_k(retrieved: list[str], relevant: set[str], k: int) -> float:
    top_k = retrieved[:k]
    dcg = 0.0
    for i, doc in enumerate(top_k, start=1):
        rel = 1.0 if doc in relevant else 0.0
        dcg += (2 ** rel - 1) / (i.bit_length()) if i > 1 else rel
    ideal = min(len(relevant), k)
    idcg = sum(1.0 / (i.bit_length()) if i > 1 else 1.0 for i in range(1, ideal + 1))
    return dcg / idcg if idcg > 0 else 0.0


def cross_lingual_mrr(results: dict) -> float:
    non_en_queries = {
        qid for qid, q in results.get("queries", {}).items()
        if q.get("language", "en") != "en"
    }
    if not non_en_queries:
        return 1.0
    mrr_sum = 0.0
    for qid in non_en_queries:
        mrr_sum += results.get("mrr", {}).get(qid, 0.0)
    return mrr_sum / len(non_en_queries)


def latency_stats(latencies: list[float]) -> dict:
    if not latencies:
        return {"p50": 0, "p95": 0, "p99": 0, "mean": 0, "min": 0, "max": 0}
    sorted_lat = sorted(latencies)
    n = len(sorted_lat)
    return {
        "p50": sorted_lat[int(n * 0.50)],
        "p95": sorted_lat[int(n * 0.95)],
        "p99": sorted_lat[int(n * 0.99)],
        "mean": mean(latencies),
        "min": min(latencies),
        "max": max(latencies),
    }


def find_index_dir() -> Path | None:
    path = DATA_DIR / "indexes" / "faiss_index.bin"
    return path.parent if path.exists() else None


def _run_full_pipeline(
    queries_list: list,
    gt_map: dict,
    hybrid: HybridSearch,
    skip_reranker: bool,
) -> dict:
    """Run full pipeline across all queries in a single event loop."""
    from src.core.query_parser import expand_with_aliases, parse_query
    from src.search.reranker import CrossEncoderReranker
    from src.matching.scorer import CandidateScorer
    from src.agents.executor import ExecutorAgent
    from src.core.profile_store import ProfileStore

    index_dir = find_index_dir()
    logger.info("Loading full pipeline components...")
    reranker = CrossEncoderReranker(timeout_ms=0)
    scorer = CandidateScorer()
    profiles = ProfileStore()
    profiles.load_offset_index(index_dir / "offset_index.json")
    executor = ExecutorAgent(hybrid, reranker, scorer, profiles)

    async def process_queries():
        metrics: dict[str, list] = {
            "p@5": [], "p@10": [], "p@20": [],
            "r@5": [], "r@10": [], "r@20": [],
            "mrr": [], "ndcg@10": [], "latencies": [],
        }
        evaluated = 0
        skipped = 0

        for q in queries_list:
            query_text = q.get("query", q.get("raw_query", ""))
            qid = q.get("query_id", q.get("id", ""))
            relevant = set(gt_map.get(qid, []))
            if not query_text or not relevant:
                skipped += 1
                continue

            t0 = time.perf_counter()
            variants = expand_with_aliases(query_text)[:3]
            all_pid_scores: dict[str, float] = {}
            for variant in variants:
                parsed = parse_query(variant)
                results = await executor.execute(parsed, top_k=30, skip_reranker=skip_reranker)
                for r in results:
                    pid = r.profile_id
                    score = r.scores.overall
                    if pid not in all_pid_scores or score > all_pid_scores[pid]:
                        all_pid_scores[pid] = score

            retrieved = [pid for pid, _ in sorted(all_pid_scores.items(), key=lambda x: -x[1])]
            elapsed = (time.perf_counter() - t0) * 1000
            metrics["latencies"].append(elapsed)

            for k in (5, 10, 20):
                metrics[f"p@{k}"].append(precision_at_k(retrieved, relevant, k))
                metrics[f"r@{k}"].append(recall_at_k(retrieved, relevant, k))
            metrics["mrr"].append(mean_reciprocal_rank(retrieved, relevant))
            metrics["ndcg@10"].append(ndcg_at_k(retrieved, relevant, 10))
            evaluated += 1

            if evaluated % 50 == 0:
                logger.info(f"  processed {evaluated}/{len(queries_list)} queries...")

        metrics["total_queries"] = evaluated
        metrics["skipped"] = skipped

        summary: dict = {}
        for metric, values in metrics.items():
            if isinstance(values, list) and values:
                summary[metric] = {
                    "mean": round(mean(values), 4),
                    "median": round(median(values), 4),
                    "min": round(min(values), 4),
                    "max": round(max(values), 4),
                }
            elif isinstance(values, (int, float)):
                summary[metric] = values

        summary["latency"] = latency_stats(metrics["latencies"])
        for k in ("p50", "p95", "p99", "mean", "min", "max"):
            if k in summary["latency"]:
                summary["latency"][k] = round(summary["latency"][k], 1)
        summary["cross_lingual_mrr"] = round(cross_lingual_mrr({
            "queries": {i: q for i, q in enumerate(queries_list)},
            "mrr": {i: v for i, v in enumerate(metrics["mrr"])},
        }), 4)

        logger.info("Full pipeline evaluation results:")
        for m in ("p@5", "p@10", "r@5", "r@10", "mrr", "ndcg@10"):
            if m in summary:
                s = summary[m]
                logger.info(f"  {m}: mean={s['mean']:.4f}, median={s['median']:.4f}")
        lat = summary["latency"]
        logger.info(f"  latency: p50={lat['p50']:.0f}ms, p95={lat['p95']:.0f}ms")
        logger.info(f"  cross-lingual MRR: {summary['cross_lingual_mrr']:.4f}")
        return summary

    return asyncio.run(process_queries())


def evaluate(
    queries_path: Path = QUERIES_PATH,
    ground_truth_path: Path = GROUND_TRUTH_PATH,
    use_full_pipeline: bool = False,
    skip_reranker: bool = False,
    sample_n: int = 0,
) -> dict:
    index_dir = find_index_dir()
    if index_dir is None:
        logger.error("No indexes found. Run 'python scripts/build_indexes.py --sample 50' first.")
        return {}

    errors: list[str] = []
    if not queries_path.exists():
        errors.append(f"Queries file not found: {queries_path}")
    if not ground_truth_path.exists():
        errors.append(f"Ground truth file not found: {ground_truth_path}")
    if errors:
        for e in errors:
            logger.error(e)
        return {}

    with open(queries_path) as f:
        queries_raw = json.load(f)
    with open(ground_truth_path) as f:
        ground_truth = json.load(f)

    queries_list = queries_raw if isinstance(queries_raw, list) else list(queries_raw.values())
    gt_map = ground_truth if isinstance(ground_truth, dict) else {}
    if sample_n > 0:
        queries_list = queries_list[:sample_n]
        logger.info(f"Sampling {sample_n} queries for quick evaluation")

    logger.info(f"Loaded {len(queries_list)} queries and {len(gt_map)} ground truth entries")

    vector_search = VectorSearch()
    vector_search.load(index_dir / "faiss_index.bin", index_dir / "faiss_id_map.json")
    bm25_search = BM25Search()
    bm25_search.load(index_dir / "bm25_index.pkl")

    from src.language.multilingual import MultilingualEmbedder
    embedder = MultilingualEmbedder()
    hybrid = HybridSearch(vector_search, bm25_search, embedder)

    if use_full_pipeline:
        return _run_full_pipeline(queries_list, gt_map, hybrid, skip_reranker)

    all_metrics: dict[str, list] = {
        "p@5": [], "p@10": [], "p@20": [],
        "r@5": [], "r@10": [], "r@20": [],
        "mrr": [],
        "ndcg@10": [],
        "latencies": [],
    }

    evaluated = 0
    skipped = 0

    for q in queries_list:
        query_text = q.get("query", q.get("raw_query", ""))
        qid = q.get("query_id", q.get("id", ""))
        relevant = set(gt_map.get(qid, []))
        if not query_text or not relevant:
            skipped += 1
            continue

        t0 = time.perf_counter()

        results = hybrid.search(query_text, top_k=50)
        retrieved = [pid for pid, _ in results]

        elapsed = (time.perf_counter() - t0) * 1000
        all_metrics["latencies"].append(elapsed)

        for k in (5, 10, 20):
            all_metrics[f"p@{k}"].append(precision_at_k(retrieved, relevant, k))
            all_metrics[f"r@{k}"].append(recall_at_k(retrieved, relevant, k))

        all_metrics["mrr"].append(mean_reciprocal_rank(retrieved, relevant))
        all_metrics["ndcg@10"].append(ndcg_at_k(retrieved, relevant, 10))
        evaluated += 1

    if evaluated == 0:
        logger.error("No queries could be evaluated (check ground truth IDs match query IDs)")
        return {}

    summary: dict = {}
    for metric, values in all_metrics.items():
        if values:
            summary[metric] = {
                "mean": round(mean(values), 4),
                "median": round(median(values), 4),
                "min": round(min(values), 4),
                "max": round(max(values), 4),
            }
        else:
            summary[metric] = {}

    summary["latency"] = latency_stats(all_metrics["latencies"])
    for k in ("p50", "p95", "p99", "mean", "min", "max"):
        if k in summary["latency"]:
            summary["latency"][k] = round(summary["latency"][k], 1)

    summary["total_queries"] = evaluated
    summary["skipped"] = skipped
    summary["cross_lingual_mrr"] = round(cross_lingual_mrr({
        "queries": {i: q for i, q in enumerate(queries_list)},
        "mrr": {i: v for i, v in enumerate(all_metrics["mrr"])},
    }), 4)

    logger.info("Evaluation results:")
    for metric in ("p@5", "p@10", "r@5", "r@10", "mrr", "ndcg@10"):
        if metric in summary:
            s = summary[metric]
            logger.info(f"  {metric}: mean={s['mean']:.4f}, median={s['median']:.4f}")
    lat = summary["latency"]
    logger.info(f"  latency: p50={lat['p50']:.0f}ms, p95={lat['p95']:.0f}ms")
    logger.info(f"  cross-lingual MRR: {summary['cross_lingual_mrr']:.4f}")

    return summary


if __name__ == "__main__":
    import argparse
    parser = argparse.ArgumentParser(description="Evaluate search quality")
    parser.add_argument("--full-pipeline", action="store_true",
                        help="Use full executor pipeline (cross-encoder + scorer)")
    parser.add_argument("--skip-reranker", action="store_true",
                        help="Skip cross-encoder, use hybrid RRF scores for ranking")
    parser.add_argument("--sample", type=int, default=0,
                        help="Evaluate only N queries for quick testing")
    args = parser.parse_args()

    result = evaluate(
        use_full_pipeline=args.full_pipeline,
        skip_reranker=args.skip_reranker,
        sample_n=args.sample,
    )
    if result:
        report_path = DATA_DIR / "evaluation_report.json"
        if args.full_pipeline and args.skip_reranker:
            report_path = DATA_DIR / "evaluation_report_hybrid_pipeline.json"
        elif args.full_pipeline:
            report_path = DATA_DIR / "evaluation_report_full.json"
        with open(report_path, "w") as f:
            json.dump(result, f, indent=2)
        logger.info(f"Report saved to {report_path}")
        print(json.dumps(result, indent=2))
    else:
        logger.error("Evaluation failed")
        sys.exit(1)