CamQuest/codesearch
0
1"""2HNSW ef_search sweep — M2 concept checkpoint.3 4Reads pre-cached query vectors from disk (run scripts/cache_query_vectors.py5first) and sweeps the HNSW beam width to chart the recall/latency tradeoff.6 7For each ef we report:8 Recall@100, MRR@10 — accuracy9 server ms/q — Qdrant's reported time per query (from response.time)10 wall ms/q — client-side wall time per query (server + network RTT)11 12The delta between server and wall ms/q is how much of "ms per query" is13actually network round-trip to Qdrant Cloud, not HNSW work.14 15Usage:16 uv run python scripts/cache_query_vectors.py # one-time17 uv run python scripts/ef_sweep.py18 uv run python scripts/ef_sweep.py --ef-values 16,32,64,128,256,51219 uv run python scripts/ef_sweep.py --max-queries 100020"""21 22from __future__ import annotations23 24import argparse25import os26import pickle27import random28import sys29import time30 31sys.path.insert(0, os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "src"))32 33import httpx34from tqdm import tqdm35 36from codesearch.config import (37 EMBEDDING_MODEL,38 QDRANT_API_KEY,39 QDRANT_COLLECTION,40 QDRANT_URL,41)42from codesearch.eval.metrics import mean_recall, mean_reciprocal_rank43 44_SAMPLE_SEED = 42 # match eval/harness.py45_SEARCH_BATCH = 50 # queries per HTTP request to Qdrant (matches dense.py)46_TOP_K = 10047CACHE_DIR = ".cache"48 49 50def cache_path(model_name: str) -> str:51 safe = model_name.replace("/", "_")52 return os.path.join(CACHE_DIR, f"query_vectors_{safe}.pkl")53 54 55def load_cached_queries() -> tuple[list[dict], list[list[float]]]:56 path = cache_path(EMBEDDING_MODEL)57 if not os.path.exists(path):58 sys.exit(59 f"[error] Cache not found at {path}.\n"60 f" Run: uv run python scripts/cache_query_vectors.py"61 )62 print(f"[1/3] Loading cached query vectors from {path}...")63 with open(path, "rb") as f:64 data = pickle.load(f)65 if data["model"] != EMBEDDING_MODEL:66 sys.exit(67 f"[error] Cache model mismatch:\n"68 f" cache built for: {data['model']!r}\n"69 f" current model: {EMBEDDING_MODEL!r}\n"70 f" Run: uv run python scripts/cache_query_vectors.py --recompute"71 )72 queries = data["queries"]73 vectors = data["vectors"].tolist()74 print(f" Loaded {len(queries):,} queries and their vectors.")75 return queries, vectors76 77 78def search_batch(79 http: httpx.Client,80 chunk: list[list[float]],81 ef: int,82) -> tuple[list[list[str]], float]:83 """One HTTP call to Qdrant for a chunk of <=_SEARCH_BATCH queries.84 Returns (hits_per_query, server_time_seconds)."""85 payload = {86 "searches": [87 {88 "query": qv,89 "limit": _TOP_K,90 "params": {"hnsw_ef": ef},91 "with_payload": ["doc_id"],92 }93 for qv in chunk94 ]95 }96 r = http.post(97 f"/collections/{QDRANT_COLLECTION}/points/query/batch",98 json=payload,99 timeout=60.0,100 )101 r.raise_for_status()102 body = r.json()103 server_time = float(body.get("time", 0.0))104 hits = [105 [p["payload"]["doc_id"] for p in resp["points"]]106 for resp in body["result"]107 ]108 return hits, server_time109 110 111def sweep_one_ef(112 http: httpx.Client,113 query_vectors: list[list[float]],114 ef: int,115) -> tuple[list[list[str]], float, float]:116 """Run all batches at a given ef. Returns (hits, wall_seconds, server_seconds)."""117 all_hits: list[list[str]] = []118 server_total = 0.0119 n_batches = (len(query_vectors) + _SEARCH_BATCH - 1) // _SEARCH_BATCH120 121 wall_t0 = time.perf_counter()122 for i in tqdm(123 range(0, len(query_vectors), _SEARCH_BATCH),124 total=n_batches,125 desc=f" ef={ef:<4}",126 leave=False,127 ):128 chunk = query_vectors[i : i + _SEARCH_BATCH]129 hits, server_t = search_batch(http, chunk, ef)130 all_hits.extend(hits)131 server_total += server_t132 wall_elapsed = time.perf_counter() - wall_t0133 return all_hits, wall_elapsed, server_total134 135 136def run_sweep(ef_values: list[int], max_queries: int) -> None:137 # ------------------------------------------------------------------138 # [1] Load cached queries + vectors139 # ------------------------------------------------------------------140 queries, vectors = load_cached_queries()141 142 if max_queries and len(queries) > max_queries:143 random.seed(_SAMPLE_SEED)144 idx = random.sample(range(len(queries)), max_queries)145 queries = [queries[i] for i in idx]146 vectors = [vectors[i] for i in idx]147 print(f" Sampled {len(queries):,} queries (seed={_SAMPLE_SEED}).")148 149 # ------------------------------------------------------------------150 # [2] Verify Qdrant collection is populated151 # ------------------------------------------------------------------152 print(f"[2/3] Connecting to Qdrant at {QDRANT_URL}")153 http = httpx.Client(154 base_url=QDRANT_URL,155 headers={"api-key": QDRANT_API_KEY},156 )157 info = http.get(f"/collections/{QDRANT_COLLECTION}").json()158 points_count = info.get("result", {}).get("points_count", 0)159 if not points_count:160 sys.exit(161 f"[error] Collection '{QDRANT_COLLECTION}' is empty.\n"162 f" Run: uv run python scripts/index_corpus.py"163 )164 print(f" Collection '{QDRANT_COLLECTION}' has {points_count:,} vectors.")165 166 # ------------------------------------------------------------------167 # [3] Warmup, then sweep168 # ------------------------------------------------------------------169 print(f"[3/3] Warming up (1 batch at ef={ef_values[0]})...")170 search_batch(http, vectors[: min(_SEARCH_BATCH, len(vectors))], ef_values[0])171 172 print(f" Sweeping ef across {ef_values}...")173 relevant_ids = [q["relevant_id"] for q in queries]174 rows: list[tuple[int, float, float, float, float]] = []175 176 for ef in tqdm(ef_values, desc="Sweep"):177 hits, wall_s, server_s = sweep_one_ef(http, vectors, ef)178 query_results = list(zip(hits, relevant_ids))179 mrr = mean_reciprocal_rank(query_results, k=10)180 recall = mean_recall(query_results, k=100)181 n = len(queries)182 rows.append((ef, mrr, recall, server_s / n * 1000, wall_s / n * 1000))183 184 http.close()185 186 # ------------------------------------------------------------------187 # Results table188 # ------------------------------------------------------------------189 print()190 print(f"Dense retrieval HNSW ef_search sweep — n={len(queries)} queries, top-{_TOP_K}")191 header = f"{'ef':>6} {'MRR@10':>9} {'Recall@100':>12} {'server ms/q':>14} {'wall ms/q':>12}"192 print(header)193 print("-" * len(header))194 for ef, mrr, recall, server_ms, wall_ms in rows:195 print(f"{ef:>6} {mrr:>9.4f} {recall:>12.4f} {server_ms:>14.2f} {wall_ms:>12.2f}")196 print()197 print("server ms/q = Qdrant's reported processing time per query (HNSW work only)")198 print("wall ms/q = client-side wall time per query (server time + network RTT)")199 print("(query embedding is cached and excluded from both)")200 201 202def main() -> None:203 parser = argparse.ArgumentParser(204 description="Sweep HNSW ef_search on the dense Qdrant index."205 )206 parser.add_argument(207 "--ef-values",208 default="32,128,512",209 help="Comma-separated ef_search values to sweep (default: 32,128,512 per PLAN.md).",210 )211 parser.add_argument(212 "--max-queries",213 type=int,214 default=500,215 help="Sample size (default: 500 — ~±2pp stderr on Recall@100, fast runtime).",216 )217 args = parser.parse_args()218 219 ef_values = [int(x) for x in args.ef_values.split(",")]220 run_sweep(ef_values, args.max_queries)221 222 223if __name__ == "__main__":224 main()225 