from __future__ import annotations
import numpy as np
from pyvicinity.ann_benchmarks import VicinityHNSW, VicinityIVFPQ
def run_harness_flow(name: str, algo, train: np.ndarray, test: np.ndarray) -> None:
algo.fit(train)
print(f"{name} fit ok: {algo}")
if isinstance(algo, VicinityIVFPQ):
algo.set_query_arguments(8, rerank_pool=100)
else:
algo.set_query_arguments(50)
ids = algo.query(test[0], 10)
print(f"{name} single-query ids[:10]: {ids.tolist()}")
algo.batch_query(test, 10)
batch = algo.get_batch_results()
print(f"{name} batch_query shape: {batch.shape}")
algo.done()
def main() -> None:
rng = np.random.default_rng(0)
train = rng.standard_normal((5_000, 32), dtype=np.float32)
test = rng.standard_normal((50, 32), dtype=np.float32)
run_harness_flow(
"hnsw", VicinityHNSW("cosine", {"M": 16, "efConstruction": 100}), train, test
)
run_harness_flow(
"ivfpq",
VicinityIVFPQ(
"cosine",
{
"num_clusters": 32,
"num_codebooks": 8,
"codebook_size": 32,
"training_sample_size": 1_000,
"kmeans_max_iter": 5,
"nprobe": 8,
},
),
train,
test,
)
if __name__ == "__main__":
main()