CamQuest/codesearch
0
1"""2Pre-compute query embeddings and cache them to disk.3 4The ef_sweep script (and any future queries-only experiment) reads from this5cache to skip the embedding step. The cache key includes the embedding model6name, so changing EMBEDDING_MODEL invalidates this cache automatically and7forces a recompute.8 9Usage:10 uv run python scripts/cache_query_vectors.py11 uv run python scripts/cache_query_vectors.py --recompute12"""13 14from __future__ import annotations15 16import argparse17import os18import pickle19import sys20import time21 22sys.path.insert(0, os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "src"))23 24import numpy as np25from sentence_transformers import SentenceTransformer26 27from codesearch.config import EMBEDDING_MODEL28from codesearch.data import load_codesearch29 30CACHE_DIR = ".cache"31 32 33def cache_path(model_name: str) -> str:34 safe = model_name.replace("/", "_")35 return os.path.join(CACHE_DIR, f"query_vectors_{safe}.pkl")36 37 38def main() -> None:39 parser = argparse.ArgumentParser(40 description="Cache query embeddings so the ef sweep can skip the encode step."41 )42 parser.add_argument(43 "--recompute",44 action="store_true",45 help="Re-embed even if the cache already exists.",46 )47 args = parser.parse_args()48 49 path = cache_path(EMBEDDING_MODEL)50 if os.path.exists(path) and not args.recompute:51 size_mb = os.path.getsize(path) / 1024 / 102452 print(f"[skip] Cache already exists at {path} ({size_mb:.1f} MB).")53 print(" Pass --recompute to rebuild.")54 return55 56 os.makedirs(CACHE_DIR, exist_ok=True)57 58 print("[1/3] Loading eval queries (test split only)...")59 _, queries = load_codesearch(n=-1, queries_only=True)60 61 print(f"[2/3] Loading embedding model: {EMBEDDING_MODEL}")62 model = SentenceTransformer(EMBEDDING_MODEL)63 64 print(f"[3/3] Embedding {len(queries):,} queries...")65 t0 = time.perf_counter()66 vectors = model.encode(67 [q["query"] for q in queries],68 normalize_embeddings=True,69 show_progress_bar=True,70 batch_size=256,71 )72 elapsed = time.perf_counter() - t073 print(f" Done in {elapsed:.1f}s ({elapsed / len(queries) * 1000:.2f} ms/query)")74 75 print(f"Writing cache to {path}...")76 with open(path, "wb") as f:77 pickle.dump(78 {79 "model": EMBEDDING_MODEL,80 "queries": queries,81 "vectors": np.asarray(vectors, dtype=np.float32),82 },83 f,84 protocol=pickle.HIGHEST_PROTOCOL,85 )86 size_mb = os.path.getsize(path) / 1024 / 102487 print(f"Cached {len(queries):,} query vectors ({size_mb:.1f} MB).")88 89 90if __name__ == "__main__":91 main()92 