Spaces:
Running
Running
| 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) | |