#!/usr/bin/env python3 """ Benchmark embedding models for Silo recommendations. Extracts media items from the dev database, embeds them with multiple models, and compares similarity quality using metadata as proxy ground truth. Usage: python3 scripts/benchmark-embeddings.py \ --db-host localhost \ --models lmstudio:text-embedding-qwen3-embedding-0.6b=http://localhost:1234 \ openai:text-embedding-3-large \ --openai-key sk-... \ --sample-size 500 """ import argparse import json import math import os import subprocess import sys import time from collections import defaultdict from dataclasses import dataclass, field from urllib.request import Request, urlopen from urllib.error import HTTPError # --------------------------------------------------------------------------- # Data extraction — single query with lateral join for credits # --------------------------------------------------------------------------- SAMPLE_SQL = """ SELECT json_agg(t)::text FROM ( SELECT mi.content_id, mi.title, mi.type, mi.year, mi.genres, mi.overview, mi.content_rating, mi.tagline, mi.studios, mi.networks, mi.countries, mi.keywords, mi.original_language, COALESCE(cr.credits, '[]'::json) AS credits FROM media_items mi LEFT JOIN LATERAL ( SELECT json_agg(json_build_object( 'name', p.name, 'kind', ip.kind, 'character', ip.character, 'sort_order', ip.sort_order ) ORDER BY ip.sort_order) AS credits FROM item_people ip JOIN people p ON p.id = ip.person_id WHERE ip.content_id = mi.content_id AND (ip.kind IN (2, 3) OR (ip.kind = 1 AND ip.sort_order <= 5)) ) cr ON true WHERE mi.status = 'matched' AND mi.overview IS NOT NULL AND mi.overview != '' AND array_length(mi.genres, 1) > 0 ORDER BY random() LIMIT {limit} ) t; """ def run_psql(db_host: str, sql: str) -> str: """Execute SQL via SSH + docker exec, piping via stdin to avoid escaping.""" ssh_cmd = "docker exec -i silo-postgres psql -U silo -d silo_fresh -At" cmd = ["ssh", f"root@{db_host}", ssh_cmd] result = subprocess.run(cmd, capture_output=True, text=True, timeout=120, input=sql) if result.returncode != 0: print(f"psql error: {result.stderr}", file=sys.stderr) sys.exit(1) return result.stdout.strip() def fetch_sample(db_host: str, sample_size: int) -> list[dict]: print(f"Fetching {sample_size} items with credits from {db_host}...") raw = run_psql(db_host, SAMPLE_SQL.format(limit=sample_size)) items = json.loads(raw) with_credits = sum(1 for it in items if it.get("credits") and len(it["credits"]) > 0) print(f" Got {len(items)} items ({with_credits} with cast/crew)") return items # --------------------------------------------------------------------------- # Text construction — mirrors BuildEmbeddingText() in # internal/recommendations/embeddings/text.go exactly # --------------------------------------------------------------------------- def build_embedding_text(item: dict) -> str: parts = [] genres = item.get("genres") or [] overview = item.get("overview") or "" type_name = "TV series" if item.get("type") == "series" else "movie" # Semantic lead: genres + type + overview (overview truncated to 1000 runes) if genres and overview: # Python slice on str counts code points, matching Go's rune semantics trunc = overview[:1000] parts.append(f"{', '.join(genres)} {type_name} about {trunc}") elif genres: parts.append(f"{', '.join(genres)} {type_name}") elif overview: parts.append(f"{type_name}. {overview[:1000]}") # Title with year year = item.get("year") or 0 if year > 0: parts.append(f"{item['title']} ({year})") else: parts.append(item["title"]) if item.get("content_rating"): parts.append(f"Rated {item['content_rating']}") if item.get("tagline"): parts.append(f'"{item["tagline"]}"') # Cast: up to 5 actors (kind=1) with character names credits = item.get("credits") or [] actors = [c for c in credits if c["kind"] == 1][:5] if actors: cast_parts = [] for a in actors: if a.get("character"): cast_parts.append(f"{a['name']} as {a['character']}") else: cast_parts.append(a["name"]) parts.append(f"Cast: {', '.join(cast_parts)}") # Directors (kind=2) directors = [c["name"] for c in credits if c["kind"] == 2] if directors: parts.append(f"Directed by {', '.join(directors)}") writers = [c["name"] for c in credits if c["kind"] == 3] if writers: parts.append(f"Written by {', '.join(writers)}") if item.get("keywords"): parts.append(f"Keywords: {', '.join(item['keywords'][:5])}") if item.get("original_language"): parts.append(f"Original language: {item['original_language']}") if item.get("studios"): parts.append(f"Studios: {', '.join(item['studios'])}") if item.get("networks"): parts.append(f"Network: {', '.join(item['networks'])}") countries = item.get("countries") or [] if countries: parts.append(f"Country: {', '.join(countries[:2])}") return ". ".join(parts) # --------------------------------------------------------------------------- # Embedding API clients # --------------------------------------------------------------------------- def embed_openai_compatible(texts: list[str], model: str, base_url: str, api_key: str = "", batch_size: int = 10) -> list[list[float]]: all_embeddings = [None] * len(texts) for i in range(0, len(texts), batch_size): batch = texts[i:i + batch_size] payload = json.dumps({"model": model, "input": batch}).encode() headers = {"Content-Type": "application/json"} if api_key: headers["Authorization"] = f"Bearer {api_key}" req = Request(f"{base_url}/v1/embeddings", data=payload, headers=headers) for attempt in range(5): try: with urlopen(req, timeout=120) as resp: data = json.loads(resp.read()) break except HTTPError as e: body = e.read().decode("utf-8", errors="replace")[:200] if e.code == 429 or e.code >= 500: wait = min(2 ** attempt, 30) print(f" Retry {attempt + 1} after {wait}s ({e.code})...") time.sleep(wait) elif e.code == 400: print(f"\n ERROR 400 from {base_url}: {body}", file=sys.stderr) print(f" Model '{model}' may not support embeddings in this server.", file=sys.stderr) return None else: raise except Exception as e: wait = min(2 ** attempt, 30) print(f" Retry {attempt + 1} after {wait}s ({e})...") time.sleep(wait) else: print(f" FAILED batch starting at index {i}", file=sys.stderr) return None for d in data["data"]: all_embeddings[i + d["index"]] = d["embedding"] done = min(i + batch_size, len(texts)) print(f" Embedded {done}/{len(texts)}", end="\r") print() return all_embeddings def embed_gemini(texts: list[str], model: str, api_key: str, batch_size: int = 10, task_type: str = "") -> list[list[float]]: all_embeddings = [] for i in range(0, len(texts), batch_size): batch = texts[i:i + batch_size] url = f"https://generativelanguage.googleapis.com/v1beta/models/{model}:batchEmbedContents?key={api_key}" requests_body = [] for t in batch: req = {"model": f"models/{model}", "content": {"parts": [{"text": t}]}} if task_type: req["taskType"] = task_type requests_body.append(req) payload = json.dumps({"requests": requests_body}).encode() headers = {"Content-Type": "application/json"} req = Request(url, data=payload, headers=headers) for attempt in range(5): try: with urlopen(req, timeout=120) as resp: data = json.loads(resp.read()) break except HTTPError as e: if e.code == 429 or e.code >= 500: wait = min(2 ** attempt, 30) print(f" Retry {attempt + 1} after {wait}s ({e.code})...") time.sleep(wait) else: raise else: print(f" FAILED batch at index {i}", file=sys.stderr) sys.exit(1) for emb in data["embeddings"]: all_embeddings.append(emb["values"]) done = min(i + batch_size, len(texts)) print(f" Embedded {done}/{len(texts)}", end="\r") print() return all_embeddings # --------------------------------------------------------------------------- # Similarity / evaluation # --------------------------------------------------------------------------- def cosine_sim(a: list[float], b: list[float]) -> float: dot = sum(x * y for x, y in zip(a, b)) na = math.sqrt(sum(x * x for x in a)) nb = math.sqrt(sum(x * x for x in b)) if na == 0 or nb == 0: return 0.0 return dot / (na * nb) def jaccard(a: set, b: set) -> float: if not a and not b: return 0.0 union = len(a | b) return len(a & b) / union if union else 0.0 @dataclass class ItemMeta: content_id: str title: str genres: set = field(default_factory=set) directors: set = field(default_factory=set) actors: set = field(default_factory=set) studios: set = field(default_factory=set) def metadata_similarity(a: ItemMeta, b: ItemMeta) -> dict: return { "genre_jaccard": jaccard(a.genres, b.genres), "shared_directors": len(a.directors & b.directors), "shared_actors": len(a.actors & b.actors), "shared_studios": len(a.studios & b.studios), "genre_overlap": len(a.genres & b.genres), } def metadata_relevance_score(meta_sim: dict) -> float: score = 0.0 score += meta_sim["genre_jaccard"] * 0.4 score += min(meta_sim["shared_directors"], 1) * 0.3 score += min(meta_sim["shared_actors"] / 2, 1) * 0.2 score += min(meta_sim["shared_studios"], 1) * 0.1 return score def evaluate_model(name: str, embeddings: list[list[float]], items_meta: list[ItemMeta], top_k: int = 10) -> dict: n = len(embeddings) # Use numpy if available for large matrices, else fall back to pure Python try: import numpy as np print(f" Computing {n}x{n} similarity matrix (numpy)...") emb_array = np.array(embeddings, dtype=np.float32) norms = np.linalg.norm(emb_array, axis=1, keepdims=True) norms[norms == 0] = 1.0 normed = emb_array / norms sim_matrix = (normed @ normed.T).tolist() except ImportError: print(f" Computing {n}x{n} similarity matrix (pure python — install numpy for speed)...") sim_matrix = [[0.0] * n for _ in range(n)] for i in range(n): for j in range(i + 1, n): s = cosine_sim(embeddings[i], embeddings[j]) sim_matrix[i][j] = s sim_matrix[j][i] = s genre_overlaps = [] director_hits = [] actor_hits = [] relevance_at_k = [] ndcg_scores = [] for i in range(n): neighbors = sorted(range(n), key=lambda j: sim_matrix[i][j], reverse=True) neighbors = [j for j in neighbors if j != i][:top_k] item_genre_overlaps = [] item_director_hits = 0 item_actor_hits = 0 item_relevance = [] for rank, j in enumerate(neighbors): ms = metadata_similarity(items_meta[i], items_meta[j]) item_genre_overlaps.append(ms["genre_overlap"]) item_director_hits += ms["shared_directors"] > 0 item_actor_hits += ms["shared_actors"] > 0 item_relevance.append(metadata_relevance_score(ms)) genre_overlaps.append(sum(item_genre_overlaps) / len(item_genre_overlaps) if item_genre_overlaps else 0) director_hits.append(item_director_hits) actor_hits.append(item_actor_hits) relevance_at_k.append(sum(item_relevance) / len(item_relevance) if item_relevance else 0) # NDCG: does the embedding ranking match the ideal metadata ranking? all_relevance = [] for j in range(n): if j == i: continue all_relevance.append(metadata_relevance_score(metadata_similarity(items_meta[i], items_meta[j]))) ideal = sorted(all_relevance, reverse=True)[:top_k] actual = item_relevance dcg = sum(r / math.log2(k + 2) for k, r in enumerate(actual)) idcg = sum(r / math.log2(k + 2) for k, r in enumerate(ideal)) ndcg_scores.append(dcg / idcg if idcg > 0 else 0.0) return { "model": name, "dimensions": len(embeddings[0]), "avg_genre_overlap_at_k": sum(genre_overlaps) / n, "avg_director_hits_at_k": sum(director_hits) / n, "avg_actor_hits_at_k": sum(actor_hits) / n, "mean_relevance_at_k": sum(relevance_at_k) / n, "mean_ndcg_at_k": sum(ndcg_scores) / n, } def print_topk_examples(name: str, embeddings: list[list[float]], items: list[dict], items_meta: list[ItemMeta], num_examples: int = 5, top_k: int = 5): n = len(embeddings) print(f"\n{'='*80}") print(f" Example recommendations: {name}") print(f"{'='*80}") indices = list(range(min(num_examples, n))) for i in indices: sims = [(j, cosine_sim(embeddings[i], embeddings[j])) for j in range(n) if j != i] sims.sort(key=lambda x: x[1], reverse=True) genres_str = ', '.join(items[i].get('genres') or []) print(f"\n {items[i]['title']} ({items[i].get('year', '?')}) [{genres_str}]") print(f" {'─' * 70}") for rank, (j, sim) in enumerate(sims[:top_k]): ms = metadata_similarity(items_meta[i], items_meta[j]) shared = [] if ms["genre_overlap"]: shared.append(f"{ms['genre_overlap']}g") if ms["shared_directors"]: shared.append(f"{ms['shared_directors']}d") if ms["shared_actors"]: shared.append(f"{ms['shared_actors']}a") tag = f" [{','.join(shared)}]" if shared else "" j_genres = ', '.join(items[j].get('genres') or []) print(f" {rank+1}. {items[j]['title']} ({items[j].get('year', '?')}) " f"[{j_genres}] sim={sim:.3f}{tag}") # --------------------------------------------------------------------------- # Main # --------------------------------------------------------------------------- def parse_model_spec(spec: str) -> dict: """Parse model spec. Supports optional +task_type suffix for Gemini. Examples: gemini:gemini-embedding-2-preview+SEMANTIC_SIMILARITY openai:text-embedding-3-large lmstudio:model-name=http://host:port """ if ":" not in spec: raise ValueError(f"Invalid model spec: {spec}. Format: provider:model-name[+task_type][=base_url]") provider, rest = spec.split(":", 1) task_type = "" if "=" in rest: model_part, base_url = rest.split("=", 1) else: model_part = rest base_url = None if "+" in model_part: model, task_type = model_part.split("+", 1) else: model = model_part if provider == "openai" and not base_url: base_url = "https://api.openai.com" return {"provider": provider, "model": model, "base_url": base_url, "task_type": task_type} def main(): parser = argparse.ArgumentParser(description="Benchmark embedding models for Silo") parser.add_argument("--db-host", default="localhost", help="Dev server SSH host") parser.add_argument("--models", nargs="+", required=True, help="Model specs: provider:model[=base_url]. " "Providers: lmstudio, openai, gemini") parser.add_argument("--openai-key", default=os.environ.get("OPENAI_API_KEY", "")) parser.add_argument("--gemini-key", default=os.environ.get("GEMINI_API_KEY", "")) parser.add_argument("--sample-size", type=int, default=500) parser.add_argument("--top-k", type=int, default=10, help="Number of neighbors to evaluate") parser.add_argument("--examples", type=int, default=5, help="Number of example items to show") parser.add_argument("--batch-size", type=int, default=20) parser.add_argument("--cache-dir", default="/tmp/silo-embed-bench", help="Cache embeddings to avoid re-computing") parser.add_argument("--no-cache", action="store_true", help="Ignore cached data") args = parser.parse_args() os.makedirs(args.cache_dir, exist_ok=True) # 1. Fetch data cache_file = os.path.join(args.cache_dir, f"sample_{args.sample_size}.json") if os.path.exists(cache_file) and not args.no_cache: print(f"Loading cached sample from {cache_file}") with open(cache_file) as f: items = json.load(f) else: items = fetch_sample(args.db_host, args.sample_size) with open(cache_file, "w") as f: json.dump(items, f) # 2. Build embedding texts (identical to Go's BuildEmbeddingText) texts = [build_embedding_text(item) for item in items] avg_len = sum(len(t) for t in texts) / len(texts) print(f"\nEmbedding texts: {len(texts)} items, avg {avg_len:.0f} chars") # Spot-check: show a sample text print(f"\n Sample text ({items[0]['title']}):") print(f" {texts[0][:200]}...") # 3. Build metadata for ground truth items_meta = [] for item in items: credits = item.get("credits") or [] m = ItemMeta( content_id=item["content_id"], title=item["title"], genres=set(item.get("genres") or []), directors={c["name"] for c in credits if c["kind"] == 2}, actors={c["name"] for c in credits if c["kind"] == 1}, studios=set(item.get("studios") or []), ) items_meta.append(m) has_directors = sum(1 for m in items_meta if m.directors) has_actors = sum(1 for m in items_meta if m.actors) print(f" Metadata: {has_actors} items with actors, {has_directors} with directors") # 4. Embed with each model model_specs = [parse_model_spec(s) for s in args.models] results = {} for spec in model_specs: task_suffix = f"+{spec['task_type']}" if spec.get('task_type') else "" label = f"{spec['provider']}:{spec['model']}{task_suffix}" cache_key = f"{spec['model']}{task_suffix}".replace("/", "_") emb_cache = os.path.join(args.cache_dir, f"embeddings_{cache_key}_{args.sample_size}.json") if os.path.exists(emb_cache) and not args.no_cache: print(f"\n[{label}] Loading cached embeddings...") with open(emb_cache) as f: embeddings = json.load(f) else: print(f"\n[{label}] Embedding {len(texts)} items...") t0 = time.time() if spec["provider"] in ("lmstudio", "openai", "ollama"): api_key = args.openai_key if spec["provider"] == "openai" else "" embeddings = embed_openai_compatible( texts, spec["model"], spec["base_url"], api_key=api_key, batch_size=args.batch_size, ) elif spec["provider"] == "gemini": embeddings = embed_gemini( texts, spec["model"], args.gemini_key, batch_size=args.batch_size, task_type=spec.get("task_type", ""), ) else: print(f"Unknown provider: {spec['provider']}", file=sys.stderr) sys.exit(1) if embeddings is None: print(f" SKIPPING {label} (embedding failed)") continue elapsed = time.time() - t0 print(f" Done in {elapsed:.1f}s ({len(texts)/elapsed:.1f} items/sec, {len(embeddings[0])} dims)") with open(emb_cache, "w") as f: json.dump(embeddings, f) # 5. Evaluate print(f"\n[{label}] Evaluating top-{args.top_k} quality...") result = evaluate_model(label, embeddings, items_meta, top_k=args.top_k) results[label] = result print_topk_examples(label, embeddings, items, items_meta, num_examples=args.examples, top_k=5) # 6. Summary print(f"\n{'='*80}") print(f" BENCHMARK RESULTS (top-{args.top_k} neighbors, {len(items)} items)") print(f"{'='*80}") header = f"{'Model':<45} {'Dims':>5} {'NDCG@K':>7} {'Rel@K':>7} {'Genre':>6} {'Dir':>5} {'Actor':>5}" print(f"\n{header}") print("─" * len(header)) sorted_results = sorted(results.values(), key=lambda r: r["mean_ndcg_at_k"], reverse=True) for r in sorted_results: print(f"{r['model']:<45} {r['dimensions']:>5} " f"{r['mean_ndcg_at_k']:>7.3f} " f"{r['mean_relevance_at_k']:>7.3f} " f"{r['avg_genre_overlap_at_k']:>6.2f} " f"{r['avg_director_hits_at_k']:>5.2f} " f"{r['avg_actor_hits_at_k']:>5.2f}") best = sorted_results[0] print(f"\nBest model by NDCG@{args.top_k}: {best['model']}") if len(sorted_results) >= 2: print(f"\nPairwise deltas vs {best['model']}:") for r in sorted_results[1:]: ndcg_diff = best["mean_ndcg_at_k"] - r["mean_ndcg_at_k"] rel_diff = best["mean_relevance_at_k"] - r["mean_relevance_at_k"] print(f" vs {r['model']}: NDCG {'+' if ndcg_diff >= 0 else ''}{ndcg_diff:.3f}, " f"Rel {'+' if rel_diff >= 0 else ''}{rel_diff:.3f}") if __name__ == "__main__": main()