use super::cluster::Cluster;
use super::file_storage::{
append_codes_for_ids, build_list_codes, checked_len, open_byte_storage, open_list_code_storage,
read_code_from_storage, read_list_codes_for_cluster, read_vector_from_storage,
IVFPQByteStorage, IVFPQListCodeStorage,
};
use super::manifest::{IVFPQManifest, PersistedFilterMetadata, PersistedIVFPQParams};
use super::opq::OptimizedProductQuantizer;
use super::persistence::{
read_bytes_exact, read_clusters, read_f32_exact, read_json, read_u32_exact, validate_manifest,
write_bytes_atomic, write_clusters_atomic, write_f32_atomic, write_json_atomic,
write_u32_atomic, write_u64_atomic, IVFPQ_FORMAT_VERSION,
};
use super::pq::ProductQuantizer;
use crate::pq_simd::{adc_batch_dispatch_into, PackedCodes4bit, PackedLUTRef};
use crate::RetrieveError;
use rand::seq::SliceRandom;
use rand::SeedableRng;
use serde::{Deserialize, Serialize};
use std::path::Path;
#[cfg(feature = "benchmark")]
use std::time::{Duration, Instant};
const SIMD_BATCH_THRESHOLD: usize = 16;
fn default_kmeans_max_iter() -> usize {
100
}
#[derive(Clone, Copy, Debug)]
struct IVFPQTrainingConfig {
sample_size: Option<usize>,
kmeans_max_iter: usize,
}
impl Default for IVFPQTrainingConfig {
fn default() -> Self {
Self {
sample_size: None,
kmeans_max_iter: default_kmeans_max_iter(),
}
}
}
fn normalize_query(query: &[f32]) -> Vec<f32> {
let norm: f32 = query.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 1e-10 {
query.iter().map(|x| x / norm).collect()
} else {
query.to_vec()
}
}
fn finish_top_k_by_distance(mut candidates: Vec<(u32, f32)>, k: usize) -> Vec<(u32, f32)> {
if k == 0 || candidates.is_empty() {
return Vec::new();
}
if candidates.len() > k {
candidates.select_nth_unstable_by(k - 1, |a, b| a.1.total_cmp(&b.1));
candidates.truncate(k);
}
candidates.sort_unstable_by(|a, b| a.1.total_cmp(&b.1));
candidates
}
fn copy_rows_by_index(vectors: &[f32], dimension: usize, indices: &[usize]) -> Vec<f32> {
let mut out = Vec::with_capacity(indices.len() * dimension);
for &idx in indices {
let start = idx * dimension;
out.extend_from_slice(&vectors[start..start + dimension]);
}
out
}
#[cfg(feature = "benchmark")]
#[derive(Clone, Debug, Default)]
pub struct IVFPQSearchProfile {
pub normalize: Duration,
pub centroid_lookup: Duration,
pub residual: Duration,
pub adc_table: Duration,
pub code_copy: Duration,
pub adc_dispatch: Duration,
pub finalizer: Duration,
pub probed_clusters: usize,
pub simd_clusters: usize,
pub scalar_clusters: usize,
pub scanned_vectors: usize,
pub candidate_count: usize,
pub returned_count: usize,
pub code_copy_bytes: usize,
}
#[cfg(feature = "benchmark")]
impl IVFPQSearchProfile {
pub fn add_assign(&mut self, other: &Self) {
self.normalize += other.normalize;
self.centroid_lookup += other.centroid_lookup;
self.residual += other.residual;
self.adc_table += other.adc_table;
self.code_copy += other.code_copy;
self.adc_dispatch += other.adc_dispatch;
self.finalizer += other.finalizer;
self.probed_clusters += other.probed_clusters;
self.simd_clusters += other.simd_clusters;
self.scalar_clusters += other.scalar_clusters;
self.scanned_vectors += other.scanned_vectors;
self.candidate_count += other.candidate_count;
self.returned_count += other.returned_count;
self.code_copy_bytes += other.code_copy_bytes;
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum Quantizer {
Product(ProductQuantizer),
Optimized(OptimizedProductQuantizer),
}
impl Quantizer {
pub fn quantize(&self, vector: &[f32]) -> Vec<u8> {
match self {
Self::Product(pq) => pq.quantize(vector),
Self::Optimized(opq) => opq.quantize(vector),
}
}
pub fn compute_adc_table(&self, query: &[f32]) -> Result<Vec<f32>, RetrieveError> {
match self {
Self::Product(pq) => pq.compute_adc_table(query),
Self::Optimized(opq) => opq.approximate_distance_table(query),
}
}
pub fn compute_adc_table_into(
&self,
query: &[f32],
table: &mut Vec<f32>,
) -> Result<(), RetrieveError> {
match self {
Self::Product(pq) => pq.compute_adc_table_into(query, table),
Self::Optimized(opq) => {
let t = opq.approximate_distance_table(query)?;
table.clear();
table.extend_from_slice(&t);
Ok(())
}
}
}
pub fn distance_with_table(&self, table: &[f32], codes: &[u8]) -> f32 {
match self {
Self::Product(pq) => pq.distance_with_table(table, codes),
Self::Optimized(opq) => opq.distance_with_table(table, codes),
}
}
fn owned_bytes(&self) -> usize {
match self {
Self::Product(pq) => pq.owned_bytes(),
Self::Optimized(opq) => opq.owned_bytes(),
}
}
}
#[derive(Debug)]
pub struct IVFPQIndex {
pub(crate) vectors: Vec<f32>,
pub(crate) dimension: usize,
pub(crate) num_vectors: usize,
doc_ids: Vec<u32>,
params: IVFPQParams,
built: bool,
clusters: Vec<Cluster>,
pub(crate) centroids: Vec<f32>,
pq: Option<Quantizer>,
pub(crate) quantized_codes: Vec<u8>,
metadata: Option<crate::filtering::MetadataStore>,
filter_field: Option<String>,
#[cfg(feature = "hnsw")]
coarse_quantizer: Option<crate::hnsw::HNSWIndex>,
}
pub struct IVFPQFileSearcher {
dimension: usize,
num_vectors: usize,
doc_ids: Vec<u32>,
params: IVFPQParams,
clusters: Vec<Cluster>,
centroids: Vec<f32>,
pq: Quantizer,
codes: IVFPQByteStorage,
list_codes: Option<IVFPQListCodeStorage>,
raw_vectors: Option<IVFPQByteStorage>,
code_buf: Vec<u8>,
raw_byte_buf: Vec<u8>,
vec_buf: Vec<f32>,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct IVFPQFileSearchDiagnostics {
pub probed_lists: usize,
pub scanned_vectors: usize,
pub code_reads: usize,
pub code_bytes: usize,
pub retained_candidates: usize,
pub raw_vector_reads: usize,
pub raw_vector_bytes: usize,
pub reranked_candidates: usize,
}
#[derive(Clone, Debug)]
pub struct IVFPQParams {
pub num_clusters: usize,
pub nprobe: usize,
pub num_codebooks: usize,
pub codebook_size: usize,
pub use_opq: bool,
pub seed: u64,
#[cfg(feature = "id-compression")]
pub id_compression: Option<crate::compression::IdCompressionMethod>,
#[cfg(feature = "id-compression")]
pub compression_threshold: usize,
}
impl Default for IVFPQParams {
fn default() -> Self {
Self {
num_clusters: 1024,
nprobe: 100,
num_codebooks: 8,
codebook_size: 256,
use_opq: false,
seed: 42,
#[cfg(feature = "id-compression")]
id_compression: None,
#[cfg(feature = "id-compression")]
compression_threshold: 100, }
}
}
impl IVFPQIndex {
pub fn set_nprobe(&mut self, nprobe: usize) {
self.params.nprobe = nprobe;
}
pub fn nprobe(&self) -> usize {
self.params.nprobe
}
pub fn num_clusters(&self) -> usize {
self.params.num_clusters
}
pub fn num_codebooks(&self) -> usize {
self.params.num_codebooks
}
pub fn codebook_size(&self) -> usize {
self.params.codebook_size
}
pub fn use_opq(&self) -> bool {
self.params.use_opq
}
pub fn is_built(&self) -> bool {
self.built
}
pub fn new(dimension: usize, params: IVFPQParams) -> Result<Self, RetrieveError> {
if dimension == 0 {
return Err(RetrieveError::InvalidParameter(
"dimension must be > 0".into(),
));
}
Ok(Self {
vectors: Vec::new(),
dimension,
num_vectors: 0,
doc_ids: Vec::new(),
params,
built: false,
clusters: Vec::new(),
centroids: Vec::new(),
pq: None,
quantized_codes: Vec::new(),
metadata: None,
filter_field: None,
#[cfg(feature = "hnsw")]
coarse_quantizer: None,
})
}
pub fn with_filtering(
dimension: usize,
params: IVFPQParams,
filter_field: impl Into<String>,
) -> Result<Self, RetrieveError> {
Ok(Self {
vectors: Vec::new(),
dimension,
num_vectors: 0,
doc_ids: Vec::new(),
params,
built: false,
clusters: Vec::new(),
centroids: Vec::new(),
pq: None,
quantized_codes: Vec::new(),
metadata: Some(crate::filtering::MetadataStore::new()),
filter_field: Some(filter_field.into()),
#[cfg(feature = "hnsw")]
coarse_quantizer: None,
})
}
pub fn add_metadata(
&mut self,
doc_id: u32,
metadata: crate::filtering::DocumentMetadata,
) -> Result<(), RetrieveError> {
if let Some(ref mut store) = self.metadata {
if let Some(ref field) = self.filter_field {
if let Some(category_val) = metadata.get(field) {
match category_val {
crate::filtering::MetadataValue::Int(n) if *n >= 0 && *n < 64 => {}
crate::filtering::MetadataValue::Int(n) => {
return Err(RetrieveError::InvalidParameter(format!(
"category ID {} exceeds bitmask limit of 63; \
use an integer in 0..63",
n
)));
}
_ => {
return Err(RetrieveError::InvalidParameter(
"category ID must be an integer in 0..63 for bitmask filtering"
.into(),
));
}
}
}
}
store.add(doc_id, metadata);
Ok(())
} else {
Err(RetrieveError::InvalidParameter(
"filtering not enabled; use IVFPQIndex::with_filtering()".into(),
))
}
}
pub fn add(&mut self, _doc_id: u32, vector: Vec<f32>) -> Result<(), RetrieveError> {
self.add_slice(_doc_id, &vector)
}
pub fn add_slice(&mut self, doc_id: u32, vector: &[f32]) -> Result<(), RetrieveError> {
if self.built {
return Err(RetrieveError::InvalidParameter(
"cannot add vectors after index is built".into(),
));
}
if vector.len() != self.dimension {
return Err(RetrieveError::DimensionMismatch {
query_dim: vector.len(),
doc_dim: self.dimension,
});
}
let norm: f32 = vector.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 1e-10 {
self.vectors.extend(vector.iter().map(|x| x / norm));
} else {
self.vectors.extend_from_slice(vector);
}
self.doc_ids.push(doc_id);
self.num_vectors += 1;
Ok(())
}
pub fn build(&mut self) -> Result<(), RetrieveError> {
self.build_with_training_config(IVFPQTrainingConfig::default())
}
pub fn build_with_training_sample(
&mut self,
training_sample_size: usize,
) -> Result<(), RetrieveError> {
self.build_with_training_options(Some(training_sample_size), default_kmeans_max_iter())
}
pub fn build_with_training_options(
&mut self,
training_sample_size: Option<usize>,
kmeans_max_iter: usize,
) -> Result<(), RetrieveError> {
self.build_with_training_config(IVFPQTrainingConfig {
sample_size: training_sample_size,
kmeans_max_iter,
})
}
fn build_with_training_config(
&mut self,
training: IVFPQTrainingConfig,
) -> Result<(), RetrieveError> {
if self.built {
return Ok(());
}
if self.num_vectors == 0 {
return Err(RetrieveError::EmptyIndex);
}
if self.num_vectors < self.params.codebook_size {
return Err(RetrieveError::InvalidParameter(format!(
"IVF-PQ requires at least codebook_size training vectors \
(got {} vectors, codebook_size = {}). Add more vectors \
before build(), or lower codebook_size in IVFPQParams.",
self.num_vectors, self.params.codebook_size
)));
}
if training.kmeans_max_iter == 0 {
return Err(RetrieveError::InvalidParameter(
"kmeans_max_iter must be greater than 0".into(),
));
}
let training_indices = self.training_indices(training.sample_size)?;
let training_count = training_indices.len();
let min_training_count = self
.params
.num_clusters
.min(self.num_vectors)
.max(self.params.codebook_size);
if training_count < min_training_count {
return Err(RetrieveError::InvalidParameter(format!(
"IVF-PQ training sample is too small (got {}, need at least {} \
for num_clusters={} and codebook_size={})",
training_count,
min_training_count,
self.params.num_clusters,
self.params.codebook_size
)));
}
let sampled_training = if training_count == self.num_vectors {
None
} else {
Some(copy_rows_by_index(
&self.vectors,
self.dimension,
&training_indices,
))
};
let training_vectors = sampled_training.as_deref().unwrap_or(&self.vectors);
let mut kmeans =
crate::partitioning::kmeans::KMeans::new(self.dimension, self.params.num_clusters)?
.with_seed(self.params.seed)
.with_max_iter(training.kmeans_max_iter);
kmeans.fit(training_vectors, training_count)?;
self.centroids = kmeans
.centroids()
.iter()
.flat_map(|c| c.iter().copied())
.collect();
#[cfg(feature = "hnsw")]
{
let num_centroids = self.centroids.len() / self.dimension;
let mut hnsw = crate::hnsw::HNSWIndex::builder(self.dimension)
.m(16)
.ef_construction(200)
.auto_normalize(true)
.build()?;
for i in 0..num_centroids {
let centroid = self.get_centroid(i);
hnsw.add_slice(i as u32, centroid)?;
}
hnsw.build()?;
self.coarse_quantizer = Some(hnsw);
}
let assignments = kmeans.assign_clusters(&self.vectors, self.num_vectors);
let mut temp_clusters: Vec<(Vec<u32>, u64)> =
vec![(Vec::new(), 0); self.params.num_clusters];
if let Some(ref metadata_store) = self.metadata {
if let Some(ref field) = self.filter_field {
for (vector_idx, &cluster_idx) in assignments.iter().enumerate() {
temp_clusters[cluster_idx].0.push(vector_idx as u32);
let actual_doc_id = self.doc_ids[vector_idx];
if let Some(metadata) = metadata_store.get(actual_doc_id) {
if let Some(crate::filtering::MetadataValue::Int(n)) = metadata.get(field) {
let category_id = *n as u64;
if category_id < 64 {
temp_clusters[cluster_idx].1 |= 1u64 << category_id;
}
}
}
}
} else {
for (vector_idx, &cluster_idx) in assignments.iter().enumerate() {
temp_clusters[cluster_idx].0.push(vector_idx as u32);
}
}
} else {
for (vector_idx, &cluster_idx) in assignments.iter().enumerate() {
temp_clusters[cluster_idx].0.push(vector_idx as u32);
}
}
#[cfg(feature = "id-compression")]
if let Some(method) = self.params.id_compression.as_ref() {
if !matches!(
method,
crate::compression::IdCompressionMethod::None
| crate::compression::IdCompressionMethod::DeltaVarint
) {
return Err(RetrieveError::Other(format!(
"unsupported IVF-PQ ID compression method: {method:?}; \
only None and DeltaVarint are implemented"
)));
}
}
self.clusters = temp_clusters
.into_iter()
.map(|(ids, bitmask)| {
#[cfg(feature = "id-compression")]
{
if let Some(ref method) = self.params.id_compression {
if ids.len() >= self.params.compression_threshold {
match method {
crate::compression::IdCompressionMethod::DeltaVarint => {
let compressor =
crate::compression::DeltaVarintCompressor::new();
let universe_size = self.num_vectors as u32;
let ids_clone = ids.clone();
Cluster::new_compressed(
ids,
bitmask,
&compressor,
universe_size,
)
.unwrap_or_else(|_| Cluster::new(ids_clone, bitmask))
}
crate::compression::IdCompressionMethod::None => {
Cluster::new(ids, bitmask)
}
#[allow(unreachable_patterns)]
_ => {
unreachable!("unsupported compression methods are preflighted")
}
}
} else {
Cluster::new(ids, bitmask)
}
} else {
Cluster::new(ids, bitmask)
}
}
#[cfg(not(feature = "id-compression"))]
{
Cluster::new(ids, bitmask)
}
})
.collect();
let mut residuals = Vec::with_capacity(self.num_vectors * self.dimension);
for (i, &cluster_idx) in assignments.iter().enumerate() {
let vec = self.get_vector(i);
let centroid = self.get_centroid(cluster_idx);
for (v, c) in vec.iter().zip(centroid.iter()) {
residuals.push(v - c);
}
}
let pq: Quantizer = if self.params.use_opq {
let mut opq = OptimizedProductQuantizer::new(
self.dimension,
self.params.num_codebooks,
self.params.codebook_size,
)?;
let training_residuals = if training_count == self.num_vectors {
None
} else {
Some(copy_rows_by_index(
&residuals,
self.dimension,
&training_indices,
))
};
let opq_training = training_residuals.as_deref().unwrap_or(&residuals);
opq.fit(opq_training, training_count, 10)?; Quantizer::Optimized(opq)
} else {
let mut pq = ProductQuantizer::new(
self.dimension,
self.params.num_codebooks,
self.params.codebook_size,
)?;
let training_residuals = if training_count == self.num_vectors {
None
} else {
Some(copy_rows_by_index(
&residuals,
self.dimension,
&training_indices,
))
};
let pq_training = training_residuals.as_deref().unwrap_or(&residuals);
pq.fit_with_seed_and_max_iter(
pq_training,
training_count,
Some(self.params.seed),
training.kmeans_max_iter,
)?;
Quantizer::Product(pq)
};
self.quantized_codes = Vec::with_capacity(self.num_vectors * self.params.num_codebooks);
for i in 0..self.num_vectors {
let residual = &residuals[i * self.dimension..(i + 1) * self.dimension];
let codes = pq.quantize(residual);
self.quantized_codes.extend_from_slice(&codes);
}
self.build_scan_caches();
self.pq = Some(pq);
self.built = true;
Ok(())
}
fn training_indices(
&self,
training_sample_size: Option<usize>,
) -> Result<Vec<usize>, RetrieveError> {
let Some(sample_size) = training_sample_size else {
return Ok((0..self.num_vectors).collect());
};
if sample_size == 0 {
return Err(RetrieveError::InvalidParameter(
"training_sample_size must be greater than 0".into(),
));
}
if sample_size >= self.num_vectors {
return Ok((0..self.num_vectors).collect());
}
let mut indices: Vec<usize> = (0..self.num_vectors).collect();
let mut rng = rand::rngs::StdRng::seed_from_u64(self.params.seed);
indices.shuffle(&mut rng);
indices.truncate(sample_size);
indices.sort_unstable();
Ok(indices)
}
fn build_scan_caches(&mut self) {
let num_cb = self.params.num_codebooks;
let quantized_codes = &self.quantized_codes;
for cluster in &mut self.clusters {
let ids = cluster.get_ids_ref();
if ids.len() < SIMD_BATCH_THRESHOLD {
cluster.set_fastscan_codes(None);
cluster.set_adc_codes(None);
continue;
}
let mut codes_batch = Vec::with_capacity(ids.len() * num_cb);
for &vector_idx in ids.as_ref() {
let start = vector_idx as usize * num_cb;
codes_batch.extend_from_slice(&quantized_codes[start..start + num_cb]);
}
if self.params.codebook_size == 16 {
cluster.set_fastscan_codes(Some(PackedCodes4bit::pack(
&codes_batch,
ids.len(),
num_cb,
)));
cluster.set_adc_codes(None);
} else {
cluster.set_fastscan_codes(None);
cluster.set_adc_codes(Some(codes_batch));
}
}
}
pub fn compact(&mut self) {
assert!(self.built, "compact() called before build()");
self.vectors = Vec::new();
}
pub fn save_to_dir(&self, output_dir: impl AsRef<Path>) -> Result<(), RetrieveError> {
if !self.built {
return Err(RetrieveError::InvalidParameter(
"cannot save unbuilt IVF-PQ index".into(),
));
}
if self.metadata.is_some() && self.filter_field.is_none() {
return Err(RetrieveError::InvalidParameter(
"IVF-PQ filter metadata requires a filter field".into(),
));
}
let pq = self
.pq
.clone()
.ok_or_else(|| RetrieveError::InvalidParameter("missing IVF-PQ quantizer".into()))?;
let output_dir = output_dir.as_ref();
std::fs::create_dir_all(output_dir)?;
let raw_vectors_present = self.vectors.len() == self.num_vectors * self.dimension;
let mut filter_metadata = if let Some(store) = &self.metadata {
store
.iter()
.map(|(&doc_id, metadata)| PersistedFilterMetadata {
doc_id,
metadata: metadata.clone(),
})
.collect::<Vec<_>>()
} else {
Vec::new()
};
filter_metadata.sort_by_key(|entry| entry.doc_id);
let manifest = IVFPQManifest {
version: IVFPQ_FORMAT_VERSION,
dimension: self.dimension,
num_vectors: self.num_vectors,
num_centroids: self.centroids.len() / self.dimension,
raw_vectors_present,
params: PersistedIVFPQParams::from(&self.params),
quantizer: pq,
filter_field: self.filter_field.clone(),
filter_metadata,
};
write_json_atomic(&output_dir.join("manifest.json"), &manifest)?;
write_f32_atomic(&output_dir.join("centroids.bin"), &self.centroids)?;
write_u32_atomic(&output_dir.join("doc_ids.bin"), &self.doc_ids)?;
write_bytes_atomic(&output_dir.join("codes.bin"), &self.quantized_codes)?;
let (list_offsets, list_codes) = build_list_codes(
&self.clusters,
&self.quantized_codes,
self.params.num_codebooks,
)?;
write_u64_atomic(&output_dir.join("list_offsets.bin"), &list_offsets)?;
write_bytes_atomic(&output_dir.join("list_codes.bin"), &list_codes)?;
write_clusters_atomic(&output_dir.join("clusters.bin"), &self.clusters)?;
if raw_vectors_present {
write_f32_atomic(&output_dir.join("raw_vectors.bin"), &self.vectors)?;
}
Ok(())
}
pub fn load_from_dir(input_dir: impl AsRef<Path>) -> Result<Self, RetrieveError> {
let input_dir = input_dir.as_ref();
let manifest: IVFPQManifest = read_json(&input_dir.join("manifest.json"))?;
validate_manifest(&manifest)?;
let params = manifest.params.into_params();
let mut index = Self::new(manifest.dimension, params)?;
index.num_vectors = manifest.num_vectors;
index.doc_ids = read_u32_exact(&input_dir.join("doc_ids.bin"), manifest.num_vectors)?;
index.centroids = read_f32_exact(
&input_dir.join("centroids.bin"),
manifest.num_centroids * manifest.dimension,
)?;
index.quantized_codes = read_bytes_exact(
&input_dir.join("codes.bin"),
manifest.num_vectors * index.params.num_codebooks,
)?;
index.clusters = read_clusters(
&input_dir.join("clusters.bin"),
index.params.num_clusters,
manifest.num_vectors,
)?;
index.vectors = if manifest.raw_vectors_present {
read_f32_exact(
&input_dir.join("raw_vectors.bin"),
manifest.num_vectors * manifest.dimension,
)?
} else {
Vec::new()
};
index.pq = Some(manifest.quantizer);
if manifest.filter_field.is_some() || !manifest.filter_metadata.is_empty() {
let filter_field = manifest.filter_field.clone().ok_or_else(|| {
RetrieveError::FormatError(
"IVF-PQ manifest has filter metadata without filter_field".into(),
)
})?;
index.filter_field = Some(filter_field);
index.metadata = Some(crate::filtering::MetadataStore::new());
for entry in manifest.filter_metadata {
index.add_metadata(entry.doc_id, entry.metadata)?;
}
}
index.built = true;
index.build_scan_caches();
Ok(index)
}
pub fn search(&self, query: &[f32], k: usize) -> Result<Vec<(u32, f32)>, RetrieveError> {
let candidates = self.search_approx_internal(query, k)?;
Ok(candidates
.into_iter()
.map(|(vector_idx, dist)| (self.doc_ids[vector_idx as usize], dist))
.collect())
}
#[cfg(feature = "benchmark")]
pub fn search_profiled(
&self,
query: &[f32],
k: usize,
) -> Result<(Vec<(u32, f32)>, IVFPQSearchProfile), RetrieveError> {
let mut profile = IVFPQSearchProfile::default();
let candidates = self.search_approx_internal_observed(query, k, Some(&mut profile))?;
let results = candidates
.into_iter()
.map(|(vector_idx, dist)| (self.doc_ids[vector_idx as usize], dist))
.collect();
Ok((results, profile))
}
pub fn search_reranked(
&self,
query: &[f32],
k: usize,
candidate_pool: usize,
) -> Result<Vec<(u32, f32)>, RetrieveError> {
if !self.built {
return Err(RetrieveError::InvalidParameter(
"index must be built before search".into(),
));
}
if query.len() != self.dimension {
return Err(RetrieveError::DimensionMismatch {
query_dim: query.len(),
doc_dim: self.dimension,
});
}
if self.vectors.is_empty() {
return Err(RetrieveError::InvalidParameter(
"search_reranked unavailable after compact()".into(),
));
}
let pool = candidate_pool.max(k);
let candidates = self.search_approx_internal(query, pool)?;
let query_normalized = normalize_query(query);
let reranked: Vec<(u32, f32)> = candidates
.into_iter()
.map(|(vector_idx, _approx_dist)| {
let vector = self.get_vector(vector_idx as usize);
let exact_dist =
crate::distance::cosine_distance_normalized(&query_normalized, vector);
(self.doc_ids[vector_idx as usize], exact_dist)
})
.collect();
Ok(finish_top_k_by_distance(reranked, k))
}
fn search_approx_internal(
&self,
query: &[f32],
k: usize,
) -> Result<Vec<(u32, f32)>, RetrieveError> {
#[cfg(feature = "benchmark")]
{
self.search_approx_internal_observed(query, k, None)
}
#[cfg(not(feature = "benchmark"))]
{
self.search_approx_internal_observed(query, k)
}
}
fn search_approx_internal_observed(
&self,
query: &[f32],
k: usize,
#[cfg(feature = "benchmark")] mut profile: Option<&mut IVFPQSearchProfile>,
) -> Result<Vec<(u32, f32)>, RetrieveError> {
if !self.built {
return Err(RetrieveError::InvalidParameter(
"index must be built before search".into(),
));
}
if query.len() != self.dimension {
return Err(RetrieveError::DimensionMismatch {
query_dim: query.len(),
doc_dim: self.dimension,
});
}
let pq = self
.pq
.as_ref()
.ok_or(RetrieveError::InvalidParameter("PQ not initialized".into()))?;
#[cfg(feature = "benchmark")]
let normalize_start = Instant::now();
let query_normalized = normalize_query(query);
#[cfg(feature = "benchmark")]
if let Some(profile) = &mut profile {
profile.normalize += normalize_start.elapsed();
}
let query = query_normalized.as_slice();
#[cfg(feature = "benchmark")]
let centroid_start = Instant::now();
let cluster_distances = self.find_nearest_centroids(query, self.params.nprobe);
#[cfg(feature = "benchmark")]
if let Some(profile) = &mut profile {
profile.centroid_lookup += centroid_start.elapsed();
profile.probed_clusters += cluster_distances.len();
}
let expected_candidates = cluster_distances
.iter()
.map(|(cluster_idx, _)| self.clusters[*cluster_idx].len())
.sum();
let mut candidates = Vec::with_capacity(expected_candidates);
let mut query_residual = vec![0.0f32; self.dimension];
let mut codes_batch = Vec::new();
let mut distances_batch = Vec::new();
let mut adc_table = Vec::new();
for (cluster_idx, _) in &cluster_distances {
let cluster = &self.clusters[*cluster_idx];
let ids = cluster.get_ids_ref();
#[cfg(feature = "benchmark")]
let residual_start = Instant::now();
let centroid = self.get_centroid(*cluster_idx);
for (i, (q, c)) in query.iter().zip(centroid.iter()).enumerate() {
query_residual[i] = q - c;
}
#[cfg(feature = "benchmark")]
if let Some(profile) = &mut profile {
profile.residual += residual_start.elapsed();
profile.scanned_vectors += ids.len();
}
#[cfg(feature = "benchmark")]
let adc_table_start = Instant::now();
pq.compute_adc_table_into(&query_residual, &mut adc_table)?;
#[cfg(feature = "benchmark")]
if let Some(profile) = &mut profile {
profile.adc_table += adc_table_start.elapsed();
}
if ids.len() >= SIMD_BATCH_THRESHOLD {
let num_cb = self.params.num_codebooks;
#[cfg(feature = "benchmark")]
if let Some(profile) = &mut profile {
profile.simd_clusters += 1;
}
let fastscan_distances;
let distances = if self.params.codebook_size == 16 {
if let Some(packed) = cluster.fastscan_codes.as_ref() {
#[cfg(feature = "benchmark")]
{
let dispatch_start = Instant::now();
fastscan_distances =
crate::pq_simd::fastscan_batch_flat(packed, &adc_table);
if let Some(profile) = &mut profile {
profile.adc_dispatch += dispatch_start.elapsed();
}
fastscan_distances.as_slice()
}
#[cfg(not(feature = "benchmark"))]
{
fastscan_distances =
crate::pq_simd::fastscan_batch_flat(packed, &adc_table);
fastscan_distances.as_slice()
}
} else {
#[cfg(feature = "benchmark")]
let copy_start = Instant::now();
codes_batch.clear();
codes_batch.reserve(ids.len() * num_cb);
for &vector_idx in ids.as_ref() {
let start = vector_idx as usize * num_cb;
codes_batch
.extend_from_slice(&self.quantized_codes[start..start + num_cb]);
}
#[cfg(feature = "benchmark")]
if let Some(profile) = &mut profile {
profile.code_copy += copy_start.elapsed();
profile.code_copy_bytes += codes_batch.len();
}
let packed = PackedCodes4bit::pack(&codes_batch, ids.len(), num_cb);
#[cfg(feature = "benchmark")]
let dispatch_start = Instant::now();
fastscan_distances =
crate::pq_simd::fastscan_batch_flat(&packed, &adc_table);
#[cfg(feature = "benchmark")]
if let Some(profile) = &mut profile {
profile.adc_dispatch += dispatch_start.elapsed();
}
fastscan_distances.as_slice()
}
} else {
let codes_scan = if let Some(codes) = cluster.adc_codes.as_ref() {
codes.as_slice()
} else {
#[cfg(feature = "benchmark")]
let copy_start = Instant::now();
codes_batch.clear();
codes_batch.reserve(ids.len() * num_cb);
for &vector_idx in ids.as_ref() {
let start = vector_idx as usize * num_cb;
codes_batch
.extend_from_slice(&self.quantized_codes[start..start + num_cb]);
}
#[cfg(feature = "benchmark")]
if let Some(profile) = &mut profile {
profile.code_copy += copy_start.elapsed();
profile.code_copy_bytes += codes_batch.len();
}
codes_batch.as_slice()
};
let packed_lut = PackedLUTRef::from_flat(
&adc_table,
self.params.num_codebooks,
self.params.codebook_size,
);
#[cfg(feature = "benchmark")]
let dispatch_start = Instant::now();
adc_batch_dispatch_into(codes_scan, num_cb, &packed_lut, &mut distances_batch);
#[cfg(feature = "benchmark")]
if let Some(profile) = &mut profile {
profile.adc_dispatch += dispatch_start.elapsed();
}
distances_batch.as_slice()
};
for (i, &vector_idx) in ids.iter().enumerate() {
candidates.push((vector_idx, distances[i]));
}
} else {
#[cfg(feature = "benchmark")]
if let Some(profile) = &mut profile {
profile.scalar_clusters += 1;
}
#[cfg(feature = "benchmark")]
let dispatch_start = Instant::now();
for &vector_idx in ids.as_ref() {
let start = vector_idx as usize * self.params.num_codebooks;
let end = start + self.params.num_codebooks;
let codes = &self.quantized_codes[start..end];
let dist = pq.distance_with_table(&adc_table, codes);
candidates.push((vector_idx, dist));
}
#[cfg(feature = "benchmark")]
if let Some(profile) = &mut profile {
profile.adc_dispatch += dispatch_start.elapsed();
}
}
}
#[cfg(feature = "benchmark")]
if let Some(profile) = &mut profile {
profile.candidate_count += candidates.len();
}
#[cfg(feature = "benchmark")]
let finalizer_start = Instant::now();
let results = finish_top_k_by_distance(candidates, k);
#[cfg(feature = "benchmark")]
if let Some(profile) = &mut profile {
profile.finalizer += finalizer_start.elapsed();
profile.returned_count += results.len();
}
Ok(results)
}
pub fn search_with_filter(
&self,
query: &[f32],
k: usize,
filter: &crate::filtering::MetadataFilter,
) -> Result<Vec<(u32, f32)>, RetrieveError> {
if !self.built {
return Err(RetrieveError::InvalidParameter(
"index must be built before search".into(),
));
}
if query.len() != self.dimension {
return Err(RetrieveError::DimensionMismatch {
query_dim: query.len(),
doc_dim: self.dimension,
});
}
let desired_category: u64 = match filter {
crate::filtering::MetadataFilter::Equals { field, value } => {
if Some(field) != self.filter_field.as_ref() {
return Err(RetrieveError::InvalidParameter(format!(
"filter field '{}' doesn't match index filter field '{:?}'",
field, self.filter_field
)));
}
match value {
crate::filtering::MetadataValue::Int(n) if *n >= 0 && *n < 64 => *n as u64,
crate::filtering::MetadataValue::Int(n) => {
return Err(RetrieveError::InvalidParameter(format!(
"category ID {} exceeds bitmask limit of 63",
n
)));
}
_ => {
return Err(RetrieveError::InvalidParameter(
"category ID must be an integer in 0..63 for bitmask filtering".into(),
));
}
}
}
_ => {
return Err(RetrieveError::InvalidParameter(
"only equality filters on filter_field are supported".into(),
));
}
};
let filter_bit = 1u64 << desired_category;
let query_norm: f32 = query.iter().map(|x| x * x).sum::<f32>().sqrt();
let query_normalized: Vec<f32> = if query_norm > 1e-10 {
query.iter().map(|x| x / query_norm).collect()
} else {
query.to_vec()
};
let query = query_normalized.as_slice();
let cluster_distances = self.find_nearest_centroids(query, self.params.nprobe);
let mut candidates = Vec::new();
let mut query_residual = vec![0.0f32; self.dimension];
let mut adc_table = Vec::new();
let pq = self
.pq
.as_ref()
.ok_or(RetrieveError::InvalidParameter("PQ not initialized".into()))?;
for (cluster_idx, _) in &cluster_distances {
let cluster = &self.clusters[*cluster_idx];
if (cluster.filter_bitmask & filter_bit) == 0 {
continue;
}
let centroid = self.get_centroid(*cluster_idx);
for (i, (q, c)) in query.iter().zip(centroid.iter()).enumerate() {
query_residual[i] = q - c;
}
pq.compute_adc_table_into(&query_residual, &mut adc_table)?;
if let Some(ref metadata_store) = self.metadata {
let ids = cluster.get_ids_ref();
for &vector_idx in ids.as_ref() {
let actual_doc_id = self.doc_ids[vector_idx as usize];
if metadata_store.matches(actual_doc_id, filter) {
let start = vector_idx as usize * self.params.num_codebooks;
let end = start + self.params.num_codebooks;
let codes = &self.quantized_codes[start..end];
let dist = pq.distance_with_table(&adc_table, codes);
candidates.push((actual_doc_id, dist));
}
}
} else {
return Err(RetrieveError::InvalidParameter(
"metadata store not initialized".into(),
));
}
}
Ok(finish_top_k_by_distance(candidates, k))
}
pub fn memory_usage(&self) -> crate::memory::MemoryReport {
let vectors_bytes = self.vectors.capacity() * std::mem::size_of::<f32>();
let cluster_bytes = self.clusters.capacity() * std::mem::size_of::<Cluster>()
+ self
.clusters
.iter()
.map(Cluster::owned_bytes)
.sum::<usize>();
#[cfg(feature = "hnsw")]
let coarse_quantizer_bytes = self
.coarse_quantizer
.as_ref()
.map(|index| index.memory_usage().total())
.unwrap_or(0);
#[cfg(not(feature = "hnsw"))]
let coarse_quantizer_bytes = 0;
let quantized_bytes = self.quantized_codes.capacity()
+ self.pq.as_ref().map(Quantizer::owned_bytes).unwrap_or(0);
let metadata_bytes = self.doc_ids.capacity() * std::mem::size_of::<u32>()
+ self.centroids.capacity() * std::mem::size_of::<f32>()
+ self
.filter_field
.as_ref()
.map(|field| field.capacity())
.unwrap_or(0)
+ cluster_bytes;
crate::memory::MemoryReport {
vectors_bytes,
graph_bytes: coarse_quantizer_bytes,
quantized_bytes,
metadata_bytes,
}
}
#[inline]
fn get_vector(&self, idx: usize) -> &[f32] {
let start = idx * self.dimension;
let end = start + self.dimension;
&self.vectors[start..end]
}
#[inline]
fn get_centroid(&self, idx: usize) -> &[f32] {
let start = idx * self.dimension;
let end = start + self.dimension;
&self.centroids[start..end]
}
fn find_nearest_centroids(&self, query: &[f32], nprobe: usize) -> Vec<(usize, f32)> {
#[cfg(feature = "hnsw")]
if let Some(ref hnsw) = self.coarse_quantizer {
let ef = nprobe * 2;
if let Ok(results) = hnsw.search(query, nprobe, ef.max(nprobe)) {
return results
.into_iter()
.map(|(id, d)| (id as usize, d))
.collect();
}
}
let num_centroids = self.centroids.len() / self.dimension;
let mut dists: Vec<(usize, f32)> = (0..num_centroids)
.map(|idx| {
let c = self.get_centroid(idx);
(idx, crate::distance::cosine_distance_normalized(query, c))
})
.collect();
let nprobe = nprobe.min(dists.len());
if nprobe < dists.len() {
dists.select_nth_unstable_by(nprobe, |a, b| a.1.total_cmp(&b.1));
dists.truncate(nprobe);
}
dists.sort_unstable_by(|a, b| a.1.total_cmp(&b.1));
dists
}
}
impl IVFPQFileSearcher {
pub fn load(input_dir: impl AsRef<Path>) -> Result<Self, RetrieveError> {
Self::load_with_storage(input_dir.as_ref(), false)
}
#[cfg(feature = "persistence")]
pub fn load_mmap(input_dir: impl AsRef<Path>) -> Result<Self, RetrieveError> {
Self::load_with_storage(input_dir.as_ref(), true)
}
#[cfg(not(feature = "persistence"))]
fn load_with_storage(input_dir: &Path, mmap: bool) -> Result<Self, RetrieveError> {
let _ = mmap;
Self::load_from_parts(input_dir, false)
}
#[cfg(feature = "persistence")]
fn load_with_storage(input_dir: &Path, mmap: bool) -> Result<Self, RetrieveError> {
Self::load_from_parts(input_dir, mmap)
}
fn load_from_parts(input_dir: &Path, mmap: bool) -> Result<Self, RetrieveError> {
let manifest: IVFPQManifest = read_json(&input_dir.join("manifest.json"))?;
validate_manifest(&manifest)?;
let params = manifest.params.clone().into_params();
let doc_ids = read_u32_exact(&input_dir.join("doc_ids.bin"), manifest.num_vectors)?;
let centroids_len = checked_len(
manifest.num_centroids,
manifest.dimension,
"IVF-PQ centroid length overflow",
)?;
let centroids = read_f32_exact(&input_dir.join("centroids.bin"), centroids_len)?;
let clusters = read_clusters(
&input_dir.join("clusters.bin"),
params.num_clusters,
manifest.num_vectors,
)?;
let codes_len = checked_len(
manifest.num_vectors,
params.num_codebooks,
"IVF-PQ codes length overflow",
)?;
let codes = open_byte_storage(&input_dir.join("codes.bin"), codes_len, mmap)?;
let list_codes = open_list_code_storage(input_dir, params.num_clusters, codes_len, mmap)?;
let raw_vectors = if manifest.raw_vectors_present {
let raw_floats = checked_len(
manifest.num_vectors,
manifest.dimension,
"IVF-PQ raw vector length overflow",
)?;
Some(open_byte_storage(
&input_dir.join("raw_vectors.bin"),
checked_len(
raw_floats,
std::mem::size_of::<f32>(),
"IVF-PQ raw vector byte length overflow",
)?,
mmap,
)?)
} else {
None
};
let raw_vector_byte_len = checked_len(
manifest.dimension,
std::mem::size_of::<f32>(),
"IVF-PQ raw vector byte length overflow",
)?;
Ok(Self {
dimension: manifest.dimension,
num_vectors: manifest.num_vectors,
doc_ids,
params,
clusters,
centroids,
pq: manifest.quantizer,
codes,
list_codes,
raw_vectors,
code_buf: Vec::new(),
raw_byte_buf: vec![0; raw_vector_byte_len],
vec_buf: vec![0.0; manifest.dimension],
})
}
pub fn set_nprobe(&mut self, nprobe: usize) {
self.params.nprobe = nprobe;
}
pub fn nprobe(&self) -> usize {
self.params.nprobe
}
pub fn num_vectors(&self) -> usize {
self.num_vectors
}
pub fn dimension(&self) -> usize {
self.dimension
}
pub fn num_clusters(&self) -> usize {
self.params.num_clusters
}
pub fn num_codebooks(&self) -> usize {
self.params.num_codebooks
}
pub fn codebook_size(&self) -> usize {
self.params.codebook_size
}
pub fn search(&mut self, query: &[f32], k: usize) -> Result<Vec<(u32, f32)>, RetrieveError> {
let candidates = self.search_approx_internal(query, k)?;
Ok(candidates
.into_iter()
.map(|(vector_idx, dist)| (self.doc_ids[vector_idx as usize], dist))
.collect())
}
pub fn search_with_diagnostics(
&mut self,
query: &[f32],
k: usize,
) -> Result<(Vec<(u32, f32)>, IVFPQFileSearchDiagnostics), RetrieveError> {
let (candidates, diagnostics) = self.search_approx_internal_with_diagnostics(query, k)?;
let results = candidates
.into_iter()
.map(|(vector_idx, dist)| (self.doc_ids[vector_idx as usize], dist))
.collect();
Ok((results, diagnostics))
}
pub fn search_reranked(
&mut self,
query: &[f32],
k: usize,
candidate_pool: usize,
) -> Result<Vec<(u32, f32)>, RetrieveError> {
Ok(self
.search_reranked_with_diagnostics(query, k, candidate_pool)?
.0)
}
pub fn search_reranked_with_diagnostics(
&mut self,
query: &[f32],
k: usize,
candidate_pool: usize,
) -> Result<(Vec<(u32, f32)>, IVFPQFileSearchDiagnostics), RetrieveError> {
if self.raw_vectors.is_none() {
return Err(RetrieveError::InvalidParameter(
"search_reranked unavailable without raw_vectors.bin".into(),
));
}
if query.len() != self.dimension {
return Err(RetrieveError::DimensionMismatch {
query_dim: query.len(),
doc_dim: self.dimension,
});
}
let pool = candidate_pool.max(k);
let (mut candidates, mut diagnostics) =
self.search_approx_internal_with_diagnostics(query, pool)?;
let raw_vector_byte_len = checked_len(
self.dimension,
std::mem::size_of::<f32>(),
"IVF-PQ diagnostic vector byte length overflow",
)?;
diagnostics.raw_vector_reads = candidates.len();
diagnostics.raw_vector_bytes = checked_len(
candidates.len(),
raw_vector_byte_len,
"IVF-PQ diagnostic raw-vector byte count overflow",
)?;
diagnostics.reranked_candidates = candidates.len();
let query_normalized = normalize_query(query);
let mut reranked = Vec::with_capacity(candidates.len());
let Some(storage) = self.raw_vectors.as_mut() else {
return Err(RetrieveError::InvalidParameter(
"search_reranked unavailable without raw_vectors.bin".into(),
));
};
if matches!(storage, IVFPQByteStorage::File(_)) {
candidates.sort_unstable_by_key(|(vector_idx, _)| *vector_idx);
}
for (vector_idx, _approx_dist) in candidates {
let vector = read_vector_from_storage(
storage,
&mut self.raw_byte_buf,
&mut self.vec_buf,
vector_idx as usize,
self.dimension,
)?;
let exact_dist = crate::distance::cosine_distance_normalized(&query_normalized, vector);
reranked.push((self.doc_ids[vector_idx as usize], exact_dist));
}
Ok((finish_top_k_by_distance(reranked, k), diagnostics))
}
fn search_approx_internal(
&mut self,
query: &[f32],
k: usize,
) -> Result<Vec<(u32, f32)>, RetrieveError> {
self.search_approx_internal_with_diagnostics(query, k)
.map(|(candidates, _)| candidates)
}
fn search_approx_internal_with_diagnostics(
&mut self,
query: &[f32],
k: usize,
) -> Result<(Vec<(u32, f32)>, IVFPQFileSearchDiagnostics), RetrieveError> {
if query.len() != self.dimension {
return Err(RetrieveError::DimensionMismatch {
query_dim: query.len(),
doc_dim: self.dimension,
});
}
let query_normalized = normalize_query(query);
let query = query_normalized.as_slice();
let cluster_distances = self.find_nearest_centroids(query, self.params.nprobe);
let mut diagnostics = IVFPQFileSearchDiagnostics {
probed_lists: cluster_distances.len(),
..IVFPQFileSearchDiagnostics::default()
};
let expected_candidates = cluster_distances
.iter()
.map(|(cluster_idx, _)| self.clusters[*cluster_idx].len())
.sum();
let mut candidates = Vec::with_capacity(expected_candidates);
let mut query_residual = vec![0.0f32; self.dimension];
let mut codes_batch = Vec::new();
let mut distances_batch = Vec::new();
let mut adc_table = Vec::new();
for (cluster_idx, _) in &cluster_distances {
let cluster = &self.clusters[*cluster_idx];
let ids = cluster.get_ids_ref();
diagnostics.scanned_vectors += ids.len();
let code_bytes = checked_len(
ids.len(),
self.params.num_codebooks,
"IVF-PQ diagnostic code byte count overflow",
)?;
let centroid = self.get_centroid(*cluster_idx);
for (i, (q, c)) in query.iter().zip(centroid.iter()).enumerate() {
query_residual[i] = q - c;
}
self.pq
.compute_adc_table_into(&query_residual, &mut adc_table)?;
if ids.len() >= SIMD_BATCH_THRESHOLD {
if let Some(list_codes) = self.list_codes.as_mut() {
diagnostics.code_reads += usize::from(!ids.is_empty());
diagnostics.code_bytes += code_bytes;
read_list_codes_for_cluster(
list_codes,
*cluster_idx,
ids.len(),
self.params.num_codebooks,
&mut codes_batch,
)?;
} else {
diagnostics.code_reads += ids.len();
diagnostics.code_bytes += code_bytes;
append_codes_for_ids(
&mut self.codes,
&mut codes_batch,
ids.as_ref(),
self.params.num_codebooks,
)?;
}
if self.params.codebook_size == 16 {
let packed =
PackedCodes4bit::pack(&codes_batch, ids.len(), self.params.num_codebooks);
let distances = crate::pq_simd::fastscan_batch_flat(&packed, &adc_table);
for (i, &vector_idx) in ids.iter().enumerate() {
candidates.push((vector_idx, distances[i]));
}
} else {
let packed_lut = PackedLUTRef::from_flat(
&adc_table,
self.params.num_codebooks,
self.params.codebook_size,
);
adc_batch_dispatch_into(
&codes_batch,
self.params.num_codebooks,
&packed_lut,
&mut distances_batch,
);
for (i, &vector_idx) in ids.iter().enumerate() {
candidates.push((vector_idx, distances_batch[i]));
}
}
} else {
diagnostics.code_reads += ids.len();
diagnostics.code_bytes += code_bytes;
for &vector_idx in ids.as_ref() {
let codes = read_code_from_storage(
&mut self.codes,
&mut self.code_buf,
vector_idx as usize,
self.params.num_codebooks,
)?;
let dist = self.pq.distance_with_table(&adc_table, codes);
candidates.push((vector_idx, dist));
}
}
}
let results = finish_top_k_by_distance(candidates, k);
diagnostics.retained_candidates = results.len();
Ok((results, diagnostics))
}
#[inline]
fn get_centroid(&self, idx: usize) -> &[f32] {
let start = idx * self.dimension;
let end = start + self.dimension;
&self.centroids[start..end]
}
fn find_nearest_centroids(&self, query: &[f32], nprobe: usize) -> Vec<(usize, f32)> {
let num_centroids = self.centroids.len() / self.dimension;
let mut dists: Vec<(usize, f32)> = (0..num_centroids)
.map(|idx| {
let c = self.get_centroid(idx);
(idx, crate::distance::cosine_distance_normalized(query, c))
})
.collect();
let nprobe = nprobe.min(dists.len());
if nprobe < dists.len() {
dists.select_nth_unstable_by(nprobe, |a, b| a.1.total_cmp(&b.1));
dists.truncate(nprobe);
}
dists.sort_unstable_by(|a, b| a.1.total_cmp(&b.1));
dists
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn finish_top_k_by_distance_keeps_sorted_prefix() {
let candidates = vec![
(10, 5.0),
(11, 1.5),
(12, 3.0),
(13, 0.5),
(14, 2.0),
(15, 4.0),
];
let out = finish_top_k_by_distance(candidates, 3);
assert_eq!(out, vec![(13, 0.5), (11, 1.5), (14, 2.0)]);
}
#[test]
fn finish_top_k_by_distance_handles_zero_and_short_inputs() {
assert!(finish_top_k_by_distance(vec![(1, 1.0)], 0).is_empty());
assert!(finish_top_k_by_distance(Vec::new(), 5).is_empty());
let out = finish_top_k_by_distance(vec![(2, 2.0), (1, 1.0)], 5);
assert_eq!(out, vec![(1, 1.0), (2, 2.0)]);
}
#[test]
fn compact_search_works() {
let dim = 16;
let n = 200;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(42);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 16,
nprobe: 4,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let v: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(i as u32, v).unwrap();
}
index.build().unwrap();
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
let before = index.search(&query, 5).unwrap();
index.compact();
assert!(index.vectors.is_empty());
let after = index.search(&query, 5).unwrap();
assert_eq!(before, after);
}
#[cfg(feature = "id-compression")]
#[test]
fn unsupported_id_compression_method_fails_build() {
let dim = 4;
let params = IVFPQParams {
num_clusters: 2,
num_codebooks: 2,
codebook_size: 4,
nprobe: 2,
id_compression: Some(crate::compression::IdCompressionMethod::EliasFano),
compression_threshold: 0,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..16_u32 {
let mut vector = vec![0.0; dim];
vector[i as usize % dim] = 1.0;
vector[((i as usize) + 1) % dim] = 0.5;
index.add(i, vector).unwrap();
}
let err = index.build().unwrap_err();
assert!(
err.to_string()
.contains("unsupported IVF-PQ ID compression method"),
"unexpected error: {err}"
);
}
#[test]
fn reranked_search_uses_exact_distances() {
let dim = 16;
let n = 200;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(7);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 16,
nprobe: 4,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
let mut vectors = Vec::new();
for i in 0..n {
let v: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(i as u32, v.clone()).unwrap();
vectors.push(v);
}
index.build().unwrap();
let query_id = 17u32;
let results = index
.search_reranked(&vectors[query_id as usize], 1, n)
.unwrap();
assert_eq!(results[0].0, query_id);
assert!(
results[0].1.abs() < 1e-5,
"self-query exact distance should be near zero, got {}",
results[0].1
);
}
#[test]
fn search_and_rerank_return_external_doc_ids() {
let dim = 16;
let n = 200;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(9);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 16,
nprobe: 4,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
let mut vectors = Vec::new();
let mut doc_ids = Vec::new();
for i in 0..n {
let doc_id = 10_000 + (i as u32 * 7);
let v: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(doc_id, v.clone()).unwrap();
vectors.push(v);
doc_ids.push(doc_id);
}
index.build().unwrap();
let query_idx = 17usize;
let query = &vectors[query_idx];
let search_results = index.search(query, 10).unwrap();
assert!(search_results.iter().all(|(id, _)| doc_ids.contains(id)));
let reranked = index.search_reranked(query, 1, n).unwrap();
assert_eq!(reranked[0].0, doc_ids[query_idx]);
}
#[test]
fn sub_4bit_codebook_search_uses_standard_adc() {
let dim = 8;
let n = 96;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(10);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 8,
nprobe: 4,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let v: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(1_000 + i as u32, v).unwrap();
}
index.build().unwrap();
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
let results = index.search(&query, 5).unwrap();
assert_eq!(results.len(), 5);
assert!(results
.iter()
.all(|(id, dist)| *id >= 1_000 && dist.is_finite()));
}
#[test]
fn four_bit_builds_prepacked_fastscan_codes() {
let dim = 16;
let n = 200;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(11);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 16,
nprobe: 4,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let v: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(i as u32, v).unwrap();
}
index.build().unwrap();
assert!(
index
.clusters
.iter()
.filter(|cluster| cluster.len() >= SIMD_BATCH_THRESHOLD)
.all(|cluster| cluster.fastscan_codes.is_some()),
"clusters large enough for batched 4-bit search should be prepacked"
);
}
#[test]
fn standard_adc_builds_prepacked_cluster_codes() {
let dim = 16;
let n = 320;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(16);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 256,
nprobe: 4,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let v: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(i as u32, v).unwrap();
}
index.build().unwrap();
assert!(
index
.clusters
.iter()
.filter(|cluster| cluster.len() >= SIMD_BATCH_THRESHOLD)
.all(|cluster| cluster.adc_codes.is_some()),
"clusters large enough for batched standard ADC should be prepacked"
);
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
let cached = index.search(&query, 10).unwrap();
for cluster in &mut index.clusters {
cluster.set_adc_codes(None);
}
assert_eq!(index.search(&query, 10).unwrap(), cached);
}
#[test]
fn memory_usage_reports_raw_quantized_and_cache_buffers() {
let dim = 16;
let n = 320;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(17);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 256,
nprobe: 4,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let v: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(i as u32, v).unwrap();
}
index.build().unwrap();
let report = index.memory_usage();
assert!(report.vectors_bytes >= n * dim * std::mem::size_of::<f32>());
assert!(report.quantized_bytes >= n * index.params.num_codebooks);
assert!(report.metadata_bytes >= n * std::mem::size_of::<u32>());
index.compact();
let compacted = index.memory_usage();
assert_eq!(compacted.vectors_bytes, 0);
assert!(compacted.quantized_bytes >= n * index.params.num_codebooks);
assert!(compacted.total() < report.total());
}
#[test]
fn sampled_training_is_deterministic_with_seed() {
let dim = 16;
let n = 180;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(15);
let vectors: Vec<Vec<f32>> = (0..n)
.map(|_| (0..dim).map(|_| rng.random::<f32>()).collect())
.collect();
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
let build = || {
let params = IVFPQParams {
num_clusters: 8,
num_codebooks: 4,
codebook_size: 16,
nprobe: 8,
seed: 77,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for (i, vector) in vectors.iter().enumerate() {
index.add_slice(i as u32, vector).unwrap();
}
index.build_with_training_sample(64).unwrap();
#[cfg(feature = "hnsw")]
{
index.coarse_quantizer = None;
}
index
};
let a = build();
let b = build();
assert_eq!(a.centroids, b.centroids);
assert_eq!(a.quantized_codes, b.quantized_codes);
assert_eq!(a.search(&query, 10).unwrap(), b.search(&query, 10).unwrap());
}
#[test]
fn sampled_training_indexes_all_inserted_vectors() {
let dim = 16;
let n = 180;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(16);
let params = IVFPQParams {
num_clusters: 8,
num_codebooks: 4,
codebook_size: 16,
nprobe: 8,
seed: 88,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
let mut vectors = Vec::new();
for i in 0..n {
let doc_id = 50_000 + i as u32;
let vector: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(doc_id, vector.clone()).unwrap();
vectors.push(vector);
}
index.build_with_training_sample(64).unwrap();
let query_idx = n - 1;
let results = index.search_reranked(&vectors[query_idx], 1, n).unwrap();
assert_eq!(results[0].0, 50_000 + query_idx as u32);
}
#[test]
fn save_load_sampled_training_preserves_search() {
let dim = 16;
let n = 180;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(17);
let params = IVFPQParams {
num_clusters: 8,
num_codebooks: 4,
codebook_size: 16,
nprobe: 8,
seed: 99,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let vector: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(10_000 + i as u32, vector).unwrap();
}
index.build_with_training_options(Some(64), 5).unwrap();
#[cfg(feature = "hnsw")]
{
index.coarse_quantizer = None;
}
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
let approx_before = index.search(&query, 10).unwrap();
let reranked_before = index.search_reranked(&query, 10, 80).unwrap();
let dir = tempfile::tempdir().unwrap();
index.save_to_dir(dir.path()).unwrap();
let loaded = IVFPQIndex::load_from_dir(dir.path()).unwrap();
assert_eq!(loaded.search(&query, 10).unwrap(), approx_before);
assert_eq!(
loaded.search_reranked(&query, 10, 80).unwrap(),
reranked_before
);
}
#[test]
fn reranked_search_requires_raw_vectors() {
let dim = 16;
let n = 200;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(8);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 16,
nprobe: 4,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let v: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(i as u32, v).unwrap();
}
index.build().unwrap();
index.compact();
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
let err = index.search_reranked(&query, 5, 50).unwrap_err();
assert!(
err.to_string().contains("search_reranked unavailable"),
"unexpected error: {err}"
);
}
#[test]
fn save_load_roundtrip_preserves_search_and_rerank() {
let dim = 16;
let n = 240;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(11);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 16,
nprobe: 4,
seed: 123,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
let mut doc_ids = Vec::new();
for i in 0..n {
let doc_id = 10_000 + i as u32;
let v: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(doc_id, v).unwrap();
doc_ids.push(doc_id);
}
index.build().unwrap();
#[cfg(feature = "hnsw")]
{
index.coarse_quantizer = None;
}
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
let approx_before = index.search(&query, 10).unwrap();
let reranked_before = index.search_reranked(&query, 10, 80).unwrap();
let dir = tempfile::tempdir().unwrap();
index.save_to_dir(dir.path()).unwrap();
let loaded = IVFPQIndex::load_from_dir(dir.path()).unwrap();
assert_eq!(loaded.search(&query, 10).unwrap(), approx_before);
assert_eq!(
loaded.search_reranked(&query, 10, 80).unwrap(),
reranked_before
);
assert!(loaded
.clusters
.iter()
.filter(|cluster| cluster.len() >= SIMD_BATCH_THRESHOLD)
.all(|cluster| cluster.fastscan_codes.is_some()));
assert!(approx_before.iter().all(|(id, _)| doc_ids.contains(id)));
}
#[test]
fn save_load_preserves_filter_metadata() {
let dim = 16;
let n = 240;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(115);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 16,
nprobe: 4,
seed: 321,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::with_filtering(dim, params, "group").unwrap();
for i in 0..n {
let doc_id = 30_000 + i as u32;
let v: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(doc_id, v).unwrap();
let mut metadata = crate::filtering::DocumentMetadata::new();
metadata.insert(
"group".to_string(),
crate::filtering::MetadataValue::Int((i % 4) as i64),
);
metadata.insert(
"label".to_string(),
crate::filtering::MetadataValue::Str(format!("doc-{i}")),
);
index.add_metadata(doc_id, metadata).unwrap();
}
index.build().unwrap();
#[cfg(feature = "hnsw")]
{
index.coarse_quantizer = None;
}
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
let filter = crate::filtering::MetadataFilter::equals("group", 2i32);
let before = index.search_with_filter(&query, 10, &filter).unwrap();
assert!(!before.is_empty());
let dir = tempfile::tempdir().unwrap();
index.save_to_dir(dir.path()).unwrap();
let loaded = IVFPQIndex::load_from_dir(dir.path()).unwrap();
assert_eq!(
loaded.search_with_filter(&query, 10, &filter).unwrap(),
before
);
}
#[test]
fn file_searcher_matches_snapshot_loaded_search() {
let dim = 16;
let n = 240;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(103);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 16,
nprobe: 4,
seed: 125,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let vector: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(20_000 + i as u32, vector).unwrap();
}
index.build().unwrap();
#[cfg(feature = "hnsw")]
{
index.coarse_quantizer = None;
}
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
let dir = tempfile::tempdir().unwrap();
index.save_to_dir(dir.path()).unwrap();
let loaded = IVFPQIndex::load_from_dir(dir.path()).unwrap();
let mut file_searcher = IVFPQFileSearcher::load(dir.path()).unwrap();
assert!(file_searcher.list_codes.is_some());
assert_eq!(file_searcher.num_vectors(), loaded.num_vectors);
assert_eq!(file_searcher.nprobe(), loaded.nprobe());
assert_eq!(
file_searcher.search(&query, 10).unwrap(),
loaded.search(&query, 10).unwrap()
);
assert_eq!(
file_searcher.search_reranked(&query, 10, 80).unwrap(),
loaded.search_reranked(&query, 10, 80).unwrap()
);
}
#[test]
fn file_rerank_diagnostics_report_raw_vector_reads() {
let dim = 16;
let n = 240;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(108);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 32,
nprobe: 4,
seed: 130,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let vector: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(80_000 + i as u32, vector).unwrap();
}
index.build().unwrap();
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
let dir = tempfile::tempdir().unwrap();
index.save_to_dir(dir.path()).unwrap();
let mut file_searcher = IVFPQFileSearcher::load(dir.path()).unwrap();
let candidate_pool = 80;
let reranked = file_searcher
.search_reranked(&query, 10, candidate_pool)
.unwrap();
let (with_diagnostics, diagnostics) = file_searcher
.search_reranked_with_diagnostics(&query, 10, candidate_pool)
.unwrap();
assert_eq!(with_diagnostics, reranked);
assert!(diagnostics.raw_vector_reads >= with_diagnostics.len());
assert!(diagnostics.raw_vector_reads <= candidate_pool);
assert_eq!(
diagnostics.reranked_candidates,
diagnostics.raw_vector_reads
);
assert_eq!(
diagnostics.raw_vector_bytes,
diagnostics.raw_vector_reads * dim * std::mem::size_of::<f32>()
);
}
#[test]
fn file_approx_diagnostics_report_code_reads() {
let dim = 16;
let n = 240;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(109);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 32,
nprobe: 3,
seed: 131,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let vector: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(90_000 + i as u32, vector).unwrap();
}
index.build().unwrap();
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
let dir = tempfile::tempdir().unwrap();
index.save_to_dir(dir.path()).unwrap();
let mut file_searcher = IVFPQFileSearcher::load(dir.path()).unwrap();
let plain = file_searcher.search(&query, 10).unwrap();
let (with_diagnostics, diagnostics) =
file_searcher.search_with_diagnostics(&query, 10).unwrap();
assert_eq!(with_diagnostics, plain);
assert_eq!(diagnostics.probed_lists, 3);
assert!(diagnostics.scanned_vectors >= with_diagnostics.len());
assert!(diagnostics.code_reads > 0);
assert_eq!(
diagnostics.code_bytes,
diagnostics.scanned_vectors * file_searcher.num_codebooks()
);
assert_eq!(diagnostics.retained_candidates, with_diagnostics.len());
assert_eq!(diagnostics.raw_vector_reads, 0);
assert_eq!(diagnostics.raw_vector_bytes, 0);
}
#[test]
fn file_searcher_loads_old_snapshots_without_list_codes() {
let dim = 16;
let n = 240;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(106);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 32,
nprobe: 4,
seed: 128,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let vector: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(60_000 + i as u32, vector).unwrap();
}
index.build().unwrap();
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
let dir = tempfile::tempdir().unwrap();
index.save_to_dir(dir.path()).unwrap();
std::fs::remove_file(dir.path().join("list_offsets.bin")).unwrap();
std::fs::remove_file(dir.path().join("list_codes.bin")).unwrap();
let loaded = IVFPQIndex::load_from_dir(dir.path()).unwrap();
let mut file_searcher = IVFPQFileSearcher::load(dir.path()).unwrap();
assert!(file_searcher.list_codes.is_none());
assert_eq!(
file_searcher.search(&query, 10).unwrap(),
loaded.search(&query, 10).unwrap()
);
}
#[test]
fn file_searcher_rejects_partial_list_code_sidecar() {
let dim = 16;
let n = 240;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(107);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 32,
nprobe: 4,
seed: 129,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let vector: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(70_000 + i as u32, vector).unwrap();
}
index.build().unwrap();
let dir = tempfile::tempdir().unwrap();
index.save_to_dir(dir.path()).unwrap();
std::fs::remove_file(dir.path().join("list_codes.bin")).unwrap();
let err = match IVFPQFileSearcher::load(dir.path()) {
Ok(_) => panic!("partial list-code sidecar should fail to load"),
Err(err) => err,
};
assert!(
err.to_string().contains("partial IVF-PQ list-code sidecar"),
"unexpected error: {err}"
);
}
#[test]
fn file_searcher_matches_standard_adc_snapshot_loaded_search() {
let dim = 16;
let n = 240;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(105);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 32,
nprobe: 4,
seed: 127,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let vector: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(40_000 + i as u32, vector).unwrap();
}
index.build().unwrap();
#[cfg(feature = "hnsw")]
{
index.coarse_quantizer = None;
}
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
let dir = tempfile::tempdir().unwrap();
index.save_to_dir(dir.path()).unwrap();
let loaded = IVFPQIndex::load_from_dir(dir.path()).unwrap();
let mut file_searcher = IVFPQFileSearcher::load(dir.path()).unwrap();
assert_eq!(
file_searcher.search(&query, 10).unwrap(),
loaded.search(&query, 10).unwrap()
);
}
#[cfg(feature = "persistence")]
#[test]
fn mmap_searcher_matches_snapshot_loaded_search() {
let dim = 16;
let n = 240;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(104);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 16,
nprobe: 4,
seed: 126,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let vector: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(30_000 + i as u32, vector).unwrap();
}
index.build().unwrap();
#[cfg(feature = "hnsw")]
{
index.coarse_quantizer = None;
}
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
let dir = tempfile::tempdir().unwrap();
index.save_to_dir(dir.path()).unwrap();
let loaded = IVFPQIndex::load_from_dir(dir.path()).unwrap();
let mut mmap_searcher = IVFPQFileSearcher::load_mmap(dir.path()).unwrap();
assert_eq!(
mmap_searcher.search(&query, 10).unwrap(),
loaded.search(&query, 10).unwrap()
);
assert_eq!(
mmap_searcher.search_reranked(&query, 10, 80).unwrap(),
loaded.search_reranked(&query, 10, 80).unwrap()
);
}
#[test]
fn save_load_compacted_index_keeps_approximate_search_only() {
let dim = 16;
let n = 200;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(12);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 16,
nprobe: 4,
seed: 124,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let v: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(i as u32, v).unwrap();
}
index.build().unwrap();
#[cfg(feature = "hnsw")]
{
index.coarse_quantizer = None;
}
index.compact();
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
let approx_before = index.search(&query, 10).unwrap();
let dir = tempfile::tempdir().unwrap();
index.save_to_dir(dir.path()).unwrap();
let loaded = IVFPQIndex::load_from_dir(dir.path()).unwrap();
assert_eq!(loaded.search(&query, 10).unwrap(), approx_before);
let err = loaded.search_reranked(&query, 10, 80).unwrap_err();
assert!(
err.to_string().contains("search_reranked unavailable"),
"unexpected error: {err}"
);
let mut file_searcher = IVFPQFileSearcher::load(dir.path()).unwrap();
assert_eq!(file_searcher.search(&query, 10).unwrap(), approx_before);
let err = file_searcher.search_reranked(&query, 10, 80).unwrap_err();
assert!(
err.to_string()
.contains("search_reranked unavailable without raw_vectors.bin"),
"unexpected error: {err}"
);
}
#[test]
fn load_rejects_future_manifest_version() {
let dim = 16;
let n = 200;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(13);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 16,
nprobe: 4,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let v: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(i as u32, v).unwrap();
}
index.build().unwrap();
let dir = tempfile::tempdir().unwrap();
index.save_to_dir(dir.path()).unwrap();
let manifest_path = dir.path().join("manifest.json");
let mut manifest: serde_json::Value =
serde_json::from_slice(&std::fs::read(&manifest_path).unwrap()).unwrap();
manifest["version"] = serde_json::json!(IVFPQ_FORMAT_VERSION + 1);
std::fs::write(
&manifest_path,
serde_json::to_vec_pretty(&manifest).unwrap(),
)
.unwrap();
let err = IVFPQIndex::load_from_dir(dir.path()).unwrap_err();
assert!(
err.to_string()
.contains("unsupported IVF-PQ format version"),
"unexpected error: {err}"
);
}
#[test]
fn load_accepts_legacy_manifest_without_filter_fields() {
let dim = 16;
let n = 200;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(18);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 16,
nprobe: 4,
seed: 141,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let v: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(i as u32, v).unwrap();
}
index.build().unwrap();
#[cfg(feature = "hnsw")]
{
index.coarse_quantizer = None;
}
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
let before = index.search(&query, 10).unwrap();
let dir = tempfile::tempdir().unwrap();
index.save_to_dir(dir.path()).unwrap();
let manifest_path = dir.path().join("manifest.json");
let mut manifest: serde_json::Value =
serde_json::from_slice(&std::fs::read(&manifest_path).unwrap()).unwrap();
let manifest_obj = manifest.as_object_mut().unwrap();
manifest_obj.remove("filter_field");
manifest_obj.remove("filter_metadata");
std::fs::write(
&manifest_path,
serde_json::to_vec_pretty(&manifest).unwrap(),
)
.unwrap();
let loaded = IVFPQIndex::load_from_dir(dir.path()).unwrap();
assert_eq!(loaded.search(&query, 10).unwrap(), before);
}
#[test]
fn load_rejects_corrupt_cluster_magic() {
let dim = 16;
let n = 200;
use rand::{Rng, SeedableRng};
let mut rng = rand::rngs::StdRng::seed_from_u64(14);
let params = IVFPQParams {
num_clusters: 4,
num_codebooks: 4,
codebook_size: 16,
nprobe: 4,
..IVFPQParams::default()
};
let mut index = IVFPQIndex::new(dim, params).unwrap();
for i in 0..n {
let v: Vec<f32> = (0..dim).map(|_| rng.random::<f32>()).collect();
index.add(i as u32, v).unwrap();
}
index.build().unwrap();
let dir = tempfile::tempdir().unwrap();
index.save_to_dir(dir.path()).unwrap();
std::fs::write(dir.path().join("clusters.bin"), b"not ivfpq").unwrap();
let err = IVFPQIndex::load_from_dir(dir.path()).unwrap_err();
assert!(
err.to_string()
.contains("invalid IVF-PQ cluster file magic"),
"unexpected error: {err}"
);
}
}