use core::f64::EPSILON;
use distances::Number;
use crate::{cakes::rnn::clustered::tree_search, utils, Tree};
use super::Hits;
const MULTIPLIER: f64 = 2.0;
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: crate::Dataset<T, U>,
{
let mut radius = EPSILON + tree.radius().as_f64() / tree.cardinality().as_f64();
let [mut confirmed, mut straddlers] = tree_search(tree.data(), tree.root(), query, U::from(radius));
let mut num_hits = confirmed
.iter()
.chain(straddlers.iter())
.map(|&(c, _)| c.cardinality)
.sum::<usize>();
while num_hits == 0 {
radius *= MULTIPLIER;
[confirmed, straddlers] = tree_search(tree.data(), tree.root(), query, U::from(radius));
num_hits = confirmed
.iter()
.chain(straddlers.iter())
.map(|&(c, _)| c.cardinality)
.sum::<usize>();
}
while num_hits < k {
let lfd = utils::mean(
&confirmed
.iter()
.chain(straddlers.iter())
.map(|&(c, _)| c.lfd)
.collect::<Vec<_>>(),
);
let factor = (k.as_f64() / num_hits.as_f64()).powf(1. / (lfd + EPSILON));
radius *= if factor < MULTIPLIER { factor } else { MULTIPLIER };
[confirmed, straddlers] = tree_search(tree.data(), tree.root(), query, U::from(radius));
num_hits = confirmed
.iter()
.chain(straddlers.iter())
.map(|&(c, _)| c.cardinality)
.sum::<usize>();
}
let mut hits = Hits::new(k);
confirmed.into_iter().chain(straddlers.into_iter()).for_each(|(c, d)| {
let indices = c.indices(tree.data());
let distances = if c.is_singleton() {
vec![d; c.cardinality]
} else {
tree.data().query_to_many(query, indices)
};
hits.push_batch(indices.iter().copied().zip(distances.into_iter()));
});
hits.extract()
}
#[cfg(test)]
mod tests {
use distances::vectors::euclidean;
use symagen::random_data;
use crate::{cakes::knn::linear, knn::tests::sort_hits, Cakes, PartitionCriteria, VecDataset};
#[test]
fn repeated_rnn() {
let (cardinality, dimensionality) = (1_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 repeated_nn = sort_hits(super::search(tree, query, k));
assert_eq!(linear_nn, repeated_nn);
}
}
}