WalidAlHassan's picture
initial
eb02943
Raw History Blame Contribute Delete
7.8 kB
"""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