Spaces:
Running
Running
| """ | |
| Offline corpus indexing β embed + upsert in parallel with a local vector cache. | |
| Pipeline: | |
| βββΊ upsert worker 1 β | |
| embed thread βββΊ bounded queue ββββββββΌββΊ upsert worker 2 βββΊ Qdrant | |
| (writes .npy memmap in-line) βββΊ upsert worker 3 β | |
| The embed thread is CPU-bound (sentence-transformers releases the GIL during | |
| torch ops). The upsert workers are I/O-bound on the Qdrant HTTP API. They run | |
| truly in parallel; total wall time β max(embed_time, upsert_time / N_WORKERS). | |
| Local cache: | |
| .cache/corpus_vectors_<model>.npy memmap, shape (n_docs, dim), float32 | |
| .cache/corpus_vectors_<model>.progress single int β docs durably embedded | |
| On re-run: | |
| - Embed picks up at sidecar progress (no cache β 0). | |
| - Upsert resumes from Qdrant's points_count. | |
| - When the cache is full but Qdrant is empty (cluster-reap recovery), no | |
| embedding happens β vectors stream from disk straight to upsert workers. | |
| ~3-5 min instead of ~15-30. | |
| - Upserts are idempotent on id, so re-pushing a cached batch that Qdrant | |
| already has is a no-op (re-overwrites the same point). | |
| Progress bars: | |
| Two stacked tqdm bars β embed (top) and upsert (bottom) β update independently. | |
| Usage: | |
| uv run python scripts/index_corpus.py # build/resume the index | |
| uv run python scripts/index_corpus.py --recreate # drop and re-upsert (keeps cache) | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import os | |
| import queue | |
| import random | |
| import sys | |
| import threading | |
| import time | |
| sys.path.insert(0, os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "src")) | |
| import numpy as np | |
| from tqdm import tqdm | |
| from qdrant_client import QdrantClient | |
| from qdrant_client.models import Distance, VectorParams, PointStruct | |
| from codesearch.config import ( | |
| QDRANT_URL, | |
| QDRANT_API_KEY, | |
| QDRANT_COLLECTION, | |
| EMBEDDING_MODEL, | |
| EMBEDDING_DIM, | |
| EMBED_INPUT, | |
| QDRANT_ON_DISK, | |
| ) | |
| from codesearch.data import load_codesearch, strip_docstring | |
| from codesearch.embedding import load_encoder | |
| _BATCH_SIZE = 256 | |
| _QUEUE_MAX = 20 # backpressure: blocks producer if upsert lags this far behind | |
| _UPSERT_WORKERS = 3 | |
| _CACHE_DIR = ".cache" | |
| _UPSERT_RETRIES = 5 | |
| def _cache_paths(model: str) -> tuple[str, str]: | |
| safe = model.replace("/", "_") | |
| # Non-default embedding inputs get a distinct cache so they never clobber the | |
| # code_tokens vectors (same model, different doc text). | |
| variant = "" if EMBED_INPUT == "code_tokens" else f"_{EMBED_INPUT}" | |
| return ( | |
| os.path.join(_CACHE_DIR, f"corpus_vectors_{safe}{variant}.npy"), | |
| os.path.join(_CACHE_DIR, f"corpus_vectors_{safe}{variant}.progress"), | |
| ) | |
| def _embed_text(doc: dict) -> str: | |
| """Doc text to embed, per EMBED_INPUT. Falls back to code_tokens when the | |
| source can't be parsed for stripping (keeps every doc embeddable).""" | |
| if EMBED_INPUT == "code_stripped": | |
| return strip_docstring(doc["code"]) or doc["code_tokens"] | |
| return doc["code_tokens"] | |
| def _read_progress(path: str) -> int: | |
| if not os.path.exists(path): | |
| return 0 | |
| try: | |
| with open(path) as f: | |
| return int((f.read() or "0").strip()) | |
| except (ValueError, OSError): | |
| return 0 | |
| def _write_progress(path: str, n: int) -> None: | |
| """Atomic write so a crash mid-write never leaves a half-written number.""" | |
| tmp = path + ".tmp" | |
| with open(tmp, "w") as f: | |
| f.write(str(n)) | |
| os.replace(tmp, path) | |
| def _open_memmap(path: str, n_rows: int, dim: int) -> np.memmap: | |
| """Open the vector cache as a memmap; create + pre-allocate if absent.""" | |
| if not os.path.exists(path): | |
| return np.memmap(path, dtype=np.float32, mode="w+", shape=(n_rows, dim)) | |
| expected = n_rows * dim * 4 | |
| actual = os.path.getsize(path) | |
| if actual != expected: | |
| raise RuntimeError( | |
| f"Cache size mismatch at {path}: {actual} bytes on disk, expected " | |
| f"{expected} ({n_rows:,} Γ {dim} Γ 4). Delete it (and the .progress " | |
| f"sidecar) to recreate." | |
| ) | |
| return np.memmap(path, dtype=np.float32, mode="r+", shape=(n_rows, dim)) | |
| def _upsert_worker( | |
| work_q: "queue.Queue", | |
| client: QdrantClient, | |
| bar: tqdm, | |
| bar_lock: threading.Lock, | |
| ) -> None: | |
| while True: | |
| item = work_q.get() | |
| if item is None: | |
| return | |
| offset, doc_ids, vectors = item | |
| points = [ | |
| PointStruct( | |
| id=offset + j, | |
| vector=vectors[j].tolist(), | |
| payload={"doc_id": doc_ids[j]}, | |
| ) | |
| for j in range(len(doc_ids)) | |
| ] | |
| for attempt in range(_UPSERT_RETRIES): | |
| try: | |
| client.upsert(collection_name=QDRANT_COLLECTION, points=points) | |
| break | |
| except Exception as e: | |
| if attempt == _UPSERT_RETRIES - 1: | |
| raise | |
| wait = 2 ** attempt + random.random() | |
| with bar_lock: | |
| tqdm.write(f"[upsert] retry in {wait:.1f}s after: {e}") | |
| time.sleep(wait) | |
| with bar_lock: | |
| bar.update(len(points)) | |
| def _produce( | |
| corpus: list[dict], | |
| model, # SentenceTransformer | UniXcoderEncoder | None (see codesearch.embedding) | |
| vectors_mmap: np.memmap, | |
| progress_path: str, | |
| embed_progress: int, | |
| upsert_start: int, | |
| work_q: "queue.Queue", | |
| embed_bar: tqdm, | |
| bar_lock: threading.Lock, | |
| ) -> None: | |
| n = len(corpus) | |
| # Phase 1: any docs already in the cache but not yet in Qdrant. | |
| # Stream them from disk straight to the upsert queue, no embedding. | |
| if upsert_start < embed_progress: | |
| for i in range(upsert_start, embed_progress, _BATCH_SIZE): | |
| end = min(i + _BATCH_SIZE, embed_progress) | |
| doc_ids = [corpus[k]["id"] for k in range(i, end)] | |
| # Copy out of the memmap so the worker isn't holding a slice that | |
| # could be invalidated when the file is closed. | |
| vecs = np.array(vectors_mmap[i:end], copy=True) | |
| work_q.put((i, doc_ids, vecs)) | |
| # Phase 2: embed remaining docs, write to cache, push to queue. | |
| if embed_progress < n: | |
| assert model is not None, "Model required when embed_progress < n" | |
| for i in range(embed_progress, n, _BATCH_SIZE): | |
| end = min(i + _BATCH_SIZE, n) | |
| batch = corpus[i:end] | |
| texts = [_embed_text(doc) for doc in batch] | |
| vecs = model.encode( | |
| texts, show_progress_bar=False, normalize_embeddings=True | |
| ).astype(np.float32) | |
| # Durably persist: write into memmap, flush, then bump sidecar. | |
| vectors_mmap[i:end] = vecs | |
| vectors_mmap.flush() | |
| _write_progress(progress_path, end) | |
| with bar_lock: | |
| embed_bar.update(end - i) | |
| work_q.put((i, [d["id"] for d in batch], vecs)) | |
| # Sentinel one per worker β clean shutdown. | |
| for _ in range(_UPSERT_WORKERS): | |
| work_q.put(None) | |
| def build_index(recreate: bool) -> None: | |
| # ------------------------------------------------------------------------- | |
| # 1. Corpus | |
| # ------------------------------------------------------------------------- | |
| corpus, queries = load_codesearch(n=-1) | |
| n = len(corpus) | |
| print(f"Corpus: {n:,} docs") | |
| # GOLDS_FIRST (M5 sample-eval support): embed the eval-gold docs before the | |
| # rest so a partial index already contains every query's target. Run-local | |
| # only β does NOT change data.load_codesearch's global order, so the BM25 | |
| # slim-index id fingerprint and the live Space are untouched. Order is | |
| # irrelevant to retrieval correctness (search is by vector; payload=doc_id). | |
| # scripts/sample_eval_partial.py reproduces this exact order from the cache. | |
| if os.getenv("GOLDS_FIRST", "false").lower() in ("1", "true", "yes"): | |
| gold_ids = {q["relevant_id"] for q in queries} | |
| golds = [d for d in corpus if d["id"] in gold_ids] | |
| rest = [d for d in corpus if d["id"] not in gold_ids] | |
| corpus = golds + rest | |
| print(f"[golds-first] {len(golds):,} gold docs first, then {len(rest):,} others") | |
| # ------------------------------------------------------------------------- | |
| # 2. Vector cache | |
| # ------------------------------------------------------------------------- | |
| os.makedirs(_CACHE_DIR, exist_ok=True) | |
| vec_path, prog_path = _cache_paths(EMBEDDING_MODEL) | |
| vectors_mmap = _open_memmap(vec_path, n, EMBEDDING_DIM) | |
| embed_progress = min(_read_progress(prog_path), n) | |
| print(f"Cache : {embed_progress:,}/{n:,} vectors already embedded ({vec_path})") | |
| # ------------------------------------------------------------------------- | |
| # 3. Qdrant collection + resume offset | |
| # ------------------------------------------------------------------------- | |
| print(f"Connecting to Qdrant at {QDRANT_URL} ...") | |
| client = QdrantClient(url=QDRANT_URL, api_key=QDRANT_API_KEY, timeout=60) | |
| existing = [c.name for c in client.get_collections().collections] | |
| if recreate and QDRANT_COLLECTION in existing: | |
| print(f"Dropping collection '{QDRANT_COLLECTION}' (--recreate).") | |
| client.delete_collection(QDRANT_COLLECTION) | |
| existing = [c for c in existing if c != QDRANT_COLLECTION] | |
| if QDRANT_COLLECTION not in existing: | |
| # on-disk vectors (QDRANT_ON_DISK) when a collection would push total | |
| # Qdrant RAM past the 1GB free tier β the default code_tokens collection | |
| # already fills ~667MB, and 768-dim models are ~1.3GB alone. Keeps vectors | |
| # on disk (HNSW graph stays in RAM) β fine for a batch eval. | |
| on_disk = QDRANT_ON_DISK | |
| print( | |
| f"Creating collection '{QDRANT_COLLECTION}' " | |
| f"(dim={EMBEDDING_DIM}, cosine, on_disk={on_disk})." | |
| ) | |
| client.create_collection( | |
| collection_name=QDRANT_COLLECTION, | |
| vectors_config=VectorParams( | |
| size=EMBEDDING_DIM, distance=Distance.COSINE, on_disk=on_disk | |
| ), | |
| ) | |
| upsert_start = 0 | |
| else: | |
| upsert_start = client.get_collection(QDRANT_COLLECTION).points_count or 0 | |
| print(f"Qdrant: {upsert_start:,}/{n:,} points already upserted") | |
| if embed_progress >= n and upsert_start >= n: | |
| print("Both cache and Qdrant are full β nothing to do.") | |
| return | |
| # ------------------------------------------------------------------------- | |
| # 4. Load model only if we still need to embed | |
| # ------------------------------------------------------------------------- | |
| if embed_progress < n: | |
| print(f"Loading model: {EMBEDDING_MODEL}") | |
| model = load_encoder(EMBEDDING_MODEL) | |
| else: | |
| model = None | |
| print("Embedding step skipped β cache is complete.") | |
| # ------------------------------------------------------------------------- | |
| # 5. Run the parallel pipeline | |
| # ------------------------------------------------------------------------- | |
| bar_lock = threading.Lock() | |
| work_q: queue.Queue = queue.Queue(maxsize=_QUEUE_MAX) | |
| embed_bar = tqdm(total=n, initial=embed_progress, desc="Embed ", position=0, unit="doc") | |
| upsert_bar = tqdm(total=n, initial=upsert_start, desc="Upsert", position=1, unit="doc") | |
| t0 = time.perf_counter() | |
| workers = [ | |
| threading.Thread( | |
| target=_upsert_worker, | |
| args=(work_q, client, upsert_bar, bar_lock), | |
| daemon=True, | |
| name=f"upsert-{i}", | |
| ) | |
| for i in range(_UPSERT_WORKERS) | |
| ] | |
| for w in workers: | |
| w.start() | |
| try: | |
| _produce( | |
| corpus, model, vectors_mmap, prog_path, | |
| embed_progress, upsert_start, | |
| work_q, embed_bar, bar_lock, | |
| ) | |
| for w in workers: | |
| w.join() | |
| finally: | |
| embed_bar.close() | |
| upsert_bar.close() | |
| # Move cursor below both bars before printing summary. | |
| print() | |
| elapsed = time.perf_counter() - t0 | |
| final = client.get_collection(QDRANT_COLLECTION).points_count or 0 | |
| print(f"Done in {elapsed:.0f}s. Collection '{QDRANT_COLLECTION}' now has {final:,} vectors.") | |
| def main() -> None: | |
| parser = argparse.ArgumentParser( | |
| description="Embed the corpus and upsert into Qdrant (parallel + cached)." | |
| ) | |
| parser.add_argument( | |
| "--recreate", | |
| action="store_true", | |
| help="Drop the Qdrant collection before indexing. Keeps the local " | |
| "vector cache β re-upsert is fast.", | |
| ) | |
| args = parser.parse_args() | |
| build_index(recreate=args.recreate) | |
| if __name__ == "__main__": | |
| main() | |