Content hash: aa24cef28097458a805c35c34323c9dbfc9323d0f9f70dca876847fd989f7055
#!/usr/bin/env python3
"""Benchmark embedding models on a small custom eval set.
Evaluates recall@k for candidate models against a golden set of
(query, [relevant_doc_indices]) pairs. No API keys required for local models.
"""
from __future__ import annotations
import time
from dataclasses import dataclass
from typing import Callable
import numpy as np
@dataclass
class EvalSample:
query: str
relevant_ids: list[int] # Indices into corpus
# āā Mock embedder (replace with real sentence-transformers / API calls) ā
def embed_simple(texts: list[str], dim: int = 384) -> np.ndarray:
"""Deterministic mock embedding for reproducibility."""
embs = np.zeros((len(texts), dim))
for i, t in enumerate(texts):
rng = np.random.default_rng(abs(hash(t)) % (2**32))
embs[i] = rng.normal(size=dim)
embs[i] /= np.linalg.norm(embs[i])
return embs
# āā Retrieval āāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāā
def cosine_retrieve(query_emb: np.ndarray, corpus_embs: np.ndarray, top_k: int) -> list[int]:
scores = np.dot(corpus_embs, query_emb) # Already normalized
return list(np.argsort(scores)[::-1][:top_k])
# āā Metrics āāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāā
def recall_at_k(retrieved: list[int], relevant: list[int], k: int) -> float:
if not relevant:
return 1.0
return len(set(retrieved[:k]) & set(relevant)) / len(relevant)
def mrr(retrieved: list[int], relevant: list[int]) -> float:
for rank, idx in enumerate(retrieved, 1):
if idx in relevant:
return 1.0 / rank
return 0.0
# āā Benchmark runner āāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāā
def benchmark(
name: str,
embed_fn: Callable,
corpus: list[str],
eval_set: list[EvalSample],
dims: list[int],
k_values: list[int],
) -> dict:
results = {}
for d in dims:
t0 = time.perf_counter()
corpus_embs = embed_fn(corpus, dim=d)
embed_time = time.perf_counter() - t0
recalls = {k: [] for k in k_values}
mrrs = []
for sample in eval_set:
q_emb = embed_fn([sample.query], dim=d)[0]
retrieved = cosine_retrieve(q_emb, corpus_embs, top_k=max(k_values))
for k in k_values:
recalls[k].append(recall_at_k(retrieved, sample.relevant_ids, k))
mrrs.append(mrr(retrieved, sample.relevant_ids))
results[d] = {
"embed_time_s": round(embed_time, 4),
**{f"recall@{k}": round(np.mean(recalls[k]), 3) for k in k_values},
"mrr": round(np.mean(mrrs), 3),
}
print(f" {name} @ {d}d: recall@3={results[d]['recall@3']:.3f}, embed_time={embed_time:.4f}s")
return results
# āā Run āāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāāā
if __name__ == "__main__":
corpus = [
"Python is great for data science",
"JavaScript powers the web",
"PostgreSQL is a relational database",
"Docker containers isolate apps",
"Machine learning needs evaluation",
"Rust is for systems programming",
"Kubernetes orchestrates containers",
"FastAPI builds REST APIs in Python",
]
eval_set = [
EvalSample("Python web framework", [7]),
EvalSample("container tools", [3, 6]),
EvalSample("data analysis language", [0, 4]),
EvalSample("database systems", [2]),
]
print("Benchmarking candidate embedding dimensions...\n")
benchmark("mock-embeddings", embed_simple, corpus, eval_set, dims=[128, 384, 768], k_values=[1, 3, 5])
print("\nā Benchmark complete. Replace embed_simple with real sentence-transformers/API calls.")