diff --git a/README.md b/README.md index 865bf38..3b40476 100644 --- a/README.md +++ b/README.md @@ -4,14 +4,133 @@ Implementation details are provided in our [technical note](Technical_Details_of_V3DB.pdf). +## Environment Setup + +All commands below assume the conda environment `V3DB` (Python 3.11) and a Rust nightly toolchain (`rust-toolchain.toml`). + +```bash +conda activate V3DB + +# Python dependencies +pip install -r requirements.txt +# or manually: +# pip install maturin numpy tqdm scipy scikit-learn matplotlib duckdb faiss-cpu +# pip install fastapi "uvicorn[standard]" # only for the interactive demo + +# MS MARCO embedding generation additionally requires: +# pip install torch transformers +``` + +Notes: + +- `faiss-cpu` is the default. If a GPU build matching your GPU architecture is available you may use `faiss-gpu-cu12` instead — `ivf_pq/util/kmeans.py` auto-detects whether GPU kernels actually run and falls back to CPU. +- Some scripts use `python -m ` (e.g. `python -m tests.pipeline`). Running them as `python tests/pipeline.py` will fail with `ModuleNotFoundError` because the project root is not on `sys.path` in that mode. + ## Build -Build and install the Python extension: +Build and install the Python extension (requires Rust nightly): ```bash maturin develop --release ``` +This installs the extension as `zk_IVF_PQ.zk_IVF_PQ`, matching the import path `from zk_IVF_PQ.zk_IVF_PQ import ...` used throughout the codebase. + +## Interactive Demo (White Theme) + +A white-themed, fully local (no external CDN) interactive web demo that compares **Standard IVF-PQ**, **ZK IVF-PQ**, and **Brute-force Ground Truth** side-by-side and can generate set-based + Merkle ZK proofs via the Rust extension. + +### Datasets + +The demo auto-discovers datasets under `data/`: + +| Name | Directory | Content | Size on disk | +|------|-----------|---------|--------------| +| SIFT small | `data/siftsmall/` | 10k × 128 dim | ~5 MB | +| SIFT1M | `data/sift/` | 1M × 128 dim | ~550 MB | +| GIST1M | `data/gist/` | 1M × 960 dim | ~5.4 GB | + +If missing, download and extract them from `ftp://ftp.irisa.fr/local/texmex/corpus/`: + +```bash +cd data +wget ftp://ftp.irisa.fr/local/texmex/corpus/sift.tar.gz && tar xzf sift.tar.gz +wget ftp://ftp.irisa.fr/local/texmex/corpus/gist.tar.gz && tar xzf gist.tar.gz +wget ftp://ftp.irisa.fr/local/texmex/corpus/siftsmall.tar.gz && tar xzf siftsmall.tar.gz +``` + +The GIST1M archive is ~2.7 GB compressed. `data/siftsmall` is sufficient for a quick smoke test; SIFT1M and GIST1M are optional. + +### Start + +```bash +conda activate V3DB +cd /home/huhaoran/V3DB + +# foreground (logs to terminal) +python -m demo.server # default 0.0.0.0:8000 +python -m demo.server 8080 # custom port + +# background (logs to /tmp/demo.log) +nohup python -m demo.server > /tmp/demo.log 2>&1 & +tail -f /tmp/demo.log +``` + +Then open `http://127.0.0.1:8000` in a browser. + +On first use of a dataset the backend builds and caches models under `data/demo_cache/{name}_{std,zk,proof,pca}.npz`: + +- Standard IVF-PQ (`ivf_pq/standard.py`) +- ZK integer version (`rescale_database` + `cluster_bound` rebalance + Merkle commitments in `ivf_pq/merkle_zk.py`) +- 2D PCA projection for visualization + +Build must run from the **project root** (`/home/huhaoran/V3DB`) so that `python -m demo.server` resolves the `demo` package. Subsequent launches load the cache in seconds and render all three datasets as `ready`. + +GIST1M (960 dim, 512 clusters + 8×256 PQ codebooks) takes ~30–40 minutes to build on CPU — the status badge polls `GET /api/status/gist` and the log shows Faiss iteration progress. + +### Stop + +```bash +pkill -f "demo.server" +# or, if started with nohup: +pkill -9 -f "demo.server" +``` + +Check whether it is still running: + +```bash +curl -s http://127.0.0.1:8000/api/datasets | python3 -m json.tool +# or +ps aux | grep demo.server | grep -v grep +``` + +### How to use + +1. Select a dataset on the left (status badges: `ready` / `building` / `idle`). The first build is triggered automatically. +2. Adjust `query id`, `n_probe` (1–64), `top_k` (5–100), and optionally enable **Generate and verify ZK proof** (set-based + Merkle, backed by the Rust extension; proving takes seconds to tens of seconds). +3. Click **Search** — the right pane shows 6 metric cards (forced to a single row), PCA scatter, latency/recall/cluster-size charts, and a three-column comparison table with hit badges. + +API health check: + +```bash +curl -s http://127.0.0.1:8000/api/datasets | python3 -m json.tool +curl -s http://127.0.0.1:8000/api/status/gist +curl -s -X POST http://127.0.0.1:8000/api/search \ + -H "Content-Type: application/json" \ + -d '{"dataset":"siftsmall","query_id":0,"n_probe":8,"top_k":10,"proof":false}' | python3 -m json.tool +``` + +Demo file layout: + +``` +demo/ + server.py # FastAPI backend (lazy loading, model cache, search + proof API) + static/ + index.html # page skeleton + style.css # white theme (no external deps) + app.js # frontend logic, pure Canvas charts (no external deps) +``` + ## Experiment 1: Retrieval Utility Evaluation ### Classic ANN (SIFT1M / GIST1M) diff --git a/demo/__init__.py b/demo/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/demo/server.py b/demo/server.py new file mode 100644 index 0000000..e73c32a --- /dev/null +++ b/demo/server.py @@ -0,0 +1,698 @@ +"""zk-IVF-PQ 交互式演示后端。 + +提供: + - 数据集状态查询 / 模型懒构建(标准 IVF-PQ + ZK 整数版 IVF-PQ + Merkle 证明结构 + PCA 投影) + - 检索接口: 标准 IVF-PQ / ZK IVF-PQ / 暴力精确检索三方对比 + - 可选的 set-based + Merkle 零知识证明生成(调用 Rust 扩展) + +启动: 在项目根目录执行 python -m demo.server (默认 0.0.0.0:8000) +""" + +from __future__ import annotations + +import sys +import threading +import time +import traceback +from dataclasses import dataclass, field +from pathlib import Path + +ROOT = Path(__file__).resolve().parent.parent +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +import numpy as np +import uvicorn +from fastapi import FastAPI, HTTPException +from fastapi.responses import FileResponse +from fastapi.staticfiles import StaticFiles +from pydantic import BaseModel + +from ivf_pq import MAX_SCALE, rescale_database, rescale_query +from ivf_pq.merkle_zk import _compute_cluster_root +from ivf_pq.standard import ivf_pq_learn as std_learn +from ivf_pq.zk import ivf_pq_learn as zk_train +from vec_data_load.sift import SIFT +from zk_IVF_PQ.zk_IVF_PQ import py_set_based_with_merkle + +DATA_DIR = ROOT / "data" +CACHE_DIR = DATA_DIR / "demo_cache" +STATIC_DIR = Path(__file__).resolve().parent / "static" + +SCALE_N = MAX_SCALE # 65536, ZK 整数化上界 +PCA_SAMPLE = 20000 + +DATASET_CFG: dict[str, dict] = { + "siftsmall": {"n_list": 256, "n_iter": 25, "M": 8, "K": 256, "cluster_bound": 128}, + "sift": {"n_list": 1024, "n_iter": 25, "M": 8, "K": 256, "cluster_bound": 2048}, + "gist": {"n_list": 512, "n_iter": 25, "M": 8, "K": 256, "cluster_bound": 4096}, +} + +DATASET_DESC = { + "siftsmall": "SIFT small (1万条 128 维)", + "sift": "SIFT1M (100 万条 128 维)", + "gist": "GIST1M (100 万条 960 维)", +} + + +def _dataset_dir_ok(name: str) -> bool: + return (DATA_DIR / name / f"{name}_base.fvecs").exists() + + +def _pow2_ceil(x: int) -> int: + c = 1 + while c < x: + c *= 2 + return c + + +def _id_groups_from_labels(labels: np.ndarray, n_list: int) -> dict[int, np.ndarray]: + labels = np.asarray(labels) + order = np.argsort(labels, kind="stable").astype(np.int64) + sorted_labels = labels[order] + bounds = np.searchsorted(sorted_labels, np.arange(n_list + 1)) + return {c: order[bounds[c] : bounds[c + 1]] for c in range(n_list)} + + +@dataclass +class ModelBundle: + name: str + base: np.ndarray # (N, D) float32 原始向量 + query_vecs: np.ndarray # (Q, D) float32 + n_list: int + M: int + K: int + # 标准浮点 IVF-PQ + std_labels: np.ndarray = None # (N,) + std_center: np.ndarray = None # (n_list, D) float32 + std_code_books: np.ndarray = None # (M, K, d) float32 + std_quant_vecs: np.ndarray = None # (N, M) int64 + std_id_groups: dict = field(default_factory=dict) + # ZK 整数版 + zk_labels: np.ndarray = None + zk_center: np.ndarray = None # (n_list, D) int64 + zk_code_books: np.ndarray = None # (M, K, d) int64 + zk_quant_vecs: np.ndarray = None # (N, M) int64 + zk_id_groups: dict = field(default_factory=dict) + changed_count: int = 0 + v_min: float = 0.0 + v_max: float = 0.0 + # Merkle 证明结构(逐簇 padding 到 capacity) + capacity: int = 0 + vpqs_all: np.ndarray = None # (n_list, capacity, M) int64 + valids_all: np.ndarray = None # (n_list, capacity) int64 + items_all: np.ndarray = None # (n_list, capacity) int64 + roots: np.ndarray = None # (n_list,) uint64 + # 可视化 + pca_comps: np.ndarray = None # (2, D) float32 + pca_mean: np.ndarray = None # (D,) float32 + pca_sample_ids: np.ndarray = None # (S,) + pca_sample_proj: np.ndarray = None # (S, 2) float32 + # 簇规模 + std_sizes: list = field(default_factory=list) + zk_sizes: list = field(default_factory=list) + + +MODELS: dict[str, ModelBundle] = {} +BUILD_STATE: dict[str, dict] = { + name: {"state": "idle", "message": "", "error": None} for name in DATASET_CFG +} +BUILD_LOCKS: dict[str, threading.Lock] = { + name: threading.Lock() for name in DATASET_CFG +} +BUILD_THREADS: dict[str, threading.Thread] = {} + + +def _set_state(name: str, state: str, message: str = "", error: str | None = None): + BUILD_STATE[name] = {"state": state, "message": message, "error": error} + print(f"[demo:{name}] {state}: {message}", flush=True) + + +def _save_npz(path: Path, **arrays): + path.parent.mkdir(parents=True, exist_ok=True) + np.savez_compressed(path, **arrays) + + +def _load_dataset(name: str) -> SIFT: + return SIFT(str(DATA_DIR / name)) + + +def _build_std_model(ds: SIFT, cfg: dict, name: str) -> None: + _set_state(name, "building", "训练标准浮点 IVF-PQ (coarse 聚类 + PQ 码本)…") + labels, center, code_books, quant_vecs, _ = std_learn( + ds.base_vecs, + n_list=cfg["n_list"], + n_iter=cfg["n_iter"], + M=cfg["M"], + K=cfg["K"], + random_state=1234, + layout=None, + ) + _save_npz( + CACHE_DIR / f"{name}_std.npz", + labels=labels.astype(np.int32), + center=center.astype(np.float32), + code_books=code_books.astype(np.float32), + quant_vecs=quant_vecs.astype(np.int16), + ) + + +def _build_zk_model(ds: SIFT, cfg: dict, name: str) -> float: + _set_state(name, "building", "整数化数据库并训练 ZK 版 IVF-PQ (含簇上界重平衡)…") + scaled, v_min, v_max = rescale_database(ds.base_vecs, SCALE_N) + labels, center, code_books, quant_vecs, _, changed_count = zk_train( + scaled, + n_list=cfg["n_list"], + n_iter=cfg["n_iter"], + M=cfg["M"], + K=cfg["K"], + random_state=1234, + cluster_bound=cfg["cluster_bound"], + layout=None, + ) + _save_npz( + CACHE_DIR / f"{name}_zk.npz", + labels=labels.astype(np.int32), + center=center.astype(np.int64), + code_books=np.rint(code_books).astype(np.int32), + quant_vecs=quant_vecs.astype(np.int16), + changed_count=np.int64(changed_count), + v_min=np.float64(v_min), + v_max=np.float64(v_max), + scale_n=np.int64(SCALE_N), + ) + del scaled + return changed_count + + +def _build_proof_structures(name: str, cfg: dict) -> None: + zk = np.load(CACHE_DIR / f"{name}_zk.npz") + labels = zk["labels"] + quant_vecs = zk["quant_vecs"].astype(np.int64) + n_list = cfg["n_list"] + M = cfg["M"] + + id_groups = _id_groups_from_labels(labels, n_list) + max_size = max((g.size for g in id_groups.values()), default=1) + capacity = _pow2_ceil(int(max_size)) + + _set_state( + name, + "building", + f"构建 Merkle 承诺结构 (capacity={capacity}, 共 {n_list} 簇)…", + ) + + vpqs_all = np.zeros((n_list, capacity, M), dtype=np.int64) + valids_all = np.zeros((n_list, capacity), dtype=np.int64) + items_all = np.zeros((n_list, capacity), dtype=np.int64) + for ci, ids in id_groups.items(): + take = min(ids.size, capacity) + if take == 0: + continue + items_all[ci, :take] = ids[:take] + valids_all[ci, :take] = 1 + vpqs_all[ci, :take] = quant_vecs[ids[:take]] + + roots = np.zeros(n_list, dtype=np.uint64) + for ci in range(n_list): + roots[ci] = _compute_cluster_root(ci, vpqs_all[ci], valids_all[ci], items_all[ci]) + if ci % max(1, n_list // 10) == 0: + _set_state( + name, + "building", + f"计算簇 Merkle 根 {ci}/{n_list} …", + ) + + _save_npz( + CACHE_DIR / f"{name}_proof.npz", + capacity=np.int64(capacity), + vpqs_all=vpqs_all.astype(np.int16), + valids_all=valids_all.astype(np.uint8), + items_all=items_all.astype(np.int32), + roots=roots, + ) + + +def _build_pca(name: str, base: np.ndarray) -> None: + _set_state(name, "building", "计算 2D PCA 投影 (用于可视化)…") + rng = np.random.default_rng(0) + sample_ids = np.sort(rng.choice(base.shape[0], size=min(PCA_SAMPLE, base.shape[0]), replace=False)) + X = base[sample_ids].astype(np.float64) + mean = X.mean(axis=0) + Xc = X - mean + # 协方差特征分解 (D 一般为 128/960, 很快) + cov = Xc.T @ Xc / max(1, Xc.shape[0] - 1) + eigvals, eigvecs = np.linalg.eigh(cov) + comps = eigvecs[:, -2:].T.astype(np.float32) # (2, D) + proj = ((base[sample_ids] - mean.astype(np.float32)) @ comps.T).astype(np.float32) + _save_npz( + CACHE_DIR / f"{name}_pca.npz", + comps=comps, + mean=mean.astype(np.float32), + sample_ids=sample_ids.astype(np.int32), + sample_proj=proj, + ) + + +def _load_bundle(name: str, ds: SIFT, cfg: dict) -> ModelBundle: + std = np.load(CACHE_DIR / f"{name}_std.npz") + zk = np.load(CACHE_DIR / f"{name}_zk.npz") + proof = np.load(CACHE_DIR / f"{name}_proof.npz") + pca = np.load(CACHE_DIR / f"{name}_pca.npz") + + n_list, M, K = cfg["n_list"], cfg["M"], cfg["K"] + std_labels = std["labels"] + zk_labels = zk["labels"] + + bundle = ModelBundle( + name=name, + base=ds.base_vecs, + query_vecs=ds.query_vecs, + n_list=n_list, + M=M, + K=K, + std_labels=std_labels, + std_center=std["center"], + std_code_books=std["code_books"], + std_quant_vecs=std["quant_vecs"].astype(np.int64), + std_id_groups=_id_groups_from_labels(std_labels, n_list), + zk_labels=zk_labels, + zk_center=zk["center"].astype(np.int64), + zk_code_books=zk["code_books"].astype(np.int64), + zk_quant_vecs=zk["quant_vecs"].astype(np.int64), + zk_id_groups=_id_groups_from_labels(zk_labels, n_list), + changed_count=int(zk["changed_count"]), + v_min=float(zk["v_min"]), + v_max=float(zk["v_max"]), + capacity=int(proof["capacity"]), + vpqs_all=proof["vpqs_all"].astype(np.int64), + valids_all=proof["valids_all"].astype(np.int64), + items_all=proof["items_all"].astype(np.int64), + roots=proof["roots"].astype(np.uint64), + pca_comps=pca["comps"], + pca_mean=pca["mean"], + pca_sample_ids=pca["sample_ids"], + pca_sample_proj=pca["sample_proj"], + std_sizes=[int(g.size) for g in _id_groups_from_labels(std_labels, n_list).values()], + zk_sizes=[int(g.size) for g in _id_groups_from_labels(zk_labels, n_list).values()], + ) + return bundle + + +def _build_worker(name: str): + cfg = DATASET_CFG[name] + try: + if not _dataset_dir_ok(name): + _set_state(name, "missing", "数据集文件不存在", error="missing") + return + ds = _load_dataset(name) + if not (CACHE_DIR / f"{name}_std.npz").exists(): + _build_std_model(ds, cfg, name) + if not (CACHE_DIR / f"{name}_zk.npz").exists(): + _build_zk_model(ds, cfg, name) + if not (CACHE_DIR / f"{name}_proof.npz").exists(): + _build_proof_structures(name, cfg) + if not (CACHE_DIR / f"{name}_pca.npz").exists(): + _build_pca(name, ds.base_vecs) + _set_state(name, "building", "加载模型到内存…") + bundle = _load_bundle(name, ds, cfg) + MODELS[name] = bundle + _set_state(name, "ready", "模型就绪") + except Exception: + traceback.print_exc() + _set_state(name, "error", "构建失败", error=traceback.format_exc()[-2000:]) + + +def _ensure_building(name: str) -> None: + if name not in DATASET_CFG: + raise HTTPException(status_code=404, detail=f"未知数据集: {name}") + if not _dataset_dir_ok(name): + raise HTTPException( + status_code=404, + detail=f"数据集 {name} 未下载: 请将 {name}.tar.gz 解压到 data/{name}/", + ) + st = BUILD_STATE[name] + if st["state"] in ("ready", "building"): + return + with BUILD_LOCKS[name]: + st = BUILD_STATE[name] + if st["state"] in ("ready", "building"): + return + t = threading.Thread(target=_build_worker, args=(name,), daemon=True) + BUILD_THREADS[name] = t + t.start() + + +# ---------------------------------------------------------------------------- +# 检索核心逻辑 +# ---------------------------------------------------------------------------- + + +def _std_search(bundle: ModelBundle, q: np.ndarray, n_probe: int, top_k: int): + t0 = time.perf_counter() + center = bundle.std_center + code_books = bundle.std_code_books + quant_vecs = bundle.std_quant_vecs + M, K, d = code_books.shape + m_indices = np.arange(M)[None, :] + + diff = center - q[None, :] + dist2 = np.einsum("ij,ij->i", diff, diff) + cluster_order = np.argsort(dist2, kind="stable") + + all_ids, all_dis2 = [], [] + for ci in cluster_order[: min(n_probe, bundle.n_list)]: + ids = bundle.std_id_groups.get(int(ci)) + if ids is None or ids.size == 0: + continue + res_query = q - center[ci] + codes = quant_vecs[ids] + recon = code_books[m_indices, codes, :].reshape(codes.shape[0], -1) + d2 = np.einsum("ij,ij->i", recon - res_query, recon - res_query) + all_ids.append(ids) + all_dis2.append(d2) + if not all_ids: + return np.empty(0, np.int64), np.empty(0), time.perf_counter() - t0 + ids = np.concatenate(all_ids) + dis = np.concatenate(all_dis2) + order = np.argsort(dis, kind="stable") + k = min(top_k, order.size) + return ids[order[:k]].astype(np.int64), dis[order[:k]], time.perf_counter() - t0 + + +def _zk_search(bundle: ModelBundle, q_scaled: np.ndarray, n_probe: int, top_k: int): + t0 = time.perf_counter() + center = bundle.zk_center + code_books = bundle.zk_code_books + quant_vecs = bundle.zk_quant_vecs + M, K, d = code_books.shape + m_indices = np.arange(M)[None, :] + + diff = center - q_scaled[None, :] + dist2 = (diff * diff).sum(axis=1, dtype=np.int64) + cluster_order = np.argsort(dist2, kind="stable") + + all_ids, all_dis2 = [], [] + for ci in cluster_order[: min(n_probe, bundle.n_list)]: + ids = bundle.zk_id_groups.get(int(ci)) + if ids is None or ids.size == 0: + continue + delta = q_scaled - center[ci] + codes = quant_vecs[ids] + recon = code_books[m_indices, codes, :].reshape(codes.shape[0], -1) + d2 = ((recon - delta) * (recon - delta)).sum(axis=1) + all_ids.append(ids) + all_dis2.append(d2) + if not all_ids: + return np.empty(0, np.int64), np.empty(0, np.int64), time.perf_counter() - t0 + ids = np.concatenate(all_ids) + dis = np.concatenate(all_dis2) + order = np.argsort(dis, kind="stable") + k = min(top_k, order.size) + return ids[order[:k]].astype(np.int64), dis[order[:k]], time.perf_counter() - t0 + + +def _brute_topk(base: np.ndarray, q: np.ndarray, top_k: int, chunk: int = 200_000): + t0 = time.perf_counter() + n = base.shape[0] + best_dis = np.full(top_k, np.inf, dtype=np.float64) + best_idx = np.full(top_k, -1, dtype=np.int64) + for s in range(0, n, chunk): + e = min(s + chunk, n) + d = base[s:e] - q[None, :] + d2 = np.einsum("ij,ij->i", d, d).astype(np.float64) + if top_k < e - s: + part = np.argpartition(d2, top_k - 1)[:top_k] + else: + part = np.arange(e - s) + cand_dis = np.concatenate([best_dis, d2[part]]) + cand_idx = np.concatenate([best_idx, (part + s).astype(np.int64)]) + order = np.argsort(cand_dis, kind="stable")[:top_k] + best_dis = cand_dis[order] + best_idx = cand_idx[order] + return best_idx, best_dis, time.perf_counter() - t0 + + +def _zk_proof(bundle: ModelBundle, q_scaled: np.ndarray, n_probe: int, top_k: int): + center = bundle.zk_center + diff = center - q_scaled[None, :] + dist2 = (diff * diff).sum(axis=1, dtype=np.int64) + order = np.argsort(dist2, kind="stable") + cluster_idxes = order[: min(n_probe, bundle.n_list)] + cluster_idx_dis = np.stack([order, dist2[order]], axis=1).astype(np.int64) + + vpqss = bundle.vpqs_all[cluster_idxes] + valids = bundle.valids_all[cluster_idxes] + itemss = bundle.items_all[cluster_idxes] + + result = py_set_based_with_merkle( + q_scaled.astype(np.int64), + center, + vpqss, + valids, + itemss, + bundle.zk_code_books, + bundle.roots, + int(top_k), + cluster_idx_dis, + [], # ordered_vpqss_item_dis 由 Rust 端内部重算 + ) + keys = ("build_time", "prove_time", "verify_time", "proof_size", "memory_used", "num_gates") + return dict(zip(keys, result)) + + +def _viz_payload( + bundle: ModelBundle, + q: np.ndarray, + probed_clusters: np.ndarray, + std_ids: np.ndarray, + zk_ids: np.ndarray, + gt_ids: np.ndarray, + max_per_cluster: int = 60, +): + comps = bundle.pca_comps + mean = bundle.pca_mean + + def proj_of(ids: np.ndarray) -> np.ndarray: + ids = np.asarray(ids, dtype=np.int64) + if ids.size == 0: + return np.empty((0, 2), dtype=np.float32) + return ((bundle.base[ids] - mean) @ comps.T).astype(np.float32) + + rng = np.random.default_rng(7) + pts = [] + for ci in probed_clusters: + ids = bundle.std_id_groups.get(int(ci)) + if ids is None or ids.size == 0: + continue + if ids.size > max_per_cluster: + ids = rng.choice(ids, size=max_per_cluster, replace=False) + proj = proj_of(ids) + for (x, y), vid in zip(proj, ids): + pts.append([round(float(x), 3), round(float(y), 3), int(ci), int(vid)]) + + def named(proj: np.ndarray, ids: np.ndarray): + return [ + [round(float(x), 3), round(float(y), 3), int(vid)] + for (x, y), vid in zip(proj, np.asarray(ids, dtype=np.int64)) + ] + + q_proj = ((q[None, :].astype(np.float32) - mean) @ comps.T)[0] + return { + "points": pts, + "query": [round(float(q_proj[0]), 3), round(float(q_proj[1]), 3)], + "std_topk": named(proj_of(std_ids), std_ids), + "zk_topk": named(proj_of(zk_ids), zk_ids), + "gt_topk": named(proj_of(gt_ids), gt_ids), + } + + +# ---------------------------------------------------------------------------- +# FastAPI 应用 +# ---------------------------------------------------------------------------- + +app = FastAPI(title="zk-IVF-PQ demo") + + +class SearchReq(BaseModel): + dataset: str + query_id: int + n_probe: int = 8 + top_k: int = 10 + proof: bool = False + + +@app.get("/api/datasets") +def list_datasets(): + out = [] + for name, cfg in DATASET_CFG.items(): + st = BUILD_STATE[name] + available = _dataset_dir_ok(name) + item = { + "name": name, + "desc": DATASET_DESC[name], + "available": available, + "state": st["state"] if available else "missing", + "message": st["message"], + "error": st["error"], + "n_list": cfg["n_list"], + "M": cfg["M"], + "K": cfg["K"], + "cluster_bound": cfg["cluster_bound"], + } + bundle = MODELS.get(name) + if bundle is not None: + item.update( + { + "N": int(bundle.base.shape[0]), + "D": int(bundle.base.shape[1]), + "Q": int(bundle.query_vecs.shape[0]), + "capacity": bundle.capacity, + "changed_count": bundle.changed_count, + } + ) + out.append(item) + return {"datasets": out} + + +@app.post("/api/prepare/{name}") +def prepare(name: str): + _ensure_building(name) + return {"ok": True} + + +@app.get("/api/status/{name}") +def status(name: str): + if name not in DATASET_CFG: + raise HTTPException(status_code=404, detail=f"未知数据集: {name}") + return BUILD_STATE[name] + + +@app.get("/api/cluster_stats/{name}") +def cluster_stats(name: str): + bundle = MODELS.get(name) + if bundle is None: + raise HTTPException(status_code=409, detail="模型尚未构建完成") + return { + "std_sizes": bundle.std_sizes, + "zk_sizes": bundle.zk_sizes, + "changed_count": bundle.changed_count, + "cluster_bound": DATASET_CFG[name]["cluster_bound"], + "capacity": bundle.capacity, + } + + +@app.post("/api/search") +def search(req: SearchReq): + name = req.dataset + bundle = MODELS.get(name) + if bundle is None: + st = BUILD_STATE[name] + if st.get("error") or st["state"] == "missing": + raise HTTPException(status_code=409, detail="模型不可用: " + st["message"]) + raise HTTPException(status_code=409, detail="模型尚未构建完成, 请先构建") + + Q = bundle.query_vecs.shape[0] + if req.query_id < 0 or req.query_id >= Q: + raise HTTPException(status_code=400, detail=f"query_id 超出范围 [0, {Q})") + n_probe = max(1, min(int(req.n_probe), bundle.n_list)) + top_k = max(1, min(int(req.top_k), 200)) + do_proof = bool(req.proof) + if do_proof and n_probe > 16: + raise HTTPException(status_code=400, detail="证明生成时 n_probe 最大支持 16") + + q = bundle.query_vecs[req.query_id].astype(np.float32) + + std_ids, std_dis, std_time = _std_search(bundle, q, n_probe, top_k) + + q_scaled = rescale_query(q, SCALE_N, bundle.v_min, bundle.v_max).astype(np.int64) + zk_ids, zk_dis, zk_time = _zk_search(bundle, q_scaled, n_probe, top_k) + + gt_ids, gt_dis, gt_time = _brute_topk(bundle.base, q, top_k) + + proof_res = None + if do_proof: + t0 = time.perf_counter() + proof_res = _zk_proof(bundle, q_scaled, n_probe, top_k) + proof_res["assemble_time"] = time.perf_counter() - t0 + proof_res = { + "build_time": float(proof_res["build_time"]), + "prove_time": float(proof_res["prove_time"]), + "verify_time": float(proof_res["verify_time"]), + "proof_size": int(proof_res["proof_size"]), + "memory_used": int(proof_res["memory_used"]), + "num_gates": int(proof_res["num_gates"]), + "verified": True, + } + + def rows(ids, dis, gt_set, zk_set=None, std_set=None): + out = [] + for i, (vid, vd) in enumerate(zip(ids, dis)): + r = { + "rank": i, + "id": int(vid), + "dist": float(vd), + "in_gt": int(vid) in gt_set, + } + if zk_set is not None: + r["in_zk"] = int(vid) in zk_set + if std_set is not None: + r["in_std"] = int(vid) in std_set + out.append(r) + return out + + gt_set = set(int(x) for x in gt_ids) + zk_set = set(int(x) for x in zk_ids) + std_set = set(int(x) for x in std_ids) + + k_eff = max(1, min(top_k, gt_ids.size)) + recall_std = len(std_set & gt_set) / k_eff + recall_zk = len(zk_set & gt_set) / k_eff + overlap = len(std_set & zk_set) / k_eff + + diff = bundle.std_center - q[None, :] + d2 = np.einsum("ij,ij->i", diff, diff) + probed = np.argsort(d2, kind="stable")[:n_probe] + + return { + "dataset": name, + "query_id": int(req.query_id), + "n_probe": int(n_probe), + "top_k": int(top_k), + "params": {"n_list": bundle.n_list, "M": bundle.M, "K": bundle.K, "scale_n": SCALE_N}, + "standard": { + "rows": rows(std_ids, std_dis, gt_set, zk_set=zk_set), + "time_ms": std_time * 1000.0, + "recall": recall_std, + }, + "zk": { + "rows": rows(zk_ids, zk_dis, gt_set, std_set=std_set), + "time_ms": zk_time * 1000.0, + "recall": recall_zk, + }, + "ground_truth": { + "rows": rows(gt_ids, gt_dis, gt_set), + "time_ms": gt_time * 1000.0, + }, + "metrics": { + "recall_std": recall_std, + "recall_zk": recall_zk, + "overlap_std_zk": overlap, + "gt_time_ms": gt_time * 1000.0, + }, + "proof": proof_res, + "viz": _viz_payload(bundle, q, probed, std_ids, zk_ids, gt_ids), + } + + +@app.get("/") +def index(): + return FileResponse(STATIC_DIR / "index.html") + + +app.mount("/", StaticFiles(directory=str(STATIC_DIR)), name="static") + + +if __name__ == "__main__": + port = int(sys.argv[1]) if len(sys.argv) > 1 else 8000 + uvicorn.run(app, host="0.0.0.0", port=port) diff --git a/demo/static/app.js b/demo/static/app.js new file mode 100644 index 0000000..e143e0d --- /dev/null +++ b/demo/static/app.js @@ -0,0 +1,649 @@ +/* zk-IVF-PQ 交互演示前端逻辑(无外部依赖) */ +"use strict"; + +const $ = (sel) => document.querySelector(sel); + +const state = { + datasets: [], + selected: null, + polling: null, + lastResult: null, +}; + +const C = { + std: "#2563eb", + zk: "#7c3aed", + gt: "#059669", + proof: "#d97706", + verify: "#059669", + query: "#dc2626", +}; + +function clusterColor(ci) { + const hue = (ci * 137.508) % 360; + return `hsl(${hue.toFixed(0)}, 62%, 52%)`; +} + +function fmtMs(ms) { + if (ms == null) return "—"; + if (ms >= 1000) return (ms / 1000).toFixed(2) + " s"; + if (ms >= 10) return ms.toFixed(0) + " ms"; + return ms.toFixed(2) + " ms"; +} + +function fmtInt(x) { + return Number(x).toLocaleString("en-US"); +} + +function showToast(msg, isError = false) { + const t = $("#toast"); + t.textContent = msg; + t.className = "toast show" + (isError ? " error" : ""); + clearTimeout(t._timer); + t._timer = setTimeout(() => (t.className = "toast" + (isError ? " error" : "")), 4000); +} + +async function api(path, opts) { + const res = await fetch(path, opts); + if (!res.ok) { + let detail = res.statusText; + try { + const body = await res.json(); + detail = body.detail || detail; + } catch (e) { /* ignore */ } + throw new Error(detail); + } + return res.json(); +} + +/* ================= 数据集 ================= */ + +async function loadDatasets() { + const data = await api("/api/datasets"); + state.datasets = data.datasets; + renderDatasets(); + updateControls(); +} + +function renderDatasets() { + const wrap = $("#dataset-list"); + wrap.innerHTML = ""; + for (const ds of state.datasets) { + const item = document.createElement("div"); + item.className = "ds-item" + (state.selected === ds.name ? " selected" : ""); + + const badgeTxt = { + idle: "待构建", + building: "构建中", + ready: "就绪", + error: "失败", + missing: "数据缺失", + }[ds.state] || ds.state; + + const meta = []; + if (ds.N != null) meta.push(`N=${fmtInt(ds.N)}`, `D=${ds.D}`, `Q=${fmtInt(ds.Q)}`); + if (ds.capacity != null) meta.push(`capacity=${ds.capacity}`); + + item.innerHTML = ` +
+ ${ds.name} + ${badgeTxt} +
+
${ds.desc} · n_list=${ds.n_list}, M=${ds.M}, K=${ds.K}
+ ${meta.length ? `
${meta.join(" · ")}
` : ""} + ${ds.state === "building" ? `
${ds.message || ""}
` : ""} + ${ds.state === "error" ? `
构建失败
` : ""} + ${ds.state === "missing" ? `
请先下载数据集并解压到 data/${ds.name}/
` : ""} + ${(ds.state === "idle") && ds.available ? `` : ""} + `; + item.addEventListener("click", () => selectDataset(ds.name)); + const btn = item.querySelector(".ds-build-btn"); + if (btn) btn.addEventListener("click", (e) => { e.stopPropagation(); prepareDataset(ds.name); }); + wrap.appendChild(item); + } +} + +function selectDataset(name) { + state.selected = name; + renderDatasets(); + updateControls(); + const ds = state.datasets.find((d) => d.name === name); + if (ds && ds.available && ds.state === "idle") { + prepareDataset(name); + } + if (ds && ds.state === "ready") { + $("#chk-proof").checked = name === "siftsmall"; + } +} + +async function prepareDataset(name) { + try { + await api(`/api/prepare/${name}`, { method: "POST" }); + showToast(`开始构建 ${name} 模型…`); + startPolling(name); + } catch (e) { + showToast(e.message, true); + } +} + +function startPolling(name) { + clearInterval(state.polling); + state.polling = setInterval(async () => { + try { + const st = await api(`/api/status/${name}`); + renderDatasets(); + if (st.state === "ready") { + clearInterval(state.polling); + state.polling = null; + await loadDatasets(); + showToast(`${name} 模型就绪`); + updateControls(); + if (state.selected === name) runSearch(); + } else if (st.state === "error" || st.state === "missing") { + clearInterval(state.polling); + state.polling = null; + showToast(`${name}: ${st.message}`, true); + } + } catch (e) { /* 网络抖动时继续轮询 */ } + }, 1500); +} + +function selectedDataset() { + return state.datasets.find((d) => d.name === state.selected); +} + +function updateControls() { + const ds = selectedDataset(); + const btn = $("#btn-search"); + const ready = ds && ds.state === "ready"; + btn.disabled = !ready; + btn.textContent = ready ? "开始检索" : ds ? "模型未就绪" : "请选择数据集"; + renderParamSummary(); +} + +function renderParamSummary() { + const ds = selectedDataset(); + const el = $("#param-summary"); + if (!ds || ds.state !== "ready" || ds.N == null) { + el.innerHTML = `请先选择并构建数据集`; + return; + } + const d = ds.M > 0 ? ds.D / ds.M : 0; + el.innerHTML = ` + 数据规模 N × D${fmtInt(ds.N)} × ${ds.D} + 查询数 Q${fmtInt(ds.Q)} + IVF 簇数 n_list${ds.n_list} + PQ 段数 M / 码字 K${ds.M} / ${ds.K}(d=${d}) + 簇容量上界(padding)${ds.capacity} + 重平衡改动向量数${fmtInt(ds.changed_count)} + 整数化上界 scale_n65536 + `; +} + +/* ================= 检索 ================= */ + +async function runSearch() { + const ds = selectedDataset(); + if (!ds || ds.state !== "ready") { + showToast("请先构建并选择数据集", true); + return; + } + const btn = $("#btn-search"); + btn.disabled = true; + btn.innerHTML = ` 检索中…`; + try { + const body = { + dataset: state.selected, + query_id: parseInt($("#query-id").value || "0", 10), + n_probe: parseInt($("#slider-nprobe").value, 10), + top_k: parseInt($("#slider-topk").value, 10), + proof: $("#chk-proof").checked, + }; + const result = await api("/api/search", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(body), + }); + state.lastResult = result; + renderAll(result); + if (body.proof) showToast("检索与证明生成完成"); + } catch (e) { + showToast(e.message, true); + } finally { + btn.disabled = false; + btn.textContent = "开始检索"; + } +} + +function renderAll(r) { + renderMetrics(r); + drawScatter(r.viz); + drawTimeChart(r); + drawRecallChart(r.metrics); + renderTable(r); + loadClusterStats(r.dataset); +} + +/* ================= 指标卡片 ================= */ + +function metricCard(label, value, sub, colorClass, extra = "") { + return ` +
+
${label}
+
${value}${extra}
+
${sub || ""}
+
`; +} + +function renderMetrics(r) { + const row = $("#metrics-row"); + const p = r.proof; + let proofCards; + if (p) { + proofCards = + metricCard("证明大小", (p.proof_size / 1024).toFixed(1), `KB · ${fmtInt(p.num_gates)} 门`, "c-neutral") + + metricCard( + `证明验证`, + `通过VERIFIED`, + `生成 ${fmtMs(p.prove_time * 1000)} · 验证 ${fmtMs(p.verify_time * 1000)} · 内存 ${(p.memory_used / 1024 ** 3).toFixed(2)} GB`, + "c-ok" + ); + } else { + proofCards = + metricCard("证明大小", "—", "未生成", "c-neutral") + + metricCard("证明验证", "—", "未生成", "c-neutral"); + } + row.innerHTML = + metricCard("标准 IVF-PQ 检索", fmtMs(r.standard.time_ms), `浮点 · n_probe=${r.n_probe}`, "c-std") + + metricCard("ZK IVF-PQ 检索", fmtMs(r.zk.time_ms), `整数域 · 与证明一致`, "c-zk") + + metricCard("标准召回率", (r.metrics.recall_std * 100).toFixed(1) + "%", `recall@${r.top_k} 对比精确检索`, "c-std") + + metricCard("ZK 召回率", (r.metrics.recall_zk * 100).toFixed(1) + "%", `recall@${r.top_k} · 标准∩ZK ${(r.metrics.overlap_std_zk * 100).toFixed(0)}%`, "c-zk") + + metricCard("精确检索耗时", fmtMs(r.metrics.gt_time_ms), `暴力扫描全库`, "c-neutral") + + proofCards; +} + +/* ================= 结果表 ================= */ + +function renderTable(r) { + const tbody = $("#result-table tbody"); + const n = Math.max(r.standard.rows.length, r.zk.rows.length, r.ground_truth.rows.length); + let html = ""; + for (let i = 0; i < n; i++) { + const g = r.ground_truth.rows[i]; + const s = r.standard.rows[i]; + const z = r.zk.rows[i]; + const hitBadge = (hit) => + hit ? `命中` : `未中`; + html += ` + ${i + 1} + ${g ? g.id : ""} + ${g ? g.dist.toFixed(1) : ""} + ${s ? s.id : ""} ${s ? hitBadge(s.in_gt) : ""} + ${s ? s.dist.toFixed(1) : ""} + ${z ? z.id : ""} ${z ? hitBadge(z.in_gt) : ""} + ${z ? fmtInt(z.dist) : ""} + `; + } + tbody.innerHTML = html; +} + +/* ================= 画布工具 ================= */ + +function setupCanvas(canvas) { + const dpr = window.devicePixelRatio || 1; + const w = canvas.clientWidth; + const h = canvas.clientHeight; + canvas.width = Math.max(1, Math.round(w * dpr)); + canvas.height = Math.max(1, Math.round(h * dpr)); + const ctx = canvas.getContext("2d"); + ctx.setTransform(dpr, 0, 0, dpr, 0, 0); + ctx.clearRect(0, 0, w, h); + return { ctx, w, h }; +} + +/* ================= 耗时柱状图(对数) ================= */ + +function drawTimeChart(r) { + const items = [ + { label: "精确检索", value: r.metrics.gt_time_ms, color: "#9ca3af" }, + { label: "标准检索", value: r.standard.time_ms, color: C.std }, + { label: "ZK 检索", value: r.zk.time_ms, color: C.zk }, + ]; + if (r.proof) { + items.push( + { label: "证明构建", value: r.proof.build_time * 1000, color: C.proof }, + { label: "证明生成", value: r.proof.prove_time * 1000, color: C.proof }, + { label: "证明验证", value: r.proof.verify_time * 1000, color: C.verify } + ); + } + const { ctx, w, h } = setupCanvas($("#time-chart")); + const padL = 10, padR = 10, padT = 24, padB = 44; + const plotW = w - padL - padR, plotH = h - padT - padB; + const logs = items.map((it) => Math.log10(Math.max(it.value, 0.01))); + const maxLog = Math.max(...logs, 0.5); + const bw = plotW / items.length; + + ctx.font = "11px sans-serif"; + ctx.textAlign = "center"; + // 网格: 0.01ms ~ 最大 + ctx.strokeStyle = "#f3f4f6"; + for (let e = -2; e <= Math.ceil(maxLog); e++) { + const y = padT + plotH * (1 - (e + 2) / (maxLog + 2)); + if (y < padT || y > padT + plotH) continue; + ctx.beginPath(); + ctx.moveTo(padL, y); + ctx.lineTo(w - padR, y); + ctx.stroke(); + ctx.fillStyle = "#9ca3af"; + ctx.textAlign = "left"; + ctx.fillText(`${Math.pow(10, e) < 1 ? Math.pow(10, e).toFixed(2) : Math.pow(10, e).toFixed(0)}`, padL + 2, y - 2); + ctx.textAlign = "center"; + } + + items.forEach((it, i) => { + const x = padL + i * bw + bw * 0.22; + const barW = bw * 0.56; + const frac = (logs[i] + 2) / (maxLog + 2); + const barH = Math.max(3, plotH * frac); + const y = padT + plotH - barH; + ctx.fillStyle = it.color; + roundRect(ctx, x, y, barW, barH, 4); + ctx.fill(); + ctx.fillStyle = "#111827"; + ctx.fillText(fmtMs(it.value), x + barW / 2, y - 6); + ctx.fillStyle = "#4b5563"; + const label = it.label; + ctx.fillText(label, x + barW / 2, padT + plotH + 16); + }); +} + +function roundRect(ctx, x, y, w, h, r) { + r = Math.min(r, w / 2, h / 2); + ctx.beginPath(); + ctx.moveTo(x + r, y); + ctx.arcTo(x + w, y, x + w, y + h, r); + ctx.arcTo(x + w, y + h, x, y + h, r); + ctx.arcTo(x, y + h, x, y, r); + ctx.arcTo(x, y, x + w, y, r); + ctx.closePath(); +} + +/* ================= 召回率柱状图 ================= */ + +function drawRecallChart(metrics) { + const items = [ + { label: "标准 IVF-PQ", value: metrics.recall_std * 100, color: C.std }, + { label: "ZK IVF-PQ", value: metrics.recall_zk * 100, color: C.zk }, + { label: "标准 ∩ ZK", value: metrics.overlap_std_zk * 100, color: "#64748b" }, + ]; + const { ctx, w, h } = setupCanvas($("#recall-chart")); + const padL = 34, padR = 10, padT = 18, padB = 40; + const plotW = w - padL - padR, plotH = h - padT - padB; + const bw = plotW / items.length; + + ctx.strokeStyle = "#f3f4f6"; + ctx.font = "11px sans-serif"; + for (let v = 0; v <= 100; v += 25) { + const y = padT + plotH * (1 - v / 100); + ctx.beginPath(); + ctx.moveTo(padL, y); + ctx.lineTo(w - padR, y); + ctx.stroke(); + ctx.fillStyle = "#9ca3af"; + ctx.textAlign = "right"; + ctx.fillText(v + "%", padL - 6, y + 4); + } + + items.forEach((it, i) => { + const x = padL + i * bw + bw * 0.22; + const barW = bw * 0.56; + const barH = Math.max(2, plotH * (it.value / 100)); + const y = padT + plotH - barH; + ctx.fillStyle = it.color; + roundRect(ctx, x, y, barW, barH, 4); + ctx.fill(); + ctx.fillStyle = "#111827"; + ctx.textAlign = "center"; + ctx.fillText(it.value.toFixed(1) + "%", x + barW / 2, y - 6); + ctx.fillStyle = "#4b5563"; + ctx.fillText(it.label, x + barW / 2, padT + plotH + 16); + }); +} + +/* ================= 簇规模直方图 ================= */ + +async function loadClusterStats(name) { + try { + const st = await api(`/api/cluster_stats/${name}`); + drawClusterChart(st); + $("#cluster-meta").innerHTML = + `重平衡改动向量数 ${fmtInt(st.changed_count)} · 簇上界 ${st.cluster_bound} · Merkle capacity ${st.capacity}`; + } catch (e) { + $("#cluster-meta").textContent = ""; + } +} + +function drawClusterChart(st) { + const { ctx, w, h } = setupCanvas($("#cluster-chart")); + const all = st.std_sizes.concat(st.zk_sizes); + if (!all.length) return; + const lo = 0, hi = Math.max(...all); + const bins = 30; + const bw = (hi - lo) / bins || 1; + const hist = (sizes) => { + const arr = new Array(bins).fill(0); + for (const s of sizes) { + const b = Math.min(bins - 1, Math.floor((s - lo) / bw)); + arr[b]++; + } + return arr; + }; + const hStd = hist(st.std_sizes); + const hZk = hist(st.zk_sizes); + const maxC = Math.max(...hStd, ...hZk, 1); + + const padL = 34, padR = 10, padT = 14, padB = 34; + const plotW = w - padL - padR, plotH = h - padT - padB; + + ctx.strokeStyle = "#f3f4f6"; + ctx.font = "11px sans-serif"; + for (let v = 0; v <= 4; v++) { + const y = padT + plotH * (1 - v / 4); + ctx.beginPath(); + ctx.moveTo(padL, y); + ctx.lineTo(w - padR, y); + ctx.stroke(); + ctx.fillStyle = "#9ca3af"; + ctx.textAlign = "right"; + ctx.fillText(String(Math.round((maxC * v) / 4)), padL - 6, y + 4); + } + + const binW = plotW / bins; + for (let b = 0; b < bins; b++) { + const x = padL + b * binW; + const bh1 = plotH * (hStd[b] / maxC); + const bh2 = plotH * (hZk[b] / maxC); + ctx.fillStyle = "rgba(37, 99, 235, 0.35)"; + ctx.fillRect(x + 1, padT + plotH - bh1, binW - 2, bh1); + ctx.fillStyle = "rgba(124, 58, 237, 0.45)"; + ctx.fillRect(x + 1 + binW / 2, padT + plotH - bh2, binW / 2 - 2, bh2); + } + + // 簇上界参考线 + const bx = padL + plotW * ((st.cluster_bound - lo) / (hi - lo || 1)); + if (bx >= padL && bx <= w - padR) { + ctx.strokeStyle = "#dc2626"; + ctx.setLineDash([4, 4]); + ctx.beginPath(); + ctx.moveTo(bx, padT); + ctx.lineTo(bx, padT + plotH); + ctx.stroke(); + ctx.setLineDash([]); + ctx.fillStyle = "#dc2626"; + ctx.textAlign = "left"; + ctx.fillText(`簇上界 ${st.cluster_bound}`, bx + 4, padT + 10); + } + + ctx.fillStyle = "#4b5563"; + ctx.textAlign = "center"; + ctx.fillText(`0`, padL, h - 12); + ctx.fillText(String(Math.round(hi)), w - padR, h - 12); + ctx.fillText("簇内向量数", padL + plotW / 2, h - 12); +} + +/* ================= PCA 散点图 ================= */ + +function drawScatter(viz) { + const canvas = $("#scatter"); + const { ctx, w, h } = setupCanvas(canvas); + const base = viz.points || []; + const extras = [viz.std_topk || [], viz.zk_topk || [], viz.gt_topk || [], [viz.query]]; + let minX = Infinity, maxX = -Infinity, minY = Infinity, maxY = -Infinity; + for (const p of base) { + minX = Math.min(minX, p[0]); maxX = Math.max(maxX, p[0]); + minY = Math.min(minY, p[1]); maxY = Math.max(maxY, p[1]); + } + for (const arr of extras) { + for (const p of arr) { + minX = Math.min(minX, p[0]); maxX = Math.max(maxX, p[0]); + minY = Math.min(minY, p[1]); maxY = Math.max(maxY, p[1]); + } + } + if (!isFinite(minX)) return; + const dx = maxX - minX || 1, dy = maxY - minY || 1; + minX -= dx * 0.06; maxX += dx * 0.06; + minY -= dy * 0.06; maxY += dy * 0.06; + const scale = Math.min(w / (maxX - minX), h / (maxY - minY)); + const ox = (w - (maxX - minX) * scale) / 2; + const oy = (h - (maxY - minY) * scale) / 2; + const X = (x) => ox + (x - minX) * scale; + const Y = (y) => h - oy - (y - minY) * scale; + + ctx.font = "11px sans-serif"; + + // 探测簇内的基向量(按簇着色) + for (const [x, y, ci] of base) { + ctx.globalAlpha = 0.65; + ctx.fillStyle = clusterColor(ci); + ctx.beginPath(); + ctx.arc(X(x), Y(y), 2.6, 0, Math.PI * 2); + ctx.fill(); + } + ctx.globalAlpha = 1; + + // 精确检索 top-k:绿色菱形 + ctx.strokeStyle = C.gt; + ctx.lineWidth = 2; + for (const [x, y] of viz.gt_topk || []) { + const cx = X(x), cy = Y(y), r = 5; + ctx.beginPath(); + ctx.moveTo(cx, cy - r); + ctx.lineTo(cx + r, cy); + ctx.lineTo(cx, cy + r); + ctx.lineTo(cx - r, cy); + ctx.closePath(); + ctx.stroke(); + } + + // 标准 top-k:蓝色圆环 + ctx.strokeStyle = C.std; + for (const [x, y] of viz.std_topk || []) { + ctx.beginPath(); + ctx.arc(X(x), Y(y), 7, 0, Math.PI * 2); + ctx.stroke(); + } + + // ZK top-k:紫色叉 + ctx.strokeStyle = C.zk; + ctx.lineWidth = 2; + for (const [x, y] of viz.zk_topk || []) { + const cx = X(x), cy = Y(y), r = 5; + ctx.beginPath(); + ctx.moveTo(cx - r, cy - r); ctx.lineTo(cx + r, cy + r); + ctx.moveTo(cx - r, cy + r); ctx.lineTo(cx + r, cy - r); + ctx.stroke(); + } + + // 查询点:红星 + const [qx, qy] = viz.query; + ctx.fillStyle = C.query; + ctx.strokeStyle = "#fff"; + ctx.lineWidth = 1.5; + drawStar(ctx, X(qx), Y(qy), 7, 5); + ctx.fill(); + ctx.stroke(); + ctx.fillStyle = C.query; + ctx.textAlign = "left"; + ctx.fillText("查询", X(qx) + 10, Y(qy) + 4); + + // 图例 + $("#scatter-legend").innerHTML = ` + 标准 top-k(圆环) + ZK top-k(叉) + 精确 top-k(菱形) + 查询向量 + 彩色点 = 探测簇内基向量(按簇着色) + `; +} + +function drawStar(ctx, cx, cy, outer, inner) { + ctx.beginPath(); + for (let i = 0; i < 10; i++) { + const r = i % 2 === 0 ? outer : inner; + const a = (Math.PI / 5) * i - Math.PI / 2; + const x = cx + r * Math.cos(a); + const y = cy + r * Math.sin(a); + i === 0 ? ctx.moveTo(x, y) : ctx.lineTo(x, y); + } + ctx.closePath(); +} + +/* ================= 初始化与事件 ================= */ + +function initEvents() { + $("#btn-search").addEventListener("click", runSearch); + $("#btn-random").addEventListener("click", () => { + const ds = selectedDataset(); + if (!ds || ds.Q == null) return; + $("#query-id").value = Math.floor(Math.random() * ds.Q); + }); + $("#slider-nprobe").addEventListener("input", (e) => { + $("#val-nprobe").textContent = e.target.value; + }); + $("#slider-topk").addEventListener("input", (e) => { + $("#val-topk").textContent = e.target.value; + }); + let resizeTimer = null; + window.addEventListener("resize", () => { + clearTimeout(resizeTimer); + resizeTimer = setTimeout(() => { + if (state.lastResult) renderAll(state.lastResult); + }, 200); + }); +} + +async function init() { + initEvents(); + try { + await loadDatasets(); + // 默认选中第一个可用数据集(优先已就绪,其次 siftsmall) + const ready = state.datasets.find((d) => d.state === "ready"); + const available = state.datasets.find((d) => d.name === "siftsmall" && d.available) + || state.datasets.find((d) => d.available); + const pick = ready || available; + if (pick) { + state.selected = pick.name; + renderDatasets(); + updateControls(); + if (pick.state !== "ready" && pick.available) prepareDataset(pick.name); + if (pick.state === "ready") { + $("#chk-proof").checked = pick.name === "siftsmall"; + runSearch(); + } + } + } catch (e) { + showToast("无法连接后端: " + e.message, true); + } +} + +init(); diff --git a/demo/static/index.html b/demo/static/index.html new file mode 100644 index 0000000..066967f --- /dev/null +++ b/demo/static/index.html @@ -0,0 +1,126 @@ + + + + + + zk-IVF-PQ 交互演示 + + + +
+
+ + zk-IVF-PQ + 零知识向量数据库检索交互演示 +
+
标准 IVF-PQ · ZK IVF-PQ · 精确检索 三方对比
+
+ +
+ + +
+
+
在左侧选择查询并点击「开始检索」查看结果
+
+ +
+
+

PCA 二维投影 · 探测簇与检索结果

+
+
+
+
+

耗时对比(对数刻度,ms)

+
+
+
+ +
+
+

召回率 @ top_k

+
+
+
+

IVF 簇规模分布(标准 vs ZK 重平衡)

+
+
+
+
+ +
+

检索结果对比

+
+ + + + + + + + + + + + + + + + +
排名精确检索(真值)标准 IVF-PQZK IVF-PQ
向量 ID距离向量 ID距离向量 ID距离
+
+
+
+
+ +
+ + + diff --git a/demo/static/screenshot.png b/demo/static/screenshot.png new file mode 100644 index 0000000..38bcbe6 Binary files /dev/null and b/demo/static/screenshot.png differ diff --git a/demo/static/style.css b/demo/static/style.css new file mode 100644 index 0000000..543a6b2 --- /dev/null +++ b/demo/static/style.css @@ -0,0 +1,278 @@ +:root { + --bg: #ffffff; + --surface: #ffffff; + --surface-2: #f8fafc; + --border: #e5e7eb; + --border-strong: #d1d5db; + --text: #111827; + --text-2: #4b5563; + --muted: #9ca3af; + --accent: #2563eb; + --accent-soft: #eff6ff; + --zk: #7c3aed; + --zk-soft: #f5f3ff; + --ok: #059669; + --ok-soft: #ecfdf5; + --warn: #d97706; + --warn-soft: #fffbeb; + --bad: #dc2626; + --bad-soft: #fef2f2; + --shadow: 0 1px 2px rgba(17, 24, 39, 0.05); + --radius: 12px; +} + +* { box-sizing: border-box; } + +html, body { + margin: 0; + padding: 0; + background: var(--bg); + color: var(--text); + font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", "PingFang SC", + "Hiragino Sans GB", "Microsoft YaHei", "Noto Sans CJK SC", sans-serif; + font-size: 14px; + line-height: 1.55; +} + +/* ---------- 顶栏 ---------- */ +.topbar { + display: flex; + align-items: center; + justify-content: space-between; + padding: 14px 24px; + border-bottom: 1px solid var(--border); + background: var(--bg); + position: sticky; + top: 0; + z-index: 10; +} +.brand { display: flex; align-items: baseline; gap: 10px; } +.logo { + width: 12px; height: 12px; border-radius: 4px; + background: linear-gradient(135deg, var(--accent), var(--zk)); + display: inline-block; transform: translateY(-1px); +} +.title { font-size: 18px; font-weight: 700; letter-spacing: 0.2px; } +.subtitle { color: var(--text-2); font-size: 13px; } +.topbar-note { color: var(--muted); font-size: 12.5px; } + +/* ---------- 布局 ---------- */ +.layout { + display: grid; + grid-template-columns: 320px 1fr; + gap: 16px; + padding: 16px 24px 40px; + max-width: 1560px; + margin: 0 auto; +} +.sidebar { display: flex; flex-direction: column; gap: 16px; } +.main { display: flex; flex-direction: column; gap: 16px; min-width: 0; } +.grid-2 { + display: grid; + grid-template-columns: 1.4fr 1fr; + gap: 16px; +} +@media (max-width: 1100px) { + .layout { grid-template-columns: 1fr; } + .grid-2 { grid-template-columns: 1fr; } +} + +/* ---------- 卡片 ---------- */ +.card { + background: var(--surface); + border: 1px solid var(--border); + border-radius: var(--radius); + padding: 16px 18px; + box-shadow: var(--shadow); +} +.card h2 { + margin: 0 0 12px; + font-size: 14.5px; + font-weight: 600; + color: var(--text); +} + +/* ---------- 数据集卡片 ---------- */ +.dataset-list { display: flex; flex-direction: column; gap: 10px; } +.ds-item { + border: 1px solid var(--border); + border-radius: 10px; + padding: 10px 12px; + cursor: pointer; + transition: border-color 0.15s, background 0.15s; + background: var(--bg); +} +.ds-item:hover { border-color: var(--border-strong); background: var(--surface-2); } +.ds-item.selected { border-color: var(--accent); background: var(--accent-soft); } +.ds-item .ds-head { display: flex; align-items: center; justify-content: space-between; gap: 8px; } +.ds-item .ds-name { font-weight: 600; } +.ds-item .ds-desc { color: var(--text-2); font-size: 12.5px; margin-top: 2px; } +.ds-item .ds-meta { color: var(--muted); font-size: 12px; margin-top: 4px; } +.badge { + display: inline-block; padding: 1px 8px; border-radius: 999px; + font-size: 11.5px; white-space: nowrap; +} +.badge.ready { background: var(--ok-soft); color: var(--ok); } +.badge.building { background: var(--warn-soft); color: var(--warn); } +.badge.idle { background: var(--surface-2); color: var(--text-2); } +.badge.error, .badge.missing { background: var(--bad-soft); color: var(--bad); } +.ds-build-btn { + margin-top: 8px; width: 100%; + border: 1px solid var(--accent); color: var(--accent); + background: var(--bg); border-radius: 8px; padding: 5px 10px; + font-size: 12.5px; cursor: pointer; +} +.ds-build-btn:hover { background: var(--accent-soft); } +.ds-build-btn:disabled { opacity: 0.5; cursor: default; } +.progress-text { margin-top: 6px; color: var(--warn); font-size: 12px; } + +/* ---------- 表单 ---------- */ +.field-label { + display: block; margin: 12px 0 6px; + font-size: 13px; color: var(--text-2); font-weight: 500; +} +.field-label .val { color: var(--accent); font-weight: 600; } +.row { display: flex; gap: 8px; } +input[type="number"] { + width: 100%; padding: 7px 10px; + border: 1px solid var(--border-strong); border-radius: 8px; + font-size: 14px; color: var(--text); background: var(--bg); +} +input[type="number"]:focus { outline: 2px solid var(--accent-soft); border-color: var(--accent); } +input[type="range"] { width: 100%; accent-color: var(--accent); cursor: pointer; } + +.switch-row { + display: flex; align-items: center; gap: 10px; + margin: 16px 0 4px; cursor: pointer; user-select: none; +} +.switch-row input { display: none; } +.switch { + width: 36px; height: 20px; border-radius: 999px; + background: var(--border-strong); position: relative; + transition: background 0.15s; flex-shrink: 0; +} +.switch::after { + content: ""; position: absolute; top: 2px; left: 2px; + width: 16px; height: 16px; border-radius: 50%; + background: #fff; box-shadow: 0 1px 2px rgba(0,0,0,0.2); + transition: left 0.15s; +} +.switch-row input:checked + .switch { background: var(--accent); } +.switch-row input:checked + .switch::after { left: 18px; } + +.btn { + border: none; border-radius: 8px; padding: 9px 14px; + font-size: 14px; font-weight: 600; cursor: pointer; + transition: filter 0.15s, background 0.15s; +} +.btn.primary { width: 100%; margin-top: 14px; background: var(--accent); color: #fff; } +.btn.primary:hover { filter: brightness(1.08); } +.btn.primary:disabled { background: var(--border-strong); cursor: default; } +.btn.ghost { + background: var(--bg); color: var(--text-2); + border: 1px solid var(--border-strong); padding: 7px 12px; +} +.btn.ghost:hover { background: var(--surface-2); } + +.kv { display: grid; grid-template-columns: auto 1fr; gap: 4px 12px; font-size: 13px; } +.kv .k { color: var(--muted); } +.kv .v { color: var(--text); font-weight: 500; text-align: right; } +.muted { color: var(--muted); } + +.note { color: var(--muted); font-size: 12px; margin: 8px 0 0; } +.explain { margin: 0; padding-left: 18px; color: var(--text-2); font-size: 12.5px; } +.explain li { margin: 4px 0; } +.explain b.zk { color: var(--zk); } + +/* ---------- 指标卡片(强制一行 6 块,紧凑不换行) ---------- */ +.metrics-row { + display: flex; + flex-wrap: nowrap; + gap: 8px; + overflow: hidden; +} +.metrics-row .metric { + flex: 1 1 0; + min-width: 0; + background: var(--surface); border: 1px solid var(--border); + border-radius: var(--radius); padding: 9px 10px; box-shadow: var(--shadow); +} +.metric .m-label { color: var(--text-2); font-size: 11px; white-space: nowrap; overflow: hidden; text-overflow: ellipsis; } +.metric .m-value { font-size: 14px; font-weight: 700; margin-top: 2px; white-space: nowrap; overflow: hidden; text-overflow: ellipsis; } +.metric .m-sub { color: var(--muted); font-size: 10px; margin-top: 2px; white-space: nowrap; overflow: hidden; text-overflow: ellipsis; line-height: 1.35; } +.metric .m-sub, .metric .m-label { display: block; } +@media (max-width: 1500px) { + .metrics-row { gap: 6px; } + .metrics-row .metric { padding: 7px 8px; } + .metric .m-value { font-size: 12.5px; } + .metric .m-label { font-size: 10px; } + .metric .m-sub { font-size: 9px; } + .verified-pill { font-size: 9px; padding: 1px 6px; margin-left: 4px; } +} +@media (max-width: 1200px) { + .metrics-row .metric { padding: 6px 7px; } + .metric .m-value { font-size: 11.5px; } + .metric .m-label { font-size: 9.5px; } + .metric .m-sub { font-size: 8.5px; } +} +.metric.c-std .m-value { color: var(--accent); } +.metric.c-zk .m-value { color: var(--zk); } +.metric.c-ok .m-value { color: var(--ok); } +.metric.c-neutral .m-value { color: var(--text); } +.metric-placeholder { grid-column: 1 / -1; text-align: center; padding: 28px; } + +.verified-pill { + display: inline-block; margin-left: 8px; padding: 1px 8px; + border-radius: 999px; background: var(--ok-soft); color: var(--ok); + font-size: 11px; font-weight: 600; vertical-align: middle; +} + +/* ---------- 图表 ---------- */ +.canvas-wrap { width: 100%; height: 300px; position: relative; } +.canvas-wrap canvas { width: 100%; height: 100%; display: block; } +.legend { display: flex; flex-wrap: wrap; gap: 12px; margin-top: 10px; font-size: 12.5px; color: var(--text-2); } +.legend .lg { display: inline-flex; align-items: center; gap: 6px; } +.legend .sw { width: 10px; height: 10px; border-radius: 3px; display: inline-block; } + +/* ---------- 结果表 ---------- */ +.table-scroll { max-height: 420px; overflow: auto; border: 1px solid var(--border); border-radius: 10px; } +table { width: 100%; border-collapse: collapse; font-size: 13px; background: var(--bg); } +thead th { + position: sticky; top: 0; background: var(--surface-2); + color: var(--text-2); font-weight: 600; text-align: center; + padding: 8px 10px; border-bottom: 1px solid var(--border); + z-index: 1; +} +thead tr:first-child th { font-size: 13px; } +thead tr:nth-child(2) th { font-size: 12px; color: var(--muted); } +tbody td { + padding: 6px 10px; border-bottom: 1px solid var(--surface-2); + text-align: center; font-variant-numeric: tabular-nums; +} +tbody tr:hover { background: var(--surface-2); } +td.rank { color: var(--muted); width: 56px; } +.td-gt { background: rgba(5, 150, 105, 0.04); } +.col-std { color: var(--accent); } +.col-zk { color: var(--zk); } +.col-gt { color: var(--ok); } +.hit { display: inline-block; padding: 0 6px; border-radius: 999px; font-size: 11px; } +.hit.yes { background: var(--ok-soft); color: var(--ok); } +.hit.no { background: var(--surface-2); color: var(--muted); } +.dist { color: var(--muted); font-size: 12px; margin-left: 4px; } + +/* ---------- Toast ---------- */ +.toast { + position: fixed; bottom: 24px; left: 50%; transform: translateX(-50%) translateY(20px); + background: #111827; color: #fff; padding: 10px 18px; border-radius: 10px; + font-size: 13px; opacity: 0; pointer-events: none; + transition: opacity 0.2s, transform 0.2s; z-index: 100; max-width: 80vw; +} +.toast.show { opacity: 1; transform: translateX(-50%) translateY(0); } +.toast.error { background: var(--bad); } + +.spinner { + display: inline-block; width: 12px; height: 12px; + border: 2px solid var(--border-strong); border-top-color: var(--accent); + border-radius: 50%; animation: spin 0.8s linear infinite; vertical-align: -2px; +} +@keyframes spin { to { transform: rotate(360deg); } } diff --git a/ivf_pq/util/kmeans.py b/ivf_pq/util/kmeans.py index c2e667b..f3be6cf 100644 --- a/ivf_pq/util/kmeans.py +++ b/ivf_pq/util/kmeans.py @@ -3,6 +3,28 @@ import faiss from typing import Dict, Tuple, Optional +_GPU_USABLE: Optional[bool] = None + + +def _gpu_usable() -> bool: + """ + 检测当前 faiss 构建是否真的能在本机 GPU 上执行 kernel。 + 某些预编译 faiss-gpu 轮子不含新架构(如 sm_120)的 kernel, + StandardGpuResources 能创建但实际计算会崩溃, 因此这里真正跑一次小计算。 + """ + global _GPU_USABLE + if _GPU_USABLE is None: + try: + res = faiss.StandardGpuResources() + probe = np.zeros((4, 4), dtype=np.float32) + index = faiss.GpuIndexFlatL2(res, 4) + index.add(probe) + index.search(probe, 1) + _GPU_USABLE = True + except Exception: + _GPU_USABLE = False + return _GPU_USABLE + def kmeans_with_ids( X: np.ndarray, @@ -51,7 +73,8 @@ def faiss_kmeans_with_ids( # 训练 KMeans # 注意:faiss.extra_wrappers.Kmeans 读取的是 ClusteringParameters.seed, # 需要通过构造参数 seed 传入(而不是 kmeans.seed 这种动态属性赋值)。 - kmeans_kwargs = dict(d=D, k=k, niter=niter, verbose=True, gpu=1) + use_gpu = _gpu_usable() + kmeans_kwargs = dict(d=D, k=k, niter=niter, verbose=True, gpu=use_gpu) if random_state is not None: kmeans_kwargs["seed"] = int(random_state) kmeans = faiss.Kmeans(**kmeans_kwargs) @@ -61,7 +84,8 @@ def faiss_kmeans_with_ids( # kmeans.niter = niter # kmeans.verbose = True kmeans.train(X32) # 得到 centroids - faiss.gpu_sync_all_devices() + if use_gpu and hasattr(faiss, "gpu_sync_all_devices"): + faiss.gpu_sync_all_devices() centers = kmeans.centroids # (k, D), float32 # 用中心建一个Index,把每个样本分到最近的中心(1-NN 到中心) diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..c7be994 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,24 @@ +# Python dependencies for zk-IVF-PQ (tested with Python 3.11) +# +# The Rust extension itself is built with: +# maturin develop --release +# and requires a Rust nightly toolchain (see rust-toolchain.toml). + +maturin>=1.9,<2.0 + +numpy>=2.0 +tqdm +scipy +scikit-learn +matplotlib +duckdb + +# faiss: on machines where a GPU build matching your GPU architecture is +# available you can use faiss-gpu-cu12 instead; ivf_pq/util/kmeans.py +# auto-detects whether GPU kernels actually run and falls back to CPU. +faiss-cpu>=1.15 + +# Only needed for MS MARCO embedding generation (tests/msmacro_emb.py, +# vec_data_load/ms_macro.py): +# torch +# transformers