Team Ai
Apppublic

CamQuest/codesearch

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