Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions snapvec/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,11 +21,11 @@

__version__ = "0.11.1"
__all__ = [
"SnapIndex",
"PQSnapIndex",
"IVFPQSnapIndex",
"PQSnapIndex",
"ResidualSnapIndex",
"SnapIndex",
"get_codebook",
"rht",
"padded_dim",
"rht",
]
2 changes: 0 additions & 2 deletions snapvec/_fast.pyi
Original file line number Diff line number Diff line change
Expand Up @@ -4,12 +4,10 @@ The real module is built from Cython and does not ship a ``.pyi``
from the compiler; this stub lets ``mypy --strict`` see the same
Python-level shapes the Cython kernels expose to callers.
"""
from __future__ import annotations

import numpy as np
from numpy.typing import NDArray


def adc_colmajor(
lut: NDArray[np.float32],
codes: NDArray[np.uint8],
Expand Down
17 changes: 10 additions & 7 deletions snapvec/_file_format.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,9 +29,13 @@
import os
import struct
import zlib
from collections.abc import Callable
from pathlib import Path
from types import TracebackType
from typing import IO, Callable
from typing import IO, TYPE_CHECKING

if TYPE_CHECKING:
from typing_extensions import Self


_TRAILER_MAGIC = b"CRC2"
Expand Down Expand Up @@ -79,7 +83,7 @@ def finalise(self) -> None:
self._f.write(struct.pack("<I", self._crc & 0xFFFFFFFF))
self._finalised = True

def __enter__(self) -> "ChecksumWriter":
def __enter__(self) -> Self:
return self

def __exit__(
Expand Down Expand Up @@ -163,16 +167,15 @@ def save_with_checksum_atomic(
"""
path = Path(path)
tmp = path.with_suffix(path.suffix + ".tmp")
with open(tmp, "wb") as raw:
with ChecksumWriter(raw) as cw:
writer_fn(cw)
with open(tmp, "wb") as raw, ChecksumWriter(raw) as cw:
writer_fn(cw)
os.replace(tmp, path)


__all__ = [
"ChecksumWriter",
"has_trailer",
"verify_checksum",
"trailer_len",
"save_with_checksum_atomic",
"trailer_len",
"verify_checksum",
]
4 changes: 2 additions & 2 deletions snapvec/_index.py
Original file line number Diff line number Diff line change
Expand Up @@ -581,7 +581,7 @@ def save(self, path: str | Path) -> None:
else:
packed = _pack(self._indices, self._mse_bits)

def _write(f: "ChecksumWriter") -> None:
def _write(f: ChecksumWriter) -> None:
f.write(_MAGIC)
f.write(struct.pack("<IIIIII", _VERSION, self.dim, self.bits, self.seed, n, flags))
f.write(struct.pack("<I", len(packed)))
Expand All @@ -599,7 +599,7 @@ def _write(f: "ChecksumWriter") -> None:
save_with_checksum_atomic(path, _write)

@classmethod
def load(cls, path: str | Path) -> "SnapIndex":
def load(cls, path: str | Path) -> SnapIndex:
"""Load index from a ``.snpv`` file.

Supports v1 (mse-only legacy) and v2 (prod/flags) formats.
Expand Down
4 changes: 2 additions & 2 deletions snapvec/_ivfpq.py
Original file line number Diff line number Diff line change
Expand Up @@ -1127,7 +1127,7 @@ def save(self, path: str | Path) -> None:
flags |= _FLAG_USE_OPQ
n = len(self._ids_by_row)

def _write(f: "ChecksumWriter") -> None:
def _write(f: ChecksumWriter) -> None:
f.write(_MAGIC)
f.write(
struct.pack(
Expand Down Expand Up @@ -1170,7 +1170,7 @@ def _write(f: "ChecksumWriter") -> None:
save_with_checksum_atomic(path, _write)

@classmethod
def load(cls, path: str | Path) -> "IVFPQSnapIndex":
def load(cls, path: str | Path) -> IVFPQSnapIndex:
path = Path(path)
verify_checksum(path) # no-op for legacy files without a trailer
with open(path, "rb") as f:
Expand Down
6 changes: 3 additions & 3 deletions snapvec/_kmeans.py
Original file line number Diff line number Diff line change
Expand Up @@ -199,9 +199,9 @@ def fit_opq_rotation(


__all__ = [
"kmeans_pp_init",
"kmeans_mse",
"assign_l2",
"probe_scores_l2_monotone",
"fit_opq_rotation",
"kmeans_mse",
"kmeans_pp_init",
"probe_scores_l2_monotone",
]
4 changes: 2 additions & 2 deletions snapvec/_pq.py
Original file line number Diff line number Diff line change
Expand Up @@ -426,7 +426,7 @@ def save(self, path: str | Path) -> None:
flags |= _FLAG_USE_OPQ
n = len(self._ids)

def _write(f: "ChecksumWriter") -> None:
def _write(f: ChecksumWriter) -> None:
f.write(_MAGIC)
f.write(
struct.pack(
Expand Down Expand Up @@ -459,7 +459,7 @@ def _write(f: "ChecksumWriter") -> None:
save_with_checksum_atomic(path, _write)

@classmethod
def load(cls, path: str | Path) -> "PQSnapIndex":
def load(cls, path: str | Path) -> PQSnapIndex:
path = Path(path)
verify_checksum(path) # no-op for legacy files without a trailer
with open(path, "rb") as f:
Expand Down
5 changes: 2 additions & 3 deletions snapvec/_residual.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,6 @@
from ._freezable import FreezableIndex
from ._rotation import padded_dim, rht


_MAX_ID_BYTES = 0xFFFF # file format stores id length as uint16


Expand Down Expand Up @@ -295,7 +294,7 @@ def save(self, path: str | Path) -> None:
flags |= 1
n = len(self._ids)

def _write(f: "ChecksumWriter") -> None:
def _write(f: ChecksumWriter) -> None:
f.write(_MAGIC)
f.write(struct.pack("<IIIIIIII", _VERSION, self.dim, self.b1,
self.b2, self.seed, n, flags, self._pdim))
Expand All @@ -318,7 +317,7 @@ def _write(f: "ChecksumWriter") -> None:
save_with_checksum_atomic(path, _write)

@classmethod
def load(cls, path: str | Path) -> "ResidualSnapIndex":
def load(cls, path: str | Path) -> ResidualSnapIndex:
path = Path(path)
verify_checksum(path) # no-op for legacy files without a trailer
with open(path, "rb") as f:
Expand Down
1 change: 0 additions & 1 deletion tests/test_adversarial.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,6 @@

from snapvec import IVFPQSnapIndex, PQSnapIndex, ResidualSnapIndex, SnapIndex


# --------------------------------------------------------------------------- #
# Empty index #
# --------------------------------------------------------------------------- #
Expand Down
4 changes: 2 additions & 2 deletions tests/test_file_format.py
Original file line number Diff line number Diff line change
Expand Up @@ -150,8 +150,8 @@ def test_truncated_trailer_falls_back_to_legacy_mode(tmp_path: Path) -> None:
# ──────────────────────────────────────────────────────────────────── #

@pytest.mark.parametrize("index_cls, ctor_kwargs, suffix", [
(SnapIndex, dict(dim=32, bits=4, normalized=True), ".snpv"),
(ResidualSnapIndex, dict(dim=32, b1=3, b2=3, normalized=True), ".snpr"),
(SnapIndex, {"dim": 32, "bits": 4, "normalized": True}, ".snpv"),
(ResidualSnapIndex, {"dim": 32, "b1": 3, "b2": 3, "normalized": True}, ".snpr"),
])
def test_trailing_crc_roundtrip_trainingfree(
index_cls, ctor_kwargs, suffix, tmp_path,
Expand Down
1 change: 0 additions & 1 deletion tests/test_properties.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,6 @@

from snapvec import IVFPQSnapIndex, PQSnapIndex, SnapIndex


PROFILE = settings(
max_examples=25,
deadline=None,
Expand Down
6 changes: 4 additions & 2 deletions tests/test_snapvec.py
Original file line number Diff line number Diff line change
Expand Up @@ -171,6 +171,7 @@ def test_legacy_v2_3bit_file_loads_via_compat_decoder(self, tmp_path):
path, then re-pack into the new tight RAM layout.
"""
import struct

from snapvec._index import _MAGIC

idx = SnapIndex(dim=128, bits=3)
Expand Down Expand Up @@ -212,7 +213,8 @@ def test_legacy_v2_prod_mode_3bit_payload_stays_aligned(self, tmp_path):
corrupt the prod correction term).
"""
import struct
from snapvec._index import _MAGIC, _FLAG_PROD

from snapvec._index import _FLAG_PROD, _MAGIC

# Real v3 prod-mode index to source the reference indices + payload.
idx = SnapIndex(dim=128, bits=4, use_prod=True)
Expand Down Expand Up @@ -473,7 +475,7 @@ def test_filter_restricts_results(self):
idx = SnapIndex(dim=DIM, bits=4)
idx.add_batch(list(range(100)), vecs)

allowed = set(range(0, 50))
allowed = set(range(50))
results = idx.search(vecs[0], k=10, filter_ids=allowed)
assert all(r[0] in allowed for r in results)

Expand Down
Loading