use distances::Number;
use crate::{cakes::rnn::clustered, utils, Cluster, Instance, Tree};
use super::Hits;
const MULTIPLIER: f64 = 2.0;
pub fn search<I, U, D>(tree: &Tree<I, U, D>, query: &I, k: usize) -> Vec<(usize, U)>
where
I: Instance,
U: Number,
D: crate::Dataset<I, U>,
{
let mut radius = f64::EPSILON + tree.radius().as_f64() / tree.cardinality().as_f64();
let [mut confirmed, mut straddlers] = clustered::tree_search(tree.data(), &tree.root, query, U::from(radius));
let mut num_confirmed = count_hits(&confirmed);
while num_confirmed == 0 {
radius *= MULTIPLIER;
[confirmed, straddlers] = clustered::tree_search(tree.data(), &tree.root, query, U::from(radius));
num_confirmed = count_hits(&confirmed);
}
while num_confirmed < k {
let lfd = utils::mean(
&confirmed
.iter()
.chain(straddlers.iter())
.map(|&(c, _)| c.lfd())
.collect::<Vec<_>>(),
);
let factor = (k.as_f64() / num_confirmed.as_f64()).powf(1. / (lfd + f64::EPSILON));
radius *= if factor < MULTIPLIER { factor } else { MULTIPLIER };
[confirmed, straddlers] = clustered::tree_search(tree.data(), &tree.root, query, U::from(radius));
num_confirmed = count_hits(&confirmed);
}
Hits::from_vec(
k,
clustered::leaf_search(&tree.data, confirmed, straddlers, query, U::from(radius)),
)
.extract()
}
fn count_hits<U: Number>(clusters: &[(&Cluster<U>, U)]) -> usize {
clusters.iter().map(|(c, _)| c.cardinality()).sum()
}