Team Ai
Apppublic

CamQuest/codesearch

sourceHugging Facemitupdated 3mo agoView on Hugging Face
0likes
cache_query_vectors.py92 linesDownload Raw Back to scripts
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