use crate::index::ann::backend::{AnnBackend, AnnBackendCheckpoint, BackendMetric};
use crate::index::hnsw::cosine_distance;
use crate::query::AiExecutionContext;
use crate::rowid::RowId;
use crate::schema::IvfOptions;
use crate::Result;
use std::collections::BTreeMap;
#[derive(Clone, serde::Serialize, serde::Deserialize)]
pub(crate) struct IvfBackend {
dim: usize,
nlist: usize,
nprobe: usize,
training_samples: usize,
centroids: Option<Vec<Vec<f32>>>,
lists: BTreeMap<usize, Vec<(RowId, Vec<f32>)>>,
pending: BTreeMap<RowId, Vec<f32>>,
seed: u64,
}
type FrozenIvf = (Vec<Vec<f32>>, BTreeMap<usize, Vec<(RowId, Vec<f32>)>>);
impl IvfBackend {
pub(crate) fn new(dim: usize, options: &IvfOptions, seed: u64) -> Self {
Self {
dim,
nlist: options.nlist,
nprobe: options.nprobe,
training_samples: options.training_samples,
centroids: None,
lists: BTreeMap::new(),
pending: BTreeMap::new(),
seed,
}
}
fn freeze_active(&self) -> Option<FrozenIvf> {
self.freeze_active_with_checkpoint(&mut || Ok(()))
.expect("infallible IVF-training checkpoint")
}
fn freeze_active_with_checkpoint(
&self,
checkpoint: &mut dyn FnMut() -> Result<()>,
) -> Result<Option<FrozenIvf>> {
if self.pending.is_empty() {
return Ok(None);
}
let all_samples: Vec<&[f32]> = self.pending.values().map(|v| v.as_slice()).collect();
let samples: Vec<&[f32]> = if all_samples.len() > self.training_samples {
let stride = all_samples.len() / self.training_samples;
let start = (splitmix64(self.seed) as usize) % stride.max(1);
(0..self.training_samples)
.map(|i| all_samples[(start + i * stride) % all_samples.len()])
.collect()
} else {
all_samples
};
let effective_nlist = self.nlist.min(samples.len());
let centroids = kmeans(&samples, self.dim, effective_nlist, self.seed, checkpoint)?;
let mut lists: BTreeMap<usize, Vec<(RowId, Vec<f32>)>> = BTreeMap::new();
for (index, (row_id, vec)) in self.pending.iter().enumerate() {
if index.is_multiple_of(64) {
checkpoint()?;
}
let cell = nearest_centroid(vec, ¢roids);
lists.entry(cell).or_default().push((*row_id, vec.clone()));
}
Ok(Some((centroids, lists)))
}
pub(crate) fn from_checkpoint(
dim: usize,
nlist: usize,
nprobe: usize,
training_samples: usize,
centroids: Vec<Vec<f32>>,
lists: BTreeMap<usize, Vec<(RowId, Vec<f32>)>>,
seed: u64,
) -> std::result::Result<Self, String> {
if dim == 0
|| nlist == 0
|| nprobe == 0
|| nprobe > nlist
|| training_samples == 0
|| centroids.is_empty()
|| centroids.len() > nlist
|| centroids
.iter()
.any(|centroid| centroid.len() != dim || centroid.iter().any(|v| !v.is_finite()))
|| lists.iter().any(|(cell, rows)| {
*cell >= centroids.len()
|| rows.iter().any(|(_, vector)| {
vector.len() != dim || vector.iter().any(|v| !v.is_finite())
})
})
{
return Err("ANN IVF checkpoint contains invalid centroids or lists".into());
}
Ok(Self {
dim,
nlist,
nprobe,
training_samples,
centroids: Some(centroids),
lists,
pending: BTreeMap::new(),
seed,
})
}
}
impl AnnBackend for IvfBackend {
fn metric(&self) -> BackendMetric {
BackendMetric::Cosine
}
fn len(&self) -> usize {
self.lists.values().map(|l| l.len()).sum::<usize>() + self.pending.len()
}
fn is_empty(&self) -> bool {
self.lists.values().all(|l| l.is_empty()) && self.pending.is_empty()
}
fn insert_validated(
&mut self,
vec: &[f32],
row_id: RowId,
_checkpoint: &mut dyn FnMut() -> Result<()>,
) -> Result<()> {
if self.centroids.is_some() {
self.centroids = None;
self.lists.clear();
}
self.pending.insert(row_id, vec.to_vec());
Ok(())
}
fn finalize(&mut self, checkpoint: &mut dyn FnMut() -> Result<()>) -> Result<()> {
if self.centroids.is_none() {
if let Some((centroids, lists)) = self.freeze_active_with_checkpoint(checkpoint)? {
self.centroids = Some(centroids);
self.lists = lists;
self.pending.clear();
}
}
checkpoint()
}
fn search(
&self,
query: &[f32],
k: usize,
_ef: usize,
context: Option<&AiExecutionContext>,
) -> Result<Vec<(RowId, f64)>> {
let mut scored: Vec<(f32, RowId)> = Vec::new();
if let Some(centroids) = &self.centroids {
if !centroids.is_empty() {
if let Some(context) = context {
context.consume(crate::query::work_units(
self.dim.saturating_mul(centroids.len()),
crate::query::FLOAT_WORK_QUANTUM,
))?;
}
let mut centroid_dists: Vec<(f32, usize)> = centroids
.iter()
.enumerate()
.map(|(i, c)| (cosine_distance(query, c), i))
.collect();
centroid_dists.sort_by(|(da, _), (db, _)| da.total_cmp(db));
let probes = self.nprobe.min(centroids.len());
for (_, cell) in centroid_dists.into_iter().take(probes) {
if let Some(list) = self.lists.get(&cell) {
for (i, (row_id, vec)) in list.iter().enumerate() {
if let Some(context) = context {
if i.is_multiple_of(64) {
let count = (list.len() - i).min(64);
context.consume(crate::query::work_units(
self.dim.saturating_mul(count),
crate::query::FLOAT_WORK_QUANTUM,
))?;
}
}
scored.push((cosine_distance(query, vec), *row_id));
}
}
}
}
}
for (i, (row_id, vec)) in self.pending.iter().enumerate() {
if let Some(context) = context {
if i.is_multiple_of(64) {
let count = (self.pending.len() - i).min(64);
context.consume(crate::query::work_units(
self.dim.saturating_mul(count),
crate::query::FLOAT_WORK_QUANTUM,
))?;
}
}
scored.push((cosine_distance(query, vec), *row_id));
}
scored.sort_by(|(da, ra), (db, rb)| da.total_cmp(db).then_with(|| ra.cmp(rb)));
Ok(scored
.into_iter()
.take(k)
.map(|(dist, row_id)| (row_id, f64::from(dist)))
.collect())
}
fn entries(&self) -> Vec<(Vec<u8>, RowId)> {
let mut out: Vec<(Vec<u8>, RowId)> = Vec::new();
for (row_id, vec) in &self.pending {
out.push((vec_to_bytes(vec), *row_id));
}
for list in self.lists.values() {
for (row_id, vec) in list {
out.push((vec_to_bytes(vec), *row_id));
}
}
out
}
fn freeze(&self) -> AnnBackendCheckpoint {
let (centroids, lists) = if let Some(centroids) = &self.centroids {
(centroids.clone(), self.lists.clone())
} else if let Some((centroids, lists)) = self.freeze_active() {
(centroids, lists)
} else {
(vec![vec![0.0f32; self.dim]], BTreeMap::new())
};
AnnBackendCheckpoint::Ivf {
dim: self.dim,
nlist: self.nlist,
nprobe: self.nprobe,
centroids,
lists,
seed: self.seed,
}
}
fn empty_active(&self) -> Box<dyn AnnBackend> {
Box::new(Self::new(
self.dim,
&IvfOptions {
nlist: self.nlist,
nprobe: self.nprobe,
training_samples: self.training_samples,
},
self.seed,
))
}
fn rebuild_from_entries(&self, entries: &[(Vec<u8>, RowId)]) -> Box<dyn AnnBackend> {
let mut pending = BTreeMap::new();
for (bytes, row_id) in entries {
if bytes.len() >= self.dim * 4 {
pending.insert(*row_id, vec_from_bytes(bytes, self.dim));
}
}
let mut rebuilt = Self {
dim: self.dim,
nlist: self.nlist,
nprobe: self.nprobe,
training_samples: self.training_samples,
centroids: None,
lists: BTreeMap::new(),
pending,
seed: self.seed,
};
if let Some((centroids, lists)) = rebuilt.freeze_active() {
rebuilt.centroids = Some(centroids);
rebuilt.lists = lists;
rebuilt.pending.clear();
}
Box::new(rebuilt)
}
fn clone_box(&self) -> Box<dyn AnnBackend> {
Box::new(self.clone())
}
}
fn kmeans(
samples: &[&[f32]],
dim: usize,
k: usize,
seed: u64,
checkpoint: &mut dyn FnMut() -> Result<()>,
) -> Result<Vec<Vec<f32>>> {
if samples.is_empty() {
return Ok(vec![vec![0.0f32; dim]]);
}
let effective_k = k.min(samples.len()).max(1);
let mut centroids = vec![vec![0.0f32; dim]; k];
let start = (splitmix64(seed) as usize) % samples.len();
for c in 0..effective_k {
let src = samples[(start + c * (samples.len() / effective_k).max(1)) % samples.len()];
centroids[c] = src.to_vec();
}
for _iter in 0..25 {
checkpoint()?;
let mut sums = vec![vec![0.0f32; dim]; k];
let mut counts = vec![0u32; k];
for (index, sample) in samples.iter().enumerate() {
if index.is_multiple_of(64) {
checkpoint()?;
}
let nearest = nearest_centroid(sample, ¢roids);
for (i, value) in sample.iter().enumerate() {
sums[nearest][i] += value;
}
counts[nearest] += 1;
}
let mut moved = 0.0f32;
for c in 0..k {
if counts[c] > 0 {
let n = counts[c] as f32;
for i in 0..dim {
let new = sums[c][i] / n;
moved = moved.max((centroids[c][i] - new).abs());
centroids[c][i] = new;
}
}
}
if moved < 1e-6 {
break;
}
}
Ok(centroids)
}
fn nearest_centroid(vec: &[f32], centroids: &[Vec<f32>]) -> usize {
let mut best = 0usize;
let mut best_dist = f32::INFINITY;
for (i, c) in centroids.iter().enumerate() {
let dist = cosine_distance(vec, c);
if dist < best_dist {
best_dist = dist;
best = i;
}
}
best
}
fn vec_to_bytes(vec: &[f32]) -> Vec<u8> {
let mut bytes = Vec::with_capacity(vec.len() * 4);
for value in vec {
bytes.extend_from_slice(&value.to_le_bytes());
}
bytes
}
fn vec_from_bytes(bytes: &[u8], dim: usize) -> Vec<f32> {
(0..dim)
.map(|i| {
let offset = i * 4;
f32::from_le_bytes([
bytes[offset],
bytes[offset + 1],
bytes[offset + 2],
bytes[offset + 3],
])
})
.collect()
}
fn splitmix64(mut z: u64) -> u64 {
z = z.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut z = z;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
#[cfg(test)]
mod tests {
use super::*;
fn backend(dim: usize, nlist: usize, nprobe: usize) -> IvfBackend {
IvfBackend::new(
dim,
&IvfOptions {
nlist,
nprobe,
..Default::default()
},
0x9E37_79B9_7F4A_7C15,
)
}
#[test]
fn empty_backend_search_returns_nothing() {
let b = backend(8, 4, 2);
assert!(b.search(&[1.0; 8], 5, 0, None).unwrap().is_empty());
assert!(b.is_empty());
}
#[test]
fn finds_nearest_exact_match() {
let mut b = backend(8, 4, 2);
b.insert_validated(
&[1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
RowId(0),
&mut || Ok(()),
)
.unwrap();
b.insert_validated(
&[0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
RowId(1),
&mut || Ok(()),
)
.unwrap();
let top = b
.search(&[1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0], 1, 0, None)
.unwrap();
assert_eq!(top[0].0, RowId(0));
}
#[test]
fn recall_at_10_against_brute_force() {
let n = 200;
let dim = 16;
let mut b = backend(dim, 16, 8);
let mut data: Vec<(Vec<f32>, RowId)> = Vec::new();
let mut seed = 2024u64;
for i in 0..n {
let mut v = vec![0f32; dim];
for x in v.iter_mut() {
seed = seed
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
let u = ((seed >> 33) as u32) as f32 / (u32::MAX as f32);
*x = u * 2.0 - 1.0;
}
data.push((v.clone(), RowId(i as u64)));
b.insert_validated(&v, RowId(i as u64), &mut || Ok(()))
.unwrap();
}
let brute_topk = |q: &[f32], k: usize| -> std::collections::HashSet<u64> {
let mut s: Vec<(f32, u64)> = data
.iter()
.map(|(v, rid)| (cosine_distance(q, v), rid.0))
.collect();
s.sort_by(|(da, ra), (db, rb)| da.total_cmp(db).then_with(|| ra.cmp(rb)));
s.into_iter().take(k).map(|(_, r)| r).collect()
};
let mut total_recall = 0.0;
let queries = 20;
for qi in 0..queries {
let q = data[qi * 9 % n].0.clone();
let truth = brute_topk(&q, 10);
let got: std::collections::HashSet<u64> = b
.search(&q, 10, 0, None)
.unwrap()
.into_iter()
.map(|(r, _)| r.0)
.collect();
total_recall += truth.intersection(&got).count() as f64 / 10.0;
}
let avg = total_recall / queries as f64;
assert!(avg >= 0.85, "IVF recall@10 too low: {avg:.2}");
}
#[test]
fn nlist_exceeding_samples_degrades_gracefully() {
let mut b = backend(8, 256, 4);
for i in 0..5u64 {
let mut v = vec![0f32; 8];
v[i as usize] = 1.0;
b.insert_validated(&v, RowId(i), &mut || Ok(())).unwrap();
}
let top = b
.search(&[1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0], 3, 0, None)
.unwrap();
assert_eq!(top.len(), 3);
assert_eq!(top[0].0, RowId(0));
}
}