use core::{cmp::Ordering, f64::EPSILON};
use distances::Number;
use crate::{utils, Dataset, RnnAlgorithm, Tree};
const MULTIPLIER: f64 = 2.0;
#[allow(clippy::module_name_repetitions)]
#[derive(Clone, Copy, Debug)]
pub enum KnnAlgorithm {
Linear,
RepeatedRnn,
}
impl Default for KnnAlgorithm {
fn default() -> Self {
Self::RepeatedRnn
}
}
impl KnnAlgorithm {
pub(crate) fn search<T, U, D>(self, query: T, k: usize, tree: &Tree<T, U, D>) -> Vec<(usize, U)>
where
T: Send + Sync + Copy,
U: Number,
D: Dataset<T, U>,
{
match self {
Self::Linear => Self::linear_search(tree.data(), query, k, tree.indices()),
Self::RepeatedRnn => Self::knn_by_rnn(tree, query, k),
}
}
pub(crate) fn linear_search<T, U, D>(data: &D, query: T, k: usize, indices: &[usize]) -> Vec<(usize, U)>
where
T: Send + Sync + Copy,
U: Number,
D: Dataset<T, U>,
{
let distances = data.query_to_many(query, indices);
let mut hits = indices.iter().copied().zip(distances.into_iter()).collect::<Vec<_>>();
hits.sort_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(Ordering::Less));
hits[..k].to_vec()
}
pub(crate) fn knn_by_rnn<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 mut radius = EPSILON + tree.radius().as_f64() / tree.cardinality().as_f64();
let [mut confirmed, mut straddlers] =
RnnAlgorithm::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] = RnnAlgorithm::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));
assert!(factor > 1.);
radius *= if factor < MULTIPLIER { factor } else { MULTIPLIER };
[confirmed, straddlers] = RnnAlgorithm::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 = confirmed
.into_iter()
.chain(straddlers.into_iter())
.flat_map(|(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)
};
indices.iter().copied().zip(distances.into_iter())
})
.collect::<Vec<_>>();
hits.sort_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(Ordering::Greater));
hits[..k].to_vec()
}
}