"""Orchestrates the per-model, per-retrieval-mode, per-query-mode benchmark run.""" import statistics import time from . import config from .data import book_to_text, load_books, load_queries from .embedder import TimedEmbedder from .fusion import reciprocal_rank_fusion from .lexical import BM25Index from .metrics import evaluate_ranking from .vectorstore import BookVectorStore def _mean(values): values = [v for v in values if v is not None] return statistics.mean(values) if values else None def _retrieve(retrieval_mode, embedder, store, bm25_index, model_name, query_embedding, query_text, top_k, fusion_depth, rrf_k): """Returns (retrieved_book_ids, extra_retrieval_seconds) where extra excludes embedding time.""" start = time.perf_counter() if retrieval_mode == "dense": retrieved = store.search(model_name, query_embedding, top_k) elif retrieval_mode == "hybrid": bm25_ranked = bm25_index.search(query_text, fusion_depth) dense_ranked = store.search(model_name, query_embedding, fusion_depth) fused = reciprocal_rank_fusion([bm25_ranked, dense_ranked], k=rrf_k) retrieved = fused[:top_k] else: raise ValueError(f"Unknown retrieval_mode: {retrieval_mode}") elapsed = time.perf_counter() - start return retrieved, elapsed def _evaluate(embedder, store, bm25_index, model_name, queries, retrieval_mode, query_mode, top_k, fusion_depth, rrf_k): per_query = [] embed_latencies_sec = [] retrieval_latencies_sec = [] for q in queries: text = q.text_for_mode(query_mode) embedding, embed_latency = embedder.encode_query(text) retrieved, extra_latency = _retrieve( retrieval_mode, embedder, store, bm25_index, model_name, embedding, text, top_k, fusion_depth, rrf_k, ) total_latency = embed_latency + extra_latency embed_latencies_sec.append(embed_latency) retrieval_latencies_sec.append(total_latency) metrics = evaluate_ranking(retrieved, q.relevance) per_query.append( { "query": text, "categories": q.categories, "embedding_latency_ms": embed_latency * 1000, "retrieval_latency_ms": total_latency * 1000, **metrics, } ) per_category = {} for entry in per_query: for cat in entry["categories"]: bucket = per_category.setdefault( cat, {"n": 0, "recall@10": [], "recall@50": [], "mrr": [], "ndcg@10": []} ) bucket["n"] += 1 for metric in ("recall@10", "recall@50", "mrr", "ndcg@10"): bucket[metric].append(entry[metric]) per_category_summary = { cat: { "n": bucket["n"], "recall@10": _mean(bucket["recall@10"]), "recall@50": _mean(bucket["recall@50"]), "mrr": _mean(bucket["mrr"]), "ndcg@10": _mean(bucket["ndcg@10"]), } for cat, bucket in per_category.items() } mean_embed_latency_sec = _mean(embed_latencies_sec) mean_retrieval_latency_sec = _mean(retrieval_latencies_sec) return { "mean_embedding_latency_ms": mean_embed_latency_sec * 1000 if mean_embed_latency_sec else None, "mean_retrieval_latency_ms": mean_retrieval_latency_sec * 1000 if mean_retrieval_latency_sec else None, "queries_per_sec": (1.0 / mean_retrieval_latency_sec) if mean_retrieval_latency_sec else None, "recall@10": _mean(e["recall@10"] for e in per_query), "recall@50": _mean(e["recall@50"] for e in per_query), "mrr": _mean(e["mrr"] for e in per_query), "ndcg@10": _mean(e["ndcg@10"] for e in per_query), "per_query": per_query, "per_category": per_category_summary, } def run_benchmark( models=None, query_modes=None, retrieval_modes=None, top_k=None, fusion_depth=None, rrf_k=None, books_path=None, queries_path=None, qdrant_path=None, ): models = models if models is not None else config.MODELS query_modes = query_modes if query_modes is not None else config.QUERY_MODES retrieval_modes = retrieval_modes if retrieval_modes is not None else config.RETRIEVAL_MODES top_k = top_k or config.TOP_K fusion_depth = fusion_depth or config.FUSION_DEPTH rrf_k = rrf_k if rrf_k is not None else config.RRF_K books = load_books(books_path or config.BOOKS_PATH) queries = load_queries(queries_path or config.QUERIES_PATH) store = BookVectorStore(qdrant_path or config.QDRANT_PATH) documents = [book_to_text(b) for b in books] book_ids = [b["book_id"] for b in books] num_docs = len(books) bm25_index = BM25Index(documents, book_ids) if "hybrid" in retrieval_modes else None results = [] for model_cfg in models: model_name = model_cfg["name"] print(f"\n=== {model_name} ===") embedder = TimedEmbedder( model_name, query_prefix=model_cfg.get("query_prefix", ""), passage_prefix=model_cfg.get("passage_prefix", ""), trust_remote_code=model_cfg.get("trust_remote_code", False), query_encode_kwargs=model_cfg.get("query_encode_kwargs"), passage_encode_kwargs=model_cfg.get("passage_encode_kwargs"), ) print( f" loaded in {embedder.load_time_sec:.2f}s | " f"dim={embedder.embedding_dim} | size={embedder.model_size_mb:.1f} MB " f"| device={embedder.device}" ) doc_embeddings, doc_embed_time = embedder.encode_passages(documents) doc_throughput = num_docs / doc_embed_time if doc_embed_time else None print( f" embedded {num_docs} docs in {doc_embed_time:.2f}s " f"({doc_throughput:.1f} docs/sec)" ) store.index(model_name, doc_embeddings, books) raw_vector_bytes = embedder.embedding_dim * 4 * num_docs qdrant_disk_mb = store.collection_disk_size_mb(model_name) peak_mem_mb = embedder.peak_encode_memory_mb(queries[0].query) model_result = { "model": model_name, "device": embedder.device, "embedding_dim": embedder.embedding_dim, "param_count": embedder.param_count, "model_size_mb": embedder.model_size_mb, "load_time_sec": embedder.load_time_sec, "doc_embed_time_sec": doc_embed_time, "doc_throughput_docs_per_sec": doc_throughput, "raw_vector_storage_mb": raw_vector_bytes / (1024 ** 2), "qdrant_collection_disk_mb": qdrant_disk_mb, "peak_query_encode_memory_mb": peak_mem_mb, "retrieval_modes": {}, } for retrieval_mode in retrieval_modes: model_result["retrieval_modes"][retrieval_mode] = {} for query_mode in query_modes: print(f" evaluating retrieval_mode={retrieval_mode} query_mode={query_mode} ...") mode_result = _evaluate( embedder, store, bm25_index, model_name, queries, retrieval_mode, query_mode, top_k, fusion_depth, rrf_k, ) model_result["retrieval_modes"][retrieval_mode][query_mode] = mode_result print( f" recall@10={mode_result['recall@10']:.3f} " f"recall@50={mode_result['recall@50']:.3f} " f"mrr={mode_result['mrr']:.3f} " f"ndcg@10={mode_result['ndcg@10']:.3f} " f"retrieval_latency={mode_result['mean_retrieval_latency_ms']:.1f}ms" ) results.append(model_result) embedder.unload() return results