benchmark_embeddings.py

script

← Back to skill

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.")