use core::cmp::{min, Ordering};
use distances::Number;
use crate::{Cluster, Dataset, Tree};
use super::Hits;
pub fn search<T, U, D>(tree: &Tree<T, U, D>, query: T, k: usize) -> Vec<(usize, U)>
where
T: Send + Sync + Copy,
U: Number,
D: Dataset<T, U>,
{
let data = tree.data();
let c = tree.root();
let d = c.distance_to_instance(data, query);
let mut grains = vec![Grain::new(c, d)];
let mut hits = Hits::new(k);
loop {
let i = Grain::partition(&mut grains, k);
let threshold = grains[i].d;
hits.pop_until(threshold);
let (insiders, non_insiders) = grains.split_at_mut(i);
let (small_insiders, insiders) = insiders
.iter_mut()
.map(|&mut g| g)
.partition::<Vec<_>, _>(|g| g.is_small(k));
grains_to_hits(data, query, &mut hits, &small_insiders);
let (leaves, non_leaves) = non_insiders
.iter_mut()
.filter(|g| !g.is_outside(threshold))
.map(|&mut g| g)
.partition::<Vec<_>, _>(|g| g.is_leaf);
if non_leaves.is_empty() {
grains_to_hits(data, query, &mut hits, &insiders);
grains_to_hits(data, query, &mut hits, &leaves);
break;
}
grains = non_leaves
.into_iter()
.chain(insiders.into_iter())
.flat_map(|g| {
g.c.children()
.unwrap_or_else(|| unreachable!("This is only called on non-leaves."))
})
.map(|c| (c, c.distance_to_instance(data, query)))
.map(|(c, d)| Grain::new(c, d))
.chain(leaves.into_iter())
.collect();
}
hits.extract()
}
#[allow(clippy::needless_for_each)]
fn grains_to_hits<T, U, D>(data: &D, query: T, hits: &mut Hits<usize, U>, grains: &[Grain<T, U>])
where
T: Send + Sync + Copy,
U: Number,
D: Dataset<T, U>,
{
grains.iter().for_each(|g| {
let indices = g.c.indices(data);
let distances = data.query_to_many(query, indices);
hits.push_batch(indices.iter().copied().zip(distances.into_iter()));
});
}
#[derive(Debug, Clone, Copy)]
struct Grain<'a, T: Send + Sync + Copy, U: Number> {
c: &'a Cluster<T, U>,
d: U,
diameter: U,
multiplicity: usize,
is_leaf: bool,
}
impl<'a, T: Send + Sync + Copy, U: Number> Grain<'a, T, U> {
fn new(c: &'a Cluster<T, U>, d: U) -> Self {
let r = c.radius;
Self {
c,
d: d + r,
diameter: r + r,
multiplicity: c.cardinality,
is_leaf: c.is_leaf(),
}
}
pub const fn is_small(&self, k: usize) -> bool {
(self.multiplicity <= k) || self.is_leaf
}
pub fn is_outside(&self, threshold: U) -> bool {
if self.d <= self.diameter {
false
} else {
threshold < (self.d - self.diameter)
}
}
pub fn partition(grains: &mut [Self], k: usize) -> usize {
Self::_partition(grains, k, 0, grains.len() - 1)
}
#[allow(clippy::many_single_char_names)]
fn _partition(grains: &mut [Self], k: usize, l: usize, r: usize) -> usize {
if l >= r {
min(l, r)
} else {
let p = Self::_partition_once(grains, l, r);
let g = grains.iter().take(p).map(|g| g.multiplicity).sum::<usize>();
match g.cmp(&k) {
Ordering::Equal => p,
Ordering::Less => Self::_partition(grains, k, p + 1, r),
Ordering::Greater => {
if (p > 0) && (g > (k + grains[p - 1].multiplicity)) {
Self::_partition(grains, k, l, p - 1)
} else {
p
}
}
}
}
}
fn _partition_once(grains: &mut [Self], l: usize, r: usize) -> usize {
let pivot = l + (r - l) / 2;
grains.swap(pivot, r);
let (mut a, mut b) = (l, l);
while b < r {
if grains[b].d < grains[r].d {
grains.swap(a, b);
a += 1;
}
b += 1;
}
grains.swap(a, r);
a
}
}
#[cfg(test)]
mod tests {
use core::f32::EPSILON;
use distances::vectors::euclidean;
use symagen::random_data;
use crate::{cakes::knn::linear, knn::tests::sort_hits, Cakes, PartitionCriteria, VecDataset};
#[test]
fn sieve_v1() {
let (cardinality, dimensionality) = (10_000, 10);
let (min_val, max_val) = (-1.0, 1.0);
let seed = 42;
let data = random_data::random_f32(cardinality, dimensionality, min_val, max_val, seed);
let data = data.iter().map(Vec::as_slice).collect::<Vec<_>>();
let data = VecDataset::new("knn-test".to_string(), data, euclidean::<_, f32>, false);
let query = random_data::random_f32(1, dimensionality, min_val, max_val, seed * 2);
let query = query[0].as_slice();
let criteria = PartitionCriteria::default();
let model = Cakes::new(data, Some(seed), criteria);
let tree = model.tree();
for k in [100, 10, 1] {
let linear_nn = sort_hits(linear::search(tree.data(), query, k, tree.indices()));
let sieve_nn = sort_hits(super::search(tree, query, k));
assert_eq!(sieve_nn.len(), k);
let d_linear = linear_nn[k - 1].1;
let d_sieve = sieve_nn[k - 1].1;
assert!(
(d_linear - d_sieve).abs() < EPSILON,
"k = {}, linear = {}, sieve = {}",
k,
d_linear,
d_sieve
);
}
}
}