use alloc::collections::BTreeMap;
use alloc::vec::Vec;
pub const MIN_ROWS_FOR_INDEX: usize = 1000;
pub const DEFAULT_NPROBE: usize = 8;
const KMEANS_ITERS: usize = 4;
#[derive(Debug, Clone)]
pub struct IvfFlatIndex {
dim: usize,
centroids: Vec<Vec<f32>>,
lists: Vec<Vec<(u64, Vec<f32>)>>,
id_to_list: BTreeMap<u64, usize>,
}
fn dot(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b).map(|(x, y)| x * y).sum()
}
fn sqrt_f32(x: f32) -> f32 {
if x <= 0.0 {
return 0.0;
}
let i = x.to_bits();
let i = 0x5f37_59df - (i >> 1);
let mut y = f32::from_bits(i);
y *= 1.5 - 0.5 * x * y * y;
y *= 1.5 - 0.5 * x * y * y;
y *= 1.5 - 0.5 * x * y * y;
y *= 1.5 - 0.5 * x * y * y;
x * y
}
fn normalize(v: &[f32]) -> Vec<f32> {
let norm = sqrt_f32(dot(v, v));
if norm == 0.0 {
v.to_vec()
} else {
v.iter().map(|x| x / norm).collect()
}
}
fn nearest_cluster(centroids: &[Vec<f32>], unit_v: &[f32]) -> usize {
centroids
.iter()
.enumerate()
.map(|(i, c)| (i, dot(c, unit_v)))
.max_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(core::cmp::Ordering::Equal))
.map(|(i, _)| i)
.unwrap_or(0)
}
fn nlist_for(n: usize) -> usize {
let mut lo = 0usize;
let mut hi = n;
while lo < hi {
let mid = lo + (hi - lo).div_ceil(2);
if mid.checked_mul(mid).is_some_and(|sq| sq <= n) {
lo = mid;
} else {
hi = mid - 1;
}
}
lo.clamp(1, 256)
}
fn seed_centroids(unit_vectors: &[Vec<f32>], nlist: usize) -> Vec<Vec<f32>> {
let n = unit_vectors.len();
(0..nlist)
.map(|i| unit_vectors[i * n / nlist].clone())
.collect()
}
impl IvfFlatIndex {
pub fn build(vectors: &[(u64, Vec<f32>)]) -> Self {
assert!(
!vectors.is_empty(),
"cannot build an index over zero vectors"
);
let dim = vectors[0].1.len();
let nlist = nlist_for(vectors.len());
let unit_vectors: Vec<Vec<f32>> = vectors.iter().map(|(_, v)| normalize(v)).collect();
let mut centroids = seed_centroids(&unit_vectors, nlist);
for _ in 0..KMEANS_ITERS {
let mut sums = alloc::vec![alloc::vec![0f32; dim]; nlist];
let mut counts = alloc::vec![0usize; nlist];
for uv in &unit_vectors {
let c = nearest_cluster(¢roids, uv);
for (s, x) in sums[c].iter_mut().zip(uv.iter()) {
*s += x;
}
counts[c] += 1;
}
for c in 0..nlist {
if counts[c] > 0 {
for (centroid_dim, sum) in centroids[c].iter_mut().zip(sums[c].iter()) {
*centroid_dim = sum / counts[c] as f32;
}
centroids[c] = normalize(¢roids[c]);
}
}
}
let mut lists: Vec<Vec<(u64, Vec<f32>)>> = alloc::vec![Vec::new(); nlist];
let mut id_to_list = BTreeMap::new();
for ((id, v), uv) in vectors.iter().zip(unit_vectors.iter()) {
let c = nearest_cluster(¢roids, uv);
lists[c].push((*id, v.clone()));
id_to_list.insert(*id, c);
}
Self {
dim,
centroids,
lists,
id_to_list,
}
}
pub fn dim(&self) -> usize {
self.dim
}
pub fn len(&self) -> usize {
self.id_to_list.len()
}
pub fn is_empty(&self) -> bool {
self.id_to_list.is_empty()
}
pub fn insert(&mut self, id: u64, v: &[f32]) {
self.remove(id);
let c = nearest_cluster(&self.centroids, &normalize(v));
self.lists[c].push((id, v.to_vec()));
self.id_to_list.insert(id, c);
}
pub fn remove(&mut self, id: u64) {
if let Some(c) = self.id_to_list.remove(&id) {
self.lists[c].retain(|(rid, _)| *rid != id);
}
}
pub fn search(&self, query: &[f32], k: usize, nprobe: usize) -> Vec<u64> {
if self.centroids.is_empty() {
return Vec::new();
}
let nprobe = nprobe.clamp(1, self.centroids.len());
let unit_query = normalize(query);
let mut ranked_centroids: Vec<(usize, f32)> = self
.centroids
.iter()
.enumerate()
.map(|(i, c)| (i, dot(c, &unit_query)))
.collect();
ranked_centroids
.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(core::cmp::Ordering::Equal));
let mut scored: Vec<(u64, f32)> = Vec::new();
for &(c, _) in ranked_centroids.iter().take(nprobe) {
for (id, v) in &self.lists[c] {
scored.push((*id, dot(v, query)));
}
}
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(core::cmp::Ordering::Equal));
scored.into_iter().take(k).map(|(id, _)| id).collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn corner_vectors(n: usize, dim: usize) -> Vec<(u64, Vec<f32>)> {
(0..n)
.map(|i| {
let mut v = alloc::vec![0f32; dim];
v[i % dim] = 1.0;
(i as u64, v)
})
.collect()
}
#[test]
fn finds_exact_match() {
let vectors = corner_vectors(2000, 32);
let idx = IvfFlatIndex::build(&vectors);
let query = vectors[7].1.clone();
let top = idx.search(&query, 1, 32);
assert_eq!(top[0], 7);
}
fn unique_one_hot_vectors(n: usize) -> Vec<(u64, Vec<f32>)> {
corner_vectors(n, n)
}
#[test]
fn insert_then_search_finds_it() {
let mut vectors = unique_one_hot_vectors(300);
let held_out = vectors.pop().unwrap();
let mut idx = IvfFlatIndex::build(&vectors);
assert_eq!(idx.len(), 299);
idx.insert(held_out.0, &held_out.1);
assert_eq!(idx.len(), 300);
let top = idx.search(&held_out.1, 1, 256);
assert_eq!(top[0], held_out.0);
}
#[test]
fn remove_drops_id_from_results() {
let vectors = corner_vectors(1200, 16);
let mut idx = IvfFlatIndex::build(&vectors);
let target = vectors[3].clone();
idx.remove(target.0);
assert_eq!(idx.len(), 1199);
let top = idx.search(&target.1, 5, 16);
assert!(!top.contains(&target.0));
}
#[test]
fn reinsert_moves_between_clusters() {
let vectors = corner_vectors(300, 301);
let mut idx = IvfFlatIndex::build(&vectors);
let mut moved = alloc::vec![0f32; 301];
moved[300] = 1.0;
idx.insert(0, &moved);
assert_eq!(idx.len(), 300);
let top = idx.search(&moved, 1, 256);
assert_eq!(top[0], 0);
}
}