use super::all_pairs::{by_distance_then_index, keep_smallest, sort_key};
use crate::matrix::kmeans::{
kmeans_rows_seeded, nearest_centroid_rows, KmeansMetric, KmeansRowsOpts,
};
use crate::matrix::knn::metric::sqdist_soa_range;
use crate::matrix::utils::generate_minibatch_intervals;
use log::info;
use nalgebra::DMatrix;
use rand::rngs::StdRng;
use rand::SeedableRng;
use rayon::prelude::*;
pub const DEFAULT_N_PROBE: usize = 16;
const PARTITION_ITER: usize = 10;
const PARTITION_ROWS_PER_CELL: usize = 256;
const PARTITION_INIT_PER_CELL: usize = 16;
const PARTITION_MIN_CHANGED: f64 = 1e-3;
const QUERY_BLOCK: usize = 256;
#[derive(Debug, Clone)]
pub struct IvfArgs {
pub k: usize,
pub n_lists: usize,
pub n_probe: usize,
pub seed: u64,
}
pub fn knn_rows_ivf(x: &DMatrix<f32>, args: &IvfArgs) -> (Vec<Vec<usize>>, Vec<Vec<f32>>) {
let (n, d) = (x.nrows(), x.ncols());
let k = args.k.min(n.saturating_sub(1));
if n == 0 || k == 0 {
return (vec![Vec::new(); n], vec![Vec::new(); n]);
}
let n_lists = if args.n_lists == 0 {
(n as f64).sqrt().ceil() as usize
} else {
args.n_lists
}
.clamp(1, n);
let t_partition = std::time::Instant::now();
let train_rows = (PARTITION_ROWS_PER_CELL * n_lists).min(n);
let opts = KmeansRowsOpts {
k: n_lists,
max_iter: PARTITION_ITER,
seed: args.seed,
metric: KmeansMetric::Euclidean,
min_changed_frac: PARTITION_MIN_CHANGED,
init_sample: PARTITION_INIT_PER_CELL * n_lists,
};
let fit = if train_rows < n {
let mut rng = StdRng::seed_from_u64(args.seed);
let mut ids = rand::seq::index::sample(&mut rng, n, train_rows).into_vec();
ids.sort_unstable();
let mut fit = kmeans_rows_seeded(&x.select_rows(&ids), &opts);
fit.labels = nearest_centroid_rows(x, &fit.centroids, opts.metric);
fit
} else {
kmeans_rows_seeded(x, &opts)
};
let labels = fit.labels;
let n_probe = args.n_probe.clamp(1, n_lists);
let mut order: Vec<u32> = (0..n as u32).collect();
order.par_sort_unstable_by_key(|&i| (labels[i as usize], i));
let mut offsets = vec![0usize; n_lists + 1];
for &l in &labels {
offsets[l + 1] += 1;
}
for c in 0..n_lists {
offsets[c + 1] += offsets[c];
}
let widest_cell = (0..n_lists)
.map(|c| offsets[c + 1] - offsets[c])
.max()
.unwrap_or(0);
let mut soa = vec![0f32; n * d];
soa.par_chunks_mut(n).enumerate().for_each(|(dim, dst)| {
let col = x.column(dim);
for (v, &i) in dst.iter_mut().zip(&order) {
*v = col[i as usize];
}
});
let cents_soa: &[f32] = fit.centroids.as_slice();
info!(
"IVF kNN: {n} rows x {d} into {n_lists} cells fitted on {train_rows} rows in {} Lloyd \
iterations ({:.1} s); probing {n_probe} per query",
fit.n_iter,
t_partition.elapsed().as_secs_f64()
);
let t_search = std::time::Instant::now();
let blocks = generate_minibatch_intervals(n, 0, Some(QUERY_BLOCK));
let bar = crate::matrix::progress::new_progress_bar(blocks.len() as u64).with_message(format!(
"IVF kNN {n} x {d}, k={k}, {n_lists} cells, {n_probe} probed"
));
let index = Index {
soa: &soa,
n,
d,
cents_soa,
n_lists,
order: &order,
offsets: &offsets,
};
let found: Vec<(Vec<usize>, Vec<f32>)> = blocks
.into_par_iter()
.map_init(
|| Scratch::new(d, widest_cell.max(n_lists), n_probe, k),
|scratch, (p0, p1)| {
let out = search_block(&index, p0, p1, n_probe, k, scratch);
bar.inc(1);
out
},
)
.flatten()
.collect();
bar.finish_and_clear();
info!(
"IVF kNN: searched {n} queries in {:.1} s",
t_search.elapsed().as_secs_f64()
);
let mut indices = vec![Vec::new(); n];
let mut distances = vec![Vec::new(); n];
for (p, (nb, ds)) in found.into_iter().enumerate() {
let i = order[p] as usize;
indices[i] = nb;
distances[i] = ds;
}
(indices, distances)
}
struct Index<'a> {
soa: &'a [f32],
n: usize,
d: usize,
cents_soa: &'a [f32],
n_lists: usize,
order: &'a [u32],
offsets: &'a [usize],
}
struct Scratch {
q: Vec<f32>,
dist: Vec<f32>,
probes: Vec<(f32, usize)>,
hits: Vec<(usize, usize)>,
cand: Vec<Vec<(f32, usize)>>,
}
impl Scratch {
fn new(d: usize, widest: usize, n_probe: usize, k: usize) -> Self {
Self {
q: vec![0f32; QUERY_BLOCK * d],
dist: vec![0f32; widest],
probes: Vec::with_capacity(n_probe + 1),
hits: Vec::with_capacity(QUERY_BLOCK * n_probe),
cand: (0..QUERY_BLOCK)
.map(|_| Vec::with_capacity(k + 1))
.collect(),
}
}
}
fn search_block(
index: &Index<'_>,
p0: usize,
p1: usize,
n_probe: usize,
k: usize,
scratch: &mut Scratch,
) -> Vec<(Vec<usize>, Vec<f32>)> {
let (n, d) = (index.n, index.d);
let b = p1 - p0;
let Scratch {
q,
dist,
probes,
hits,
cand,
} = scratch;
hits.clear();
for (i, qi) in q[..b * d].chunks_exact_mut(d).enumerate() {
for (dim, qd) in qi.iter_mut().enumerate() {
*qd = index.soa[dim * n + p0 + i];
}
sqdist_soa_range(index.cents_soa, index.n_lists, qi, 0, index.n_lists, dist);
probes.clear();
for (c, &dd) in dist[..index.n_lists].iter().enumerate() {
keep_smallest(probes, (sort_key(dd), c), n_probe);
}
hits.extend(probes.iter().map(|&(_, c)| (c, i)));
cand[i].clear();
}
hits.sort_unstable();
for group in hits.chunk_by(|a, b| a.0 == b.0) {
let c = group[0].0;
let (lo, hi) = (index.offsets[c], index.offsets[c + 1]);
for &(_, i) in group {
let p = p0 + i;
sqdist_soa_range(index.soa, n, &q[i * d..(i + 1) * d], lo, hi, dist);
let list = &mut cand[i];
for (j, &dd) in dist[..hi - lo].iter().enumerate() {
let pos = lo + j;
if pos != p && (list.len() < k || dd < list[k - 1].0) {
keep_smallest(list, (sort_key(dd), index.order[pos] as usize), k);
}
}
}
}
cand[..b]
.iter_mut()
.map(|list| {
list.sort_unstable_by(by_distance_then_index);
list.iter().map(|&(d2, j)| (j, d2.sqrt())).unzip()
})
.collect()
}
#[cfg(test)]
#[path = "ivf_tests.rs"]
mod tests;