CamQuest/codesearch
0
1"""2Measure Recall@K of candidate pools — answers "what's the reranker's ceiling?"3 4For each query, the cross-encoder can only re-rank documents that the upstream5retriever put in the pool. So Recall@K of the pool is the upper bound on what6any reranker can achieve at top-K. This script reports that ceiling for four7pool strategies, at three pool sizes:8 9 - BM25 top-K10 - Dense top-K11 - RRF(BM25 top-K, Dense top-K) top-K (k=60, missing rank = K+1)12 - Set union (BM25 top-K ∪ Dense top-K) — absolute ceiling for any fusion13 14The gap between RRF and set-union tells how much ranking quality matters above15pool composition. The gap between RRF and the better of (BM25, Dense) tells16whether fusion is pulling its weight at that K.17 18Usage:19 uv run python scripts/cache_query_vectors.py # one-time, if not done20 uv run python scripts/recall_at_pool.py21 uv run python scripts/recall_at_pool.py --max-queries -1 # full corpus22"""23 24from __future__ import annotations25 26import argparse27import os28import pickle29import random30import sys31 32sys.path.insert(0, os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "src"))33 34import httpx35from tqdm import tqdm36 37from codesearch.config import (38 EMBEDDING_MODEL,39 QDRANT_API_KEY,40 QDRANT_COLLECTION,41 QDRANT_URL,42)43from codesearch.data import load_codesearch44from codesearch.retrievers.bm25 import BM25Retriever45from codesearch.retrievers.bm25_index import BM25Index46from codesearch.retrievers.hybrid import rrf_fuse47 48_SAMPLE_SEED = 4249_MAX_K = 10050_POOL_SIZES = (20, 50, 100)51_SEARCH_BATCH = 5052CACHE_DIR = ".cache"53BM25_CACHE_DIR = ".cache/bm25"54 55 56def cache_path(model_name: str) -> str:57 safe = model_name.replace("/", "_")58 return os.path.join(CACHE_DIR, f"query_vectors_{safe}.pkl")59 60 61def load_cached() -> tuple[list[dict], list[list[float]]]:62 path = cache_path(EMBEDDING_MODEL)63 if not os.path.exists(path):64 sys.exit(65 f"[error] Cache not found at {path}.\n"66 f" Run: uv run python scripts/cache_query_vectors.py"67 )68 with open(path, "rb") as f:69 data = pickle.load(f)70 if data["model"] != EMBEDDING_MODEL:71 sys.exit(72 f"[error] Cache model mismatch. Re-run scripts/cache_query_vectors.py --recompute"73 )74 return data["queries"], data["vectors"].tolist()75 76 77def dense_search_batch(http: httpx.Client, vectors: list[list[float]], ef: int = 128) -> list[list[str]]:78 hits_all: list[list[str]] = []79 n_batches = (len(vectors) + _SEARCH_BATCH - 1) // _SEARCH_BATCH80 for i in tqdm(range(0, len(vectors), _SEARCH_BATCH), total=n_batches, desc="Dense search"):81 chunk = vectors[i : i + _SEARCH_BATCH]82 payload = {83 "searches": [84 {85 "query": qv,86 "limit": _MAX_K,87 "params": {"hnsw_ef": ef},88 "with_payload": ["doc_id"],89 }90 for qv in chunk91 ]92 }93 r = http.post(94 f"/collections/{QDRANT_COLLECTION}/points/query/batch",95 json=payload,96 timeout=60.0,97 )98 r.raise_for_status()99 for resp in r.json()["result"]:100 hits_all.append([p["payload"]["doc_id"] for p in resp["points"]])101 return hits_all102 103 104def _id_dicts(ids: list[str]) -> list[dict]:105 """Wrap doc-ids as minimal hit-dicts for rrf_fuse (which takes dicts)."""106 return [{"id": i} for i in ids]107 108 109def recall(hit_lists: list[list[str]], relevant: list[str], k: int) -> float:110 """Fraction of queries whose relevant doc appears in the top-k of its hit list."""111 n = len(hit_lists)112 if n == 0:113 return 0.0114 return sum(1 for hits, rel in zip(hit_lists, relevant) if rel in hits[:k]) / n115 116 117def main() -> None:118 parser = argparse.ArgumentParser(119 description="Recall@K of candidate pools (BM25, Dense, RRF, set-union ceiling)."120 )121 parser.add_argument(122 "--max-queries",123 type=int,124 default=2000,125 help="Sample size (default: 2000; -1 = full eval set).",126 )127 args = parser.parse_args()128 129 # [1/4] Cached queries + vectors130 print("[1/4] Loading cached query vectors...")131 queries, vectors = load_cached()132 print(f" Loaded {len(queries):,} queries.")133 134 if args.max_queries and args.max_queries > 0 and len(queries) > args.max_queries:135 random.seed(_SAMPLE_SEED)136 idx = random.sample(range(len(queries)), args.max_queries)137 queries = [queries[i] for i in idx]138 vectors = [vectors[i] for i in idx]139 print(f" Sampled {len(queries):,} queries (seed={_SAMPLE_SEED}).")140 141 relevant = [q["relevant_id"] for q in queries]142 143 # [2/4] BM25 — from cache if available, else build from scratch144 if BM25Index.exists(BM25_CACHE_DIR):145 print(f"[2/4] Loading cached BM25 index from {BM25_CACHE_DIR}...")146 bm25 = BM25Retriever.from_cache(BM25_CACHE_DIR)147 print(f" Loaded BM25 over {len(bm25.corpus):,} docs.")148 else:149 print("[2/4] No BM25 cache found — building from scratch (~2-3 min).")150 print(f" Tip: run scripts/cache_bm25.py to skip this next time.")151 corpus, _ = load_codesearch(n=-1)152 bm25 = BM25Retriever(corpus)153 154 # [3/4] BM25 search155 print(f"[3/4] Running BM25 on {len(queries):,} queries (top-{_MAX_K})...")156 bm25_results = bm25.retrieve_batch([q["query"] for q in queries], top_k=_MAX_K)157 bm25_hits = [[h["id"] for h in row] for row in bm25_results]158 159 # [4/4] Dense search via Qdrant REST160 print(f"[4/4] Running dense (Qdrant) on {len(queries):,} queries (top-{_MAX_K})...")161 http = httpx.Client(base_url=QDRANT_URL, headers={"api-key": QDRANT_API_KEY})162 dense_hits = dense_search_batch(http, vectors)163 http.close()164 165 # Sanity check166 assert len(bm25_hits) == len(dense_hits) == len(relevant)167 168 # Compute recall at each pool size169 print()170 print(f"Pool Recall@K (n={len(queries):,} queries, seed={_SAMPLE_SEED})")171 print(172 f"{'K':>5} {'BM25':>8} {'Dense':>8} {'RRF':>8} {'Union':>8}"173 f" {'RRF gain':>9} {'ceiling gap':>11}"174 )175 print(" " + "─" * 70)176 for k in _POOL_SIZES:177 r_bm25 = recall(bm25_hits, relevant, k)178 r_dense = recall(dense_hits, relevant, k)179 180 rrf_hits = [181 [h["id"] for h in rrf_fuse(_id_dicts(b_ids[:k]), _id_dicts(d_ids[:k]), top_k=k)]182 for b_ids, d_ids in zip(bm25_hits, dense_hits)183 ]184 r_rrf = recall(rrf_hits, relevant, k)185 186 # Set-union ceiling: GT is in BM25 top-k OR Dense top-k187 r_union = sum(188 1189 for b_ids, d_ids, rel in zip(bm25_hits, dense_hits, relevant)190 if rel in b_ids[:k] or rel in d_ids[:k]191 ) / len(queries)192 193 better_single = max(r_bm25, r_dense)194 rrf_gain = r_rrf - better_single # vs. best single retriever195 ceiling_gap = r_union - r_rrf # what RRF leaves on the table196 197 print(198 f"{k:>5} {r_bm25:>8.4f} {r_dense:>8.4f} {r_rrf:>8.4f} {r_union:>8.4f}"199 f" {rrf_gain:>+9.4f} {ceiling_gap:>+11.4f}"200 )201 202 print()203 print("Read:")204 print(" - 'RRF gain' = Recall(RRF@K) − Recall(best of BM25/Dense @K).")205 print(" Positive → fusion finds GT the better single list missed.")206 print(" - 'ceiling gap' = Recall(set-union@K) − Recall(RRF@K).")207 print(" Positive → GT is in the union but RRF didn't surface it.")208 print(" The reranker can recover this gap if the GT is anywhere")209 print(" in the candidate set it sees.")210 211 212if __name__ == "__main__":213 main()214 