Team Ai
Apppublic

CamQuest/codesearch

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