diff --git a/docs/user-guide/concurrency.md b/docs/user-guide/concurrency.md new file mode 100644 index 0000000..40062b6 --- /dev/null +++ b/docs/user-guide/concurrency.md @@ -0,0 +1,98 @@ +# Concurrency + +`snapvec` indexes are **single-writer, multi-reader** within a single +process. + +## What's safe + +- Multiple threads calling `search()` on the **same** index concurrently, + **after** `freeze()` (see below). Search paths allocate their own + scratch buffers and, once the index is frozen, only read shared state. +- Multiple processes opening **different** index files and querying + them independently. +- A single writer and any number of readers, as long as you never + overlap a writer with a reader on the same instance. + +## Freeze before sharing across threads + +`SnapIndex.search()` lazily materialises an internal float16 centroid +cache on the first query. If two threads both issue their *first* +search concurrently, they race on that cache assignment -- concurrent +`search()` is only safe once the cache exists. + +The library exposes `freeze()` precisely to pre-warm that state: + +```python +idx = SnapIndex(dim=384, bits=4) +idx.add_batch(ids, vectors) +idx.freeze() # pre-warms the cache, makes concurrent search safe + +# Now you can hand idx to multiple reader threads. +``` + +After `freeze()`, mutations (`add_batch`, `delete`) raise, so freeze +also doubles as an "I'm done writing" signal. `PQSnapIndex`, +`IVFPQSnapIndex`, and `ResidualSnapIndex` have the same contract: +call `freeze()` -- or at least issue one warm-up `search()` on the +main thread -- before fanning out. + +## What's not safe + +- Two threads calling `add_batch`, `delete`, or `fit` on the same + index concurrently. There is no internal lock; the library assumes + the caller serializes mutations. +- One thread mutating while another searches. Even when the mutation + looks atomic at the Python level (for example, appending to a list), + internal arrays are resized and re-sorted without coordination. + +## Recommended pattern + +If your application overlaps readers and writers, every public call +must acquire the same lock. A simple wrapper: + +```python +import threading + +class SafeIndex: + def __init__(self, idx): + self._idx = idx + self._lock = threading.Lock() + + def add_batch(self, ids, vectors): + with self._lock: + self._idx.add_batch(ids, vectors) + + def delete(self, id_): + with self._lock: + return self._idx.delete(id_) + + def search(self, query, k=10, **kwargs): + with self._lock: + return self._idx.search(query, k=k, **kwargs) +``` + +If you *never* mutate the index during reads (typical for a +build-then-serve workflow: one `add_batch` at startup, many `search()` +calls forever after), the reader lock can be skipped -- `search()` +only touches immutable shared state (codes, centroids) and +thread-local scratch. If you need higher read concurrency *and* +occasional writes, use a `threading.RLock` plus a read/write wrapper +(for example, the `readerwriterlock` package) at the application +layer. + +## Cross-process access + +The on-disk format is designed for cold reload, not shared access: + +- `save()` writes to `.tmp` then renames, so a concurrent reader + calling `load(path)` either sees the old file or the new one, never + a partial write. +- Nothing prevents two processes from opening the same file and writing + back. If you need multi-process writes, put a file lock (for example, + `fcntl.flock`) around the `save()` call in your application layer. + +## Future work + +Native single-writer protection via an internal `threading.Lock`, and a +delta-buffer mode for low-latency incremental updates, are tracked on +the roadmap for a future release. diff --git a/mkdocs.yml b/mkdocs.yml index 4314760..b0ad368 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -80,6 +80,7 @@ nav: - ResidualSnapIndex: user-guide/residual.md - Save and load: user-guide/save-load.md - Filtered search: user-guide/filter-search.md + - Concurrency: user-guide/concurrency.md - Architecture: architecture.md - Benchmarks: benchmarks.md - API reference: diff --git a/pyproject.toml b/pyproject.toml index d5846cf..b0d855e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -41,6 +41,7 @@ Issues = "https://github.com/stffns/snapvec/issues" dev = [ "pytest", "pytest-cov", + "hypothesis>=6.100", "mypy>=1.8", "ruff", "Cython>=3.0", diff --git a/snapvec/_index.py b/snapvec/_index.py index 0439b57..cb2540b 100644 --- a/snapvec/_index.py +++ b/snapvec/_index.py @@ -406,6 +406,8 @@ def search( materialising the full float16 cache. This trades peak RAM for additional compute — useful when N > ~500k vectors. """ + if k < 1: + raise ValueError(f"k must be >= 1; got {k}") if not self._ids: return [] diff --git a/tests/test_adversarial.py b/tests/test_adversarial.py new file mode 100644 index 0000000..bf71d0e --- /dev/null +++ b/tests/test_adversarial.py @@ -0,0 +1,224 @@ +"""Adversarial edge-case tests. + +Tiny dims, tiny N, degenerate distributions, empty inputs. These are +the cases where off-by-one bugs and implicit shape assumptions tend to +surface. +""" +from __future__ import annotations + +import numpy as np +import pytest + +from snapvec import IVFPQSnapIndex, PQSnapIndex, ResidualSnapIndex, SnapIndex + + +# --------------------------------------------------------------------------- # +# Empty index # +# --------------------------------------------------------------------------- # + + +def test_empty_snapindex_search_returns_empty() -> None: + idx = SnapIndex(dim=8, bits=4, seed=0) + q = np.zeros(8, dtype=np.float32) + q[0] = 1.0 + assert idx.search(q, k=5) == [] + assert len(idx) == 0 + + +def test_empty_pqsnapindex_search_returns_empty() -> None: + vecs = np.random.default_rng(0).standard_normal((64, 8)).astype(np.float32) + idx = PQSnapIndex(dim=8, M=4, K=8, seed=0) + idx.fit(vecs) + # Never call add_batch -- index is fitted but empty. + q = vecs[0] + assert idx.search(q, k=5) == [] + assert len(idx) == 0 + + +# --------------------------------------------------------------------------- # +# Single-vector corpus # +# --------------------------------------------------------------------------- # + + +def test_snapindex_n1() -> None: + """n=1 is a legal if-degenerate corpus. search(k>=1) returns 1 hit.""" + v = np.random.default_rng(0).standard_normal((1, 16)).astype(np.float32) + idx = SnapIndex(dim=16, bits=4, seed=0) + idx.add_batch(["only"], v) + hits = idx.search(v[0], k=5) + assert len(hits) == 1 + assert hits[0][0] == "only" + + +# --------------------------------------------------------------------------- # +# k larger than n # +# --------------------------------------------------------------------------- # + + +def test_search_k_larger_than_n_returns_n() -> None: + vecs = np.random.default_rng(0).standard_normal((3, 16)).astype(np.float32) + idx = SnapIndex(dim=16, bits=4, seed=0) + idx.add_batch([0, 1, 2], vecs) + hits = idx.search(vecs[0], k=100) + assert len(hits) == 3 + + +# --------------------------------------------------------------------------- # +# Zero-norm inputs # +# --------------------------------------------------------------------------- # + + +def test_search_with_zero_query_returns_empty() -> None: + """Zero-norm query can't be normalized; library returns [] instead of NaN hits.""" + vecs = np.random.default_rng(0).standard_normal((20, 16)).astype(np.float32) + idx = SnapIndex(dim=16, bits=4, seed=0) + idx.add_batch(list(range(20)), vecs) + q_zero = np.zeros(16, dtype=np.float32) + assert idx.search(q_zero, k=5) == [] + + +# --------------------------------------------------------------------------- # +# All-same-vector corpus (degenerate clusters) # +# --------------------------------------------------------------------------- # + + +def test_snapindex_all_same_vector() -> None: + """Every vector identical -> search should still return k distinct ids.""" + v = np.ones((1, 32), dtype=np.float32) + vecs = np.tile(v, (10, 1)) + idx = SnapIndex(dim=32, bits=4, seed=0) + idx.add_batch(list(range(10)), vecs) + hits = idx.search(v[0], k=5) + assert len(hits) == 5 + ids = [h[0] for h in hits] + assert len(set(ids)) == len(ids) # distinct + + +# --------------------------------------------------------------------------- # +# Filter edge cases # +# --------------------------------------------------------------------------- # + + +def test_filter_with_only_unknown_ids_returns_empty() -> None: + vecs = np.random.default_rng(0).standard_normal((50, 16)).astype(np.float32) + idx = SnapIndex(dim=16, bits=4, seed=0) + idx.add_batch(list(range(50)), vecs) + hits = idx.search(vecs[0], k=5, filter_ids={"never-added", "also-never"}) + assert hits == [] + + +def test_filter_with_empty_set_returns_empty() -> None: + vecs = np.random.default_rng(0).standard_normal((50, 16)).astype(np.float32) + idx = SnapIndex(dim=16, bits=4, seed=0) + idx.add_batch(list(range(50)), vecs) + hits = idx.search(vecs[0], k=5, filter_ids=set()) + assert hits == [] + + +# --------------------------------------------------------------------------- # +# Delete-all # +# --------------------------------------------------------------------------- # + + +def test_snapindex_delete_all_then_search() -> None: + vecs = np.random.default_rng(0).standard_normal((5, 16)).astype(np.float32) + idx = SnapIndex(dim=16, bits=4, seed=0) + idx.add_batch(list(range(5)), vecs) + for i in range(5): + assert idx.delete(i) is True + assert len(idx) == 0 + assert idx.search(vecs[0], k=3) == [] + + +# --------------------------------------------------------------------------- # +# Aggressive compression (bits=2) # +# --------------------------------------------------------------------------- # + + +def test_snapindex_bits2_basic_recall() -> None: + """bits=2 still returns *something* sensible on clustered data.""" + rng = np.random.default_rng(0) + centers = rng.standard_normal((5, 32)).astype(np.float32) * 3 + assign = rng.integers(0, 5, size=200) + jitter = rng.standard_normal((200, 32)).astype(np.float32) * 0.2 + corpus = centers[assign] + jitter + + idx = SnapIndex(dim=32, bits=2, seed=0) + idx.add_batch(list(range(200)), corpus) + + # Self-query should rank the exact corpus row near the top on such + # strongly clustered data. + hits = idx.search(corpus[0], k=5) + assert len(hits) == 5 + returned = [h[0] for h in hits] + # Because clusters have ~40 members and bits=2 is aggressive, + # we don't assert hits[0] == 0. We only assert the top-5 are all + # from the same cluster as the query. + query_cluster = assign[0] + top_clusters = [assign[i] for i in returned] + assert top_clusters.count(query_cluster) >= 3 + + +# --------------------------------------------------------------------------- # +# k=0 is an error # +# --------------------------------------------------------------------------- # + + +def test_search_k_zero_raises() -> None: + vecs = np.random.default_rng(0).standard_normal((10, 16)).astype(np.float32) + idx = SnapIndex(dim=16, bits=4, seed=0) + idx.add_batch(list(range(10)), vecs) + with pytest.raises(ValueError): + idx.search(vecs[0], k=0) + + +# --------------------------------------------------------------------------- # +# IVF-PQ extreme nprobe # +# --------------------------------------------------------------------------- # + + +def test_ivfpq_nprobe_equals_nlist_is_full_scan() -> None: + """With nprobe=nlist, IVF-PQ must visit every cluster.""" + rng = np.random.default_rng(0) + corpus = rng.standard_normal((300, 32)).astype(np.float32) + idx = IVFPQSnapIndex(dim=32, nlist=8, M=4, K=16, seed=0) + idx.fit(corpus) + idx.add_batch(list(range(300)), corpus) + + hits = idx.search(corpus[0], k=10, nprobe=8) + assert len(hits) == 10 + + +def test_ivfpq_nprobe_out_of_range_raises() -> None: + rng = np.random.default_rng(0) + corpus = rng.standard_normal((300, 32)).astype(np.float32) + idx = IVFPQSnapIndex(dim=32, nlist=8, M=4, K=16, seed=0) + idx.fit(corpus) + idx.add_batch(list(range(300)), corpus) + with pytest.raises(ValueError): + idx.search(corpus[0], k=5, nprobe=0) + with pytest.raises(ValueError): + idx.search(corpus[0], k=5, nprobe=99) + + +# --------------------------------------------------------------------------- # +# Residual rerank # +# --------------------------------------------------------------------------- # + + +def test_residual_rerank_saturates_near_full_recall() -> None: + """ResidualSnapIndex with a generous rerank_M should match full scan on + clustered data.""" + rng = np.random.default_rng(0) + centers = rng.standard_normal((8, 32)).astype(np.float32) * 3 + assign = rng.integers(0, 8, size=200) + jitter = rng.standard_normal((200, 32)).astype(np.float32) * 0.2 + corpus = centers[assign] + jitter + + idx = ResidualSnapIndex(dim=32, b1=3, b2=3, seed=0) + idx.add_batch(list(range(200)), corpus) + + full = [h[0] for h in idx.search(corpus[0], k=5, rerank_M=None)] + reranked = [h[0] for h in idx.search(corpus[0], k=5, rerank_M=50)] + # Both modes should agree on most of the top-5. + assert len(set(full) & set(reranked)) >= 3 diff --git a/tests/test_determinism.py b/tests/test_determinism.py new file mode 100644 index 0000000..5aac719 --- /dev/null +++ b/tests/test_determinism.py @@ -0,0 +1,127 @@ +"""Determinism tests. + +Building an index twice with the same seed and the same inputs must +produce bit-identical outputs. If this ever breaks, it is almost +always a non-deterministic code path that will also cause subtle +recall drift between runs. +""" +from __future__ import annotations + +import hashlib +from pathlib import Path + +import numpy as np +import pytest + +from snapvec import IVFPQSnapIndex, PQSnapIndex, ResidualSnapIndex, SnapIndex + + +def _sha256(path: Path) -> str: + return hashlib.sha256(path.read_bytes()).hexdigest() + + +def _clustered(n: int, dim: int, *, seed: int, n_clusters: int = 8) -> np.ndarray: + rng = np.random.default_rng(seed) + centers = rng.standard_normal((n_clusters, dim)).astype(np.float32) * 3 + assign = rng.integers(0, n_clusters, size=n) + jitter = rng.standard_normal((n, dim)).astype(np.float32) * 0.3 + return centers[assign] + jitter + + +def test_snapindex_save_is_bitwise_deterministic(tmp_path: Path) -> None: + """Same seed + same inputs => byte-identical .snpv file.""" + vecs = _clustered(200, 64, seed=42) + + def build(path: Path) -> None: + idx = SnapIndex(dim=64, bits=4, seed=7) + idx.add_batch(list(range(200)), vecs) + idx.save(path) + + a = tmp_path / "a.snpv" + b = tmp_path / "b.snpv" + build(a) + build(b) + assert _sha256(a) == _sha256(b) + + +def test_pqsnapindex_save_is_bitwise_deterministic(tmp_path: Path) -> None: + """PQSnapIndex.fit + add_batch with the same seed => same bytes.""" + vecs = _clustered(500, 64, seed=1) + + def build(path: Path) -> None: + idx = PQSnapIndex(dim=64, M=8, K=16, seed=3) + idx.fit(vecs) + idx.add_batch(list(range(500)), vecs) + idx.save(path) + + a = tmp_path / "a.snpq" + b = tmp_path / "b.snpq" + build(a) + build(b) + assert _sha256(a) == _sha256(b) + + +def test_residualsnapindex_save_is_bitwise_deterministic(tmp_path: Path) -> None: + vecs = _clustered(200, 64, seed=9) + + def build(path: Path) -> None: + idx = ResidualSnapIndex(dim=64, b1=3, b2=3, seed=11) + idx.add_batch(list(range(200)), vecs) + idx.save(path) + + a = tmp_path / "a.snpr" + b = tmp_path / "b.snpr" + build(a) + build(b) + assert _sha256(a) == _sha256(b) + + +def test_ivfpqsnapindex_save_is_bitwise_deterministic(tmp_path: Path) -> None: + """IVF-PQ fit involves kmeans; same seed must still produce same bytes.""" + vecs = _clustered(500, 64, seed=21) + + def build(path: Path) -> None: + idx = IVFPQSnapIndex(dim=64, nlist=8, M=8, K=16, seed=13) + idx.fit(vecs) + idx.add_batch(list(range(500)), vecs) + idx.save(path) + + a = tmp_path / "a.snpi" + b = tmp_path / "b.snpi" + build(a) + build(b) + assert _sha256(a) == _sha256(b) + + +@pytest.mark.parametrize( + ("index_cls", "ext", "build_kwargs"), + [ + (SnapIndex, ".snpv", {"bits": 4}), + (PQSnapIndex, ".snpq", {"M": 8, "K": 16}), + (ResidualSnapIndex, ".snpr", {"b1": 3, "b2": 3}), + (IVFPQSnapIndex, ".snpi", {"nlist": 8, "M": 8, "K": 16}), + ], +) +def test_search_results_are_deterministic( + tmp_path: Path, + index_cls: type, + ext: str, + build_kwargs: dict[str, object], +) -> None: + """search() returns identical (id, score) tuples across two builds.""" + vecs = _clustered(300, 32, seed=5) + + def build() -> object: + idx = index_cls(dim=32, seed=17, **build_kwargs) + if hasattr(idx, "fit"): + idx.fit(vecs) + idx.add_batch(list(range(300)), vecs) + return idx + + idx_a = build() + idx_b = build() + + query = vecs[7] + hits_a = idx_a.search(query, k=5) + hits_b = idx_b.search(query, k=5) + assert hits_a == hits_b diff --git a/tests/test_properties.py b/tests/test_properties.py new file mode 100644 index 0000000..1237e77 --- /dev/null +++ b/tests/test_properties.py @@ -0,0 +1,277 @@ +"""Property-based tests. + +These exercise invariants that should hold for *any* input within a +modest range. They complement the hand-written unit tests by catching +regressions on inputs nobody thought to write down. +""" +from __future__ import annotations + +from pathlib import Path + +import numpy as np +import pytest +from hypothesis import HealthCheck, given, settings +from hypothesis import strategies as st +from numpy.typing import NDArray + +from snapvec import IVFPQSnapIndex, PQSnapIndex, SnapIndex + + +PROFILE = settings( + max_examples=25, + deadline=None, + suppress_health_check=[HealthCheck.too_slow], +) + + +def _corpus(n: int, dim: int, seed: int) -> NDArray[np.float32]: + """Clustered corpus so PQ / IVF-PQ tests exercise real structure.""" + rng = np.random.default_rng(seed) + n_clusters = max(2, n // 10) + centers = rng.standard_normal((n_clusters, dim)).astype(np.float32) * 3 + assign = rng.integers(0, n_clusters, size=n) + jitter = rng.standard_normal((n, dim)).astype(np.float32) * 0.3 + return centers[assign] + jitter + + +# --------------------------------------------------------------------------- # +# SnapIndex invariants # +# --------------------------------------------------------------------------- # + + +@PROFILE +@given( + dim=st.sampled_from([8, 16, 32, 64]), + n=st.integers(min_value=1, max_value=100), + seed=st.integers(min_value=0, max_value=2**16), + bits=st.sampled_from([2, 3, 4]), +) +def test_snap_add_then_len_matches(dim: int, n: int, seed: int, bits: int) -> None: + """len(idx) equals the number of distinct ids added.""" + vecs = _corpus(n, dim, seed) + idx = SnapIndex(dim=dim, bits=bits, seed=0) + idx.add_batch(list(range(n)), vecs) + assert len(idx) == n + + +@PROFILE +@given( + dim=st.sampled_from([8, 16, 32]), + n=st.integers(min_value=2, max_value=50), + seed=st.integers(min_value=0, max_value=2**16), +) +def test_snap_search_returns_at_most_k(dim: int, n: int, seed: int) -> None: + """search(q, k) returns at most k hits, sorted by descending score.""" + vecs = _corpus(n, dim, seed) + idx = SnapIndex(dim=dim, bits=4, seed=0) + idx.add_batch(list(range(n)), vecs) + + for k in (1, min(10, n), n + 5): + hits = idx.search(vecs[0], k=k) + assert len(hits) <= k + assert len(hits) <= n + scores = [h[1] for h in hits] + assert scores == sorted(scores, reverse=True) + + +@PROFILE +@given( + dim=st.sampled_from([8, 16, 32]), + n=st.integers(min_value=3, max_value=50), + seed=st.integers(min_value=0, max_value=2**16), +) +def test_snap_delete_reduces_len(dim: int, n: int, seed: int) -> None: + """Deleting an existing id reduces len by exactly 1.""" + vecs = _corpus(n, dim, seed) + idx = SnapIndex(dim=dim, bits=4, seed=0) + idx.add_batch(list(range(n)), vecs) + + removed = idx.delete(0) + assert removed is True + assert len(idx) == n - 1 + + # Deleting a non-existent id is a no-op. + assert idx.delete(10**9) is False + assert len(idx) == n - 1 + + +@PROFILE +@given( + dim=st.sampled_from([8, 16, 32]), + n=st.integers(min_value=2, max_value=50), + seed=st.integers(min_value=0, max_value=2**16), +) +def test_snap_save_load_preserves_search( + dim: int, n: int, seed: int, tmp_path_factory: pytest.TempPathFactory +) -> None: + """Round-trip through save/load returns bit-identical search results.""" + vecs = _corpus(n, dim, seed) + idx = SnapIndex(dim=dim, bits=4, seed=0) + idx.add_batch(list(range(n)), vecs) + + path: Path = tmp_path_factory.mktemp("prop") / "idx.snpv" + idx.save(path) + loaded = SnapIndex.load(path) + + before = idx.search(vecs[0], k=min(5, n)) + after = loaded.search(vecs[0], k=min(5, n)) + assert before == after + + +@PROFILE +@given( + dim=st.sampled_from([8, 16]), + n=st.integers(min_value=5, max_value=40), + seed=st.integers(min_value=0, max_value=2**16), +) +def test_snap_filter_hits_are_in_filter_set(dim: int, n: int, seed: int) -> None: + """With filter_ids=S, every returned hit is in S.""" + vecs = _corpus(n, dim, seed) + idx = SnapIndex(dim=dim, bits=4, seed=0) + idx.add_batch(list(range(n)), vecs) + + # Pick a sparse filter: first third of the ids. + filter_set = set(range(max(1, n // 3))) + hits = idx.search(vecs[0], k=min(5, n), filter_ids=filter_set) + for doc_id, _score in hits: + assert doc_id in filter_set + + +# --------------------------------------------------------------------------- # +# PQSnapIndex invariants # +# --------------------------------------------------------------------------- # + + +@PROFILE +@given( + dim_pair=st.sampled_from([(16, 4), (32, 8), (64, 8)]), + n=st.integers(min_value=20, max_value=100), + seed=st.integers(min_value=0, max_value=2**16), +) +def test_pq_fit_then_add_then_len( + dim_pair: tuple[int, int], n: int, seed: int +) -> None: + """PQSnapIndex after fit() + add_batch reports the right len().""" + dim, M = dim_pair + vecs = _corpus(n, dim, seed) + idx = PQSnapIndex(dim=dim, M=M, K=16, seed=0) + idx.fit(vecs) + idx.add_batch(list(range(n)), vecs) + assert len(idx) == n + + +@PROFILE +@given( + dim_pair=st.sampled_from([(16, 4), (32, 8)]), + n=st.integers(min_value=20, max_value=80), + seed=st.integers(min_value=0, max_value=2**16), +) +def test_pq_save_load_preserves_search( + dim_pair: tuple[int, int], + n: int, + seed: int, + tmp_path_factory: pytest.TempPathFactory, +) -> None: + """PQSnapIndex save/load preserves search output.""" + dim, M = dim_pair + vecs = _corpus(n, dim, seed) + idx = PQSnapIndex(dim=dim, M=M, K=16, seed=0) + idx.fit(vecs) + idx.add_batch(list(range(n)), vecs) + + path: Path = tmp_path_factory.mktemp("prop") / "idx.snpq" + idx.save(path) + loaded = PQSnapIndex.load(path) + + before = idx.search(vecs[0], k=min(5, n)) + after = loaded.search(vecs[0], k=min(5, n)) + assert before == after + + +@PROFILE +@given( + dim_pair=st.sampled_from([(16, 4), (32, 8)]), + n=st.integers(min_value=20, max_value=80), + seed=st.integers(min_value=0, max_value=2**16), +) +def test_pq_delete_reduces_len( + dim_pair: tuple[int, int], n: int, seed: int +) -> None: + """Deleting an existing id in PQSnapIndex reduces len by exactly 1.""" + dim, M = dim_pair + vecs = _corpus(n, dim, seed) + idx = PQSnapIndex(dim=dim, M=M, K=16, seed=0) + idx.fit(vecs) + idx.add_batch(list(range(n)), vecs) + + assert idx.delete(0) is True + assert len(idx) == n - 1 + assert idx.delete(10**9) is False + assert len(idx) == n - 1 + + +# --------------------------------------------------------------------------- # +# IVFPQSnapIndex invariants # +# --------------------------------------------------------------------------- # + + +@PROFILE +@given( + n=st.integers(min_value=100, max_value=300), + nprobe=st.sampled_from([1, 2, 4, 8]), + seed=st.integers(min_value=0, max_value=2**16), +) +def test_ivfpq_search_respects_k(n: int, nprobe: int, seed: int) -> None: + """IVFPQSnapIndex returns <= k results, sorted by score descending.""" + dim, M, K, nlist = 16, 4, 16, 8 + vecs = _corpus(n, dim, seed) + idx = IVFPQSnapIndex(dim=dim, nlist=nlist, M=M, K=K, seed=0) + idx.fit(vecs) + idx.add_batch(list(range(n)), vecs) + + for k in (1, 5, 20): + hits = idx.search(vecs[0], k=k, nprobe=nprobe) + assert len(hits) <= k + scores = [h[1] for h in hits] + assert scores == sorted(scores, reverse=True) + + +@PROFILE +@given( + n=st.integers(min_value=100, max_value=200), + seed=st.integers(min_value=0, max_value=2**16), +) +def test_ivfpq_unknown_filter_returns_empty(n: int, seed: int) -> None: + """filter_ids with no matching id returns [].""" + dim, M, K, nlist = 16, 4, 16, 8 + vecs = _corpus(n, dim, seed) + idx = IVFPQSnapIndex(dim=dim, nlist=nlist, M=M, K=K, seed=0) + idx.fit(vecs) + idx.add_batch(list(range(n)), vecs) + + hits = idx.search(vecs[0], k=5, filter_ids={"nope-nope-nope"}) + assert hits == [] + + +@PROFILE +@given( + n=st.integers(min_value=100, max_value=200), + seed=st.integers(min_value=0, max_value=2**16), +) +def test_ivfpq_save_load_preserves_search( + n: int, seed: int, tmp_path_factory: pytest.TempPathFactory +) -> None: + """IVFPQSnapIndex save/load preserves search output.""" + dim, M, K, nlist = 16, 4, 16, 8 + vecs = _corpus(n, dim, seed) + idx = IVFPQSnapIndex(dim=dim, nlist=nlist, M=M, K=K, seed=0) + idx.fit(vecs) + idx.add_batch(list(range(n)), vecs) + + path: Path = tmp_path_factory.mktemp("prop") / "idx.snpi" + idx.save(path) + loaded = IVFPQSnapIndex.load(path) + + before = idx.search(vecs[0], k=5, nprobe=4) + after = loaded.search(vecs[0], k=5, nprobe=4) + assert before == after