Team Ai
Apppublic

CamQuest/codesearch

sourceHugging Facemitupdated 3mo agoView on Hugging Face
0likes
sample_eval_partial.py140 linesDownload Raw Back to scripts
1"""M5 sub-experiment B — early directional read from a PARTIAL golds-first index.2 3The full UniXcoder index runs golds-first (index_corpus.py GOLDS_FIRST=1), so the4first ~22k embedded docs are every eval query's gold target; docs after that are5train-split distractors whose pool grows as indexing proceeds. This script does a6fair, exact brute-force cosine A/B over the SAME candidate pool for both models,7reading vectors straight from the .cache/*.npy files — no Qdrant, no HNSW recall8loss, no CPU contention with the running indexer.9 10Pool at run time = the first P = unixcoder embed-progress docs (golds-first order).11UniXcoder doc vectors come from its cache[0:P]; MiniLM doc vectors are pulled from12the full MiniLM cache for the SAME doc set (mapped by doc_id). Queries evaluated =13those whose gold is within the pool. Metrics: MRR@10 / nDCG@10 / Recall@100.14 15Usage:16    .venv/bin/python scripts/sample_eval_partial.py            # eval at current P17    .venv/bin/python scripts/sample_eval_partial.py --min-docs 44000   # require P first18"""19from __future__ import annotations20 21import argparse22import os23import sys24 25import numpy as np26 27sys.path.insert(0, "src")28from codesearch.data import load_codesearch  # noqa: E40229from codesearch.eval.metrics import (  # noqa: E40230    mean_ndcg,31    mean_recall,32    mean_reciprocal_rank,33)34from codesearch.embedding import load_encoder  # noqa: E40235 36CACHE = ".cache"37UNIX_MODEL = "microsoft/unixcoder-base"38MINILM_MODEL = "all-MiniLM-L6-v2"39UNIX_NPY = f"{CACHE}/corpus_vectors_microsoft_unixcoder-base.npy"40UNIX_PROG = f"{CACHE}/corpus_vectors_microsoft_unixcoder-base.progress"41MINILM_NPY = f"{CACHE}/corpus_vectors_all-MiniLM-L6-v2.npy"42UNIX_DIM, MINILM_DIM = 768, 38443 44 45def _read_progress(path: str) -> int:46    with open(path) as f:47        return int(f.read().strip() or "0")48 49 50def _golds_first_corpus(corpus, queries):51    """Reproduce index_corpus.py's GOLDS_FIRST reorder exactly."""52    gold_ids = {q["relevant_id"] for q in queries}53    golds = [d for d in corpus if d["id"] in gold_ids]54    rest = [d for d in corpus if d["id"] not in gold_ids]55    return golds + rest56 57 58def _topk_ids(query_vecs, doc_vecs, doc_ids, k=100, chunk=512):59    """Brute-force cosine (vectors are pre-normalized): top-k doc_ids per query."""60    out = []61    for i in range(0, len(query_vecs), chunk):62        sims = query_vecs[i : i + chunk] @ doc_vecs.T           # (b, P)63        kk = min(k, sims.shape[1])64        part = np.argpartition(-sims, kk - 1, axis=1)[:, :kk]65        for r in range(sims.shape[0]):66            order = part[r][np.argsort(-sims[r, part[r]])]67            out.append([doc_ids[j] for j in order])68    return out69 70 71def _eval_model(name, npy_path, dim, doc_rows, query_texts, relevant_ids, encoder):72    """doc_rows: np.ndarray of cache-row indices for the pool docs (in pool order)."""73    mm = np.memmap(npy_path, dtype=np.float32, mode="r").reshape(-1, dim)74    doc_vecs = np.ascontiguousarray(mm[doc_rows])                # (P, dim)75    print(f"[{name}] encoding {len(query_texts):,} queries...", flush=True)76    qv = encoder.encode(query_texts, normalize_embeddings=True, batch_size=128)77    qv = np.asarray(qv, dtype=np.float32)78    retrieved = _topk_ids(qv, doc_vecs, POOL_IDS, k=100)79    qr = list(zip(retrieved, relevant_ids))80    return (81        mean_reciprocal_rank(qr, k=10),82        mean_ndcg(qr, k=10),83        mean_recall(qr, k=100),84    )85 86 87def main() -> int:88    ap = argparse.ArgumentParser()89    ap.add_argument("--min-docs", type=int, default=0,90                    help="Require at least this many embedded docs before evaluating.")91    args = ap.parse_args()92 93    P = _read_progress(UNIX_PROG)94    print(f"UniXcoder embed-progress P = {P:,} docs")95    if P < args.min_docs:96        print(f"Not enough yet (need {args.min_docs:,}); exiting.")97        return 098 99    corpus, queries = load_codesearch(n=-1)100    orig_row = {d["id"]: i for i, d in enumerate(corpus)}        # doc_id -> MiniLM cache row101    gf = _golds_first_corpus(corpus, queries)                    # golds-first (== unixcoder cache order)102    pool = gf[:P]103    global POOL_IDS104    POOL_IDS = [d["id"] for d in pool]105    pool_id_set = set(POOL_IDS)106 107    # rows into each cache for the pool docs, in pool order108    unix_rows = np.arange(P)                                     # unixcoder cache is already golds-first109    minilm_rows = np.array([orig_row[i] for i in POOL_IDS])      # map to original-order MiniLM cache110 111    subset = [q for q in queries if q["relevant_id"] in pool_id_set]112    texts = [q["query"] for q in subset]113    rel = [q["relevant_id"] for q in subset]114    print(f"Pool: {P:,} docs  |  eval queries (gold in pool): {len(subset):,}")115    if not subset:116        print("No queries have their gold in the pool yet — wait for more golds.")117        return 0118 119    unix_enc = load_encoder(UNIX_MODEL)120    minilm_enc = load_encoder(MINILM_MODEL)121 122    u_mrr, u_ndcg, u_rec = _eval_model("unixcoder", UNIX_NPY, UNIX_DIM, unix_rows, texts, rel, unix_enc)123    m_mrr, m_ndcg, m_rec = _eval_model("minilm", MINILM_NPY, MINILM_DIM, minilm_rows, texts, rel, minilm_enc)124 125    print("\n" + "=" * 62)126    print(f"PARTIAL-INDEX A/B  (pool={P:,} docs, n={len(subset):,} queries, brute-force cosine)")127    print("NOTE: reduced pool inflates absolute scores; the MODEL DELTA is the signal.")128    print("=" * 62)129    print(f"{'Model':<16}{'MRR@10':>10}{'nDCG@10':>10}{'Recall@100':>12}")130    print("-" * 62)131    print(f"{'MiniLM (base)':<16}{m_mrr:>10.4f}{m_ndcg:>10.4f}{m_rec:>12.4f}")132    print(f"{'UniXcoder':<16}{u_mrr:>10.4f}{u_ndcg:>10.4f}{u_rec:>12.4f}")133    print(f"{'Δ (Unix-Mini)':<16}{u_mrr-m_mrr:>+10.4f}{u_ndcg-m_ndcg:>+10.4f}{u_rec-m_rec:>+12.4f}")134    print("=" * 62)135    return 0136 137 138if __name__ == "__main__":139    sys.exit(main())140