use crate::pq::ProductQuantizer;
use crate::cosine_similarity_prenorm;
pub struct IVFIndex {
pub centroids: Vec<f32>,
pub num_clusters: usize,
pub dim: usize,
pub inverted_lists: Vec<Vec<usize>>,
}
impl IVFIndex {
pub fn train(vectors: &[Vec<f32>], dim: usize, num_clusters: usize) -> Self {
let n = vectors.len();
let actual_clusters = num_clusters.min(n);
let mut centroids = vec![0.0f32; actual_clusters * dim];
let step = if n > actual_clusters { n / actual_clusters } else { 1 };
for c in 0..actual_clusters {
let src_idx = (c * step) % n;
let dst = &mut centroids[c * dim..(c + 1) * dim];
dst.copy_from_slice(&vectors[src_idx]);
}
let mut assignments = vec![0usize; n];
let max_iters = 15;
for _ in 0..max_iters {
let mut changed = false;
for (i, vec) in vectors.iter().enumerate() {
let best = nearest_centroid(vec, ¢roids, actual_clusters, dim);
if assignments[i] != best {
assignments[i] = best;
changed = true;
}
}
if !changed {
break;
}
let mut counts = vec![0u32; actual_clusters];
centroids.fill(0.0);
for (i, vec) in vectors.iter().enumerate() {
let c = assignments[i];
counts[c] += 1;
let offset = c * dim;
for d in 0..dim {
centroids[offset + d] += vec[d];
}
}
for (c, &count) in counts.iter().enumerate().take(actual_clusters) {
if count > 0 {
let offset = c * dim;
let cnt = count as f32;
for d in 0..dim {
centroids[offset + d] /= cnt;
}
}
}
}
let mut inverted_lists = vec![Vec::new(); actual_clusters];
for (i, &c) in assignments.iter().enumerate() {
inverted_lists[c].push(i);
}
Self {
centroids,
num_clusters: actual_clusters,
dim,
inverted_lists,
}
}
pub fn assign(&self, vector: &[f32]) -> usize {
nearest_centroid(vector, &self.centroids, self.num_clusters, self.dim)
}
pub fn search(
&self,
query: &[f32],
vectors: &[Vec<f32>],
norms: &[f32],
tombstones: &[u8],
nprobe: usize,
k: usize,
) -> Vec<(usize, f32)> {
let probe_clusters = self.nearest_clusters(query, nprobe);
let query_norm = rustyhdf5_accel::vector_norm(query);
let mut results: Vec<(usize, f32)> = Vec::new();
for cluster_id in probe_clusters {
for &idx in &self.inverted_lists[cluster_id] {
if idx < tombstones.len() && tombstones[idx] != 0 {
continue;
}
let vec_norm = if idx < norms.len() {
norms[idx]
} else {
rustyhdf5_accel::vector_norm(&vectors[idx])
};
let score =
cosine_similarity_prenorm(query, query_norm, &vectors[idx], vec_norm);
results.push((idx, score));
}
}
results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
results.truncate(k);
results
}
fn nearest_clusters(&self, query: &[f32], nprobe: usize) -> Vec<usize> {
let mut dists: Vec<(usize, f32)> = (0..self.num_clusters)
.map(|c| {
let centroid = &self.centroids[c * self.dim..(c + 1) * self.dim];
let sim = rustyhdf5_accel::cosine_similarity(query, centroid);
(c, sim)
})
.collect();
dists.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
dists.iter().take(nprobe).map(|&(c, _)| c).collect()
}
pub fn is_balanced(&self) -> bool {
if self.inverted_lists.is_empty() {
return true;
}
let total: usize = self.inverted_lists.iter().map(|l| l.len()).sum();
let avg = total as f32 / self.inverted_lists.len() as f32;
let max_size = self.inverted_lists.iter().map(|l| l.len()).max().unwrap_or(0);
max_size as f32 <= avg * 3.0
}
pub fn to_hdf5_data(&self) -> (&[f32], Vec<i64>, Vec<i64>, [i64; 2]) {
let mut offsets = Vec::with_capacity(self.num_clusters + 1);
let mut data = Vec::new();
let mut offset = 0i64;
for list in &self.inverted_lists {
offsets.push(offset);
for &idx in list {
data.push(idx as i64);
}
offset += list.len() as i64;
}
offsets.push(offset);
(
&self.centroids,
offsets,
data,
[self.num_clusters as i64, self.dim as i64],
)
}
pub fn from_hdf5_data(
centroids: Vec<f32>,
offsets: &[i64],
data: &[i64],
metadata: [i64; 2],
) -> Self {
let num_clusters = metadata[0] as usize;
let dim = metadata[1] as usize;
let mut inverted_lists = Vec::with_capacity(num_clusters);
for c in 0..num_clusters {
let start = offsets[c] as usize;
let end = offsets[c + 1] as usize;
let list: Vec<usize> = data[start..end].iter().map(|&v| v as usize).collect();
inverted_lists.push(list);
}
Self {
centroids,
num_clusters,
dim,
inverted_lists,
}
}
}
pub struct IVFPQIndex {
pub ivf: IVFIndex,
pub pq: ProductQuantizer,
pub codes: Vec<u8>,
}
impl IVFPQIndex {
pub fn build(
vectors: &[Vec<f32>],
dim: usize,
num_clusters: usize,
num_subvectors: usize,
num_centroids: usize,
) -> Self {
let ivf = IVFIndex::train(vectors, dim, num_clusters);
let pq = ProductQuantizer::train(vectors, dim, num_subvectors, num_centroids);
let codes = pq.encode_all(vectors);
Self { ivf, pq, codes }
}
#[allow(clippy::too_many_arguments)]
pub fn search(
&self,
query: &[f32],
vectors: &[Vec<f32>],
norms: &[f32],
tombstones: &[u8],
nprobe: usize,
candidates: usize,
k: usize,
) -> Vec<(usize, f32)> {
let probe_clusters = self.ivf.nearest_clusters(query, nprobe);
let table = self.pq.precompute_distance_table(query);
let mut pq_results: Vec<(usize, f32)> = Vec::new();
for cluster_id in probe_clusters {
for &idx in &self.ivf.inverted_lists[cluster_id] {
if idx < tombstones.len() && tombstones[idx] != 0 {
continue;
}
let code_start = idx * self.pq.num_subvectors;
let code_end = code_start + self.pq.num_subvectors;
let codes = &self.codes[code_start..code_end];
let dist = self.pq.asymmetric_distance_with_table(&table, codes);
pq_results.push((idx, dist));
}
}
pq_results.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
pq_results.truncate(candidates);
let query_norm = rustyhdf5_accel::vector_norm(query);
let mut reranked: Vec<(usize, f32)> = pq_results
.iter()
.map(|&(idx, _)| {
let vec_norm = if idx < norms.len() {
norms[idx]
} else {
rustyhdf5_accel::vector_norm(&vectors[idx])
};
(
idx,
cosine_similarity_prenorm(query, query_norm, &vectors[idx], vec_norm),
)
})
.collect();
reranked.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
reranked.truncate(k);
reranked
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SearchStrategy {
BruteForce,
BruteForceNorms,
IVFPQ,
}
pub fn auto_strategy(num_vectors: usize) -> SearchStrategy {
if num_vectors < 10_000 {
SearchStrategy::BruteForce
} else if num_vectors <= 100_000 {
SearchStrategy::BruteForceNorms
} else {
SearchStrategy::IVFPQ
}
}
fn nearest_centroid(vector: &[f32], centroids: &[f32], num_clusters: usize, dim: usize) -> usize {
let mut best = 0;
let mut best_sim = f32::NEG_INFINITY;
for c in 0..num_clusters {
let centroid = ¢roids[c * dim..(c + 1) * dim];
let sim = rustyhdf5_accel::cosine_similarity(vector, centroid);
if sim > best_sim {
best_sim = sim;
best = c;
}
}
best
}
#[cfg(test)]
mod tests {
use super::*;
fn make_vectors(n: usize, dim: usize, seed: u32) -> Vec<Vec<f32>> {
let mut s = seed;
let mut next = || -> f32 {
s = s.wrapping_mul(1103515245).wrapping_add(12345);
((s >> 16) as f32) / 65536.0 - 0.5
};
(0..n).map(|_| (0..dim).map(|_| next()).collect()).collect()
}
#[test]
fn ivf_clustering_produces_clusters() {
let dim = 32;
let vectors = make_vectors(200, dim, 42);
let ivf = IVFIndex::train(&vectors, dim, 10);
assert_eq!(ivf.num_clusters, 10);
let total: usize = ivf.inverted_lists.iter().map(|l| l.len()).sum();
assert_eq!(total, 200);
let non_empty = ivf.inverted_lists.iter().filter(|l| !l.is_empty()).count();
assert!(non_empty > 0);
}
#[test]
fn ivf_balanced_clusters() {
let dim = 32;
let vectors = make_vectors(1000, dim, 42);
let ivf = IVFIndex::train(&vectors, dim, 10);
assert!(ivf.is_balanced(), "clusters should be reasonably balanced");
}
#[test]
fn ivf_search_nprobe_all_matches_brute_force() {
let dim = 32;
let vectors = make_vectors(100, dim, 42);
let norms: Vec<f32> = vectors.iter().map(|v| rustyhdf5_accel::vector_norm(v)).collect();
let tombstones = vec![0u8; 100];
let query = vectors[0].clone();
let ivf = IVFIndex::train(&vectors, dim, 5);
let ivf_results = ivf.search(&query, &vectors, &norms, &tombstones, 5, 10);
let query_norm = rustyhdf5_accel::vector_norm(&query);
let mut brute: Vec<(usize, f32)> = vectors
.iter()
.enumerate()
.map(|(i, v)| {
(
i,
cosine_similarity_prenorm(&query, query_norm, v, norms[i]),
)
})
.collect();
brute.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
brute.truncate(10);
let ivf_ids: Vec<usize> = ivf_results.iter().map(|r| r.0).collect();
let brute_ids: Vec<usize> = brute.iter().map(|r| r.0).collect();
assert_eq!(ivf_ids, brute_ids, "nprobe=all should match brute force");
}
#[test]
fn ivf_search_nprobe_1_returns_results() {
let dim = 32;
let vectors = make_vectors(200, dim, 42);
let norms: Vec<f32> = vectors.iter().map(|v| rustyhdf5_accel::vector_norm(v)).collect();
let tombstones = vec![0u8; 200];
let query = vectors[0].clone();
let ivf = IVFIndex::train(&vectors, dim, 10);
let results = ivf.search(&query, &vectors, &norms, &tombstones, 1, 10);
assert!(!results.is_empty(), "nprobe=1 should still find results");
}
#[test]
fn ivf_pq_combined_search_recall() {
let dim = 64;
let n = 500;
let vectors = make_vectors(n, dim, 42);
let norms: Vec<f32> = vectors.iter().map(|v| rustyhdf5_accel::vector_norm(v)).collect();
let tombstones = vec![0u8; n];
let query = vectors[0].clone();
let index = IVFPQIndex::build(&vectors, dim, 10, 8, 64);
let results = index.search(&query, &vectors, &norms, &tombstones, 5, 100, 10);
let query_norm = rustyhdf5_accel::vector_norm(&query);
let mut exact: Vec<(usize, f32)> = vectors
.iter()
.enumerate()
.map(|(i, v)| {
(
i,
cosine_similarity_prenorm(&query, query_norm, v, norms[i]),
)
})
.collect();
exact.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
let exact_top10: Vec<usize> = exact.iter().take(10).map(|r| r.0).collect();
let ivfpq_ids: Vec<usize> = results.iter().map(|r| r.0).collect();
let overlap = exact_top10.iter().filter(|i| ivfpq_ids.contains(i)).count();
assert!(
overlap >= 8,
"IVF-PQ recall too low: {overlap}/10 overlap"
);
}
#[test]
fn auto_strategy_selection() {
assert_eq!(auto_strategy(100), SearchStrategy::BruteForce);
assert_eq!(auto_strategy(9_999), SearchStrategy::BruteForce);
assert_eq!(auto_strategy(10_000), SearchStrategy::BruteForceNorms);
assert_eq!(auto_strategy(50_000), SearchStrategy::BruteForceNorms);
assert_eq!(auto_strategy(100_000), SearchStrategy::BruteForceNorms);
assert_eq!(auto_strategy(100_001), SearchStrategy::IVFPQ);
}
#[test]
fn ivf_hdf5_roundtrip() {
let dim = 16;
let vectors = make_vectors(50, dim, 42);
let ivf = IVFIndex::train(&vectors, dim, 5);
let (centroids, offsets, data, meta) = ivf.to_hdf5_data();
let ivf2 = IVFIndex::from_hdf5_data(centroids.to_vec(), &offsets, &data, meta);
assert_eq!(ivf.num_clusters, ivf2.num_clusters);
assert_eq!(ivf.dim, ivf2.dim);
for c in 0..ivf.num_clusters {
assert_eq!(ivf.inverted_lists[c], ivf2.inverted_lists[c]);
}
}
#[test]
fn ivf_assign_consistent() {
let dim = 16;
let vectors = make_vectors(100, dim, 42);
let ivf = IVFIndex::train(&vectors, dim, 5);
for (i, v) in vectors.iter().enumerate() {
let cluster = ivf.assign(v);
assert!(
ivf.inverted_lists[cluster].contains(&i),
"vector {i} should be in cluster {cluster}"
);
}
}
#[test]
fn ivf_respects_tombstones() {
let dim = 16;
let vectors = make_vectors(50, dim, 42);
let norms: Vec<f32> = vectors.iter().map(|v| rustyhdf5_accel::vector_norm(v)).collect();
let mut tombstones = vec![0u8; 50];
tombstones[0] = 1;
tombstones[1] = 1;
let ivf = IVFIndex::train(&vectors, dim, 5);
let results = ivf.search(&vectors[2], &vectors, &norms, &tombstones, 5, 50);
assert!(results.iter().all(|r| r.0 != 0 && r.0 != 1));
}
}