use core::cmp::Ordering;
use distances::Number;
use rayon::prelude::*;
use crate::{knn, rnn, Dataset, Instance, PartitionCriteria, Tree};
#[derive(Debug)]
pub struct Cakes<I: Instance, U: Number, D: Dataset<I, U>> {
pub(crate) tree: Tree<I, U, D>,
pub(crate) best_knn: Option<knn::Algorithm>,
}
impl<I: Instance, U: Number, D: Dataset<I, U>> Cakes<I, U, D> {
pub fn new(data: D, seed: Option<u64>, criteria: &PartitionCriteria<U>) -> Self {
Self {
tree: Tree::new(data, seed).partition(criteria),
best_knn: None,
}
}
pub const fn data(&self) -> &D {
self.tree.data()
}
pub const fn tree(&self) -> &Tree<I, U, D> {
&self.tree
}
pub const fn depth(&self) -> usize {
self.tree.depth
}
pub const fn radius(&self) -> U {
self.tree.radius()
}
#[must_use]
pub fn auto_tune(mut self, k: usize, tuning_depth: usize) -> Self {
let queries = self
.tree
.root
.subtree()
.into_iter()
.filter(|&c| c.depth() == tuning_depth || c.is_leaf() && c.depth() < tuning_depth)
.map(|c| &self.tree.data[c.arg_center()])
.collect::<Vec<_>>();
(self.best_knn, _, _) = knn::Algorithm::variants()
.iter()
.map(|&algorithm| {
let start = std::time::Instant::now();
let hits = self.batch_knn_search(&queries, k, algorithm);
let elapsed = start.elapsed().as_secs_f32();
(Some(algorithm), hits, elapsed)
})
.min_by(|(_, _, a), (_, _, b)| a.partial_cmp(b).unwrap_or(Ordering::Greater))
.unwrap_or_else(|| unreachable!("There are several variants of knn-search"));
self
}
pub fn batch_rnn_search(&self, queries: &[&I], radius: U, algorithm: rnn::Algorithm) -> Vec<Vec<(usize, U)>> {
queries
.par_iter()
.map(|query| self.rnn_search(query, radius, algorithm))
.collect()
}
pub fn rnn_search(&self, query: &I, radius: U, algorithm: rnn::Algorithm) -> Vec<(usize, U)> {
algorithm.search(query, radius, &self.tree)
}
pub fn batch_knn_search(&self, queries: &[&I], k: usize, algorithm: knn::Algorithm) -> Vec<Vec<(usize, U)>> {
queries
.par_iter()
.map(|query| self.knn_search(query, k, algorithm))
.collect()
}
pub fn knn_search(&self, query: &I, k: usize, algorithm: knn::Algorithm) -> Vec<(usize, U)> {
algorithm.search(&self.tree, query, k)
}
pub fn batch_tuned_knn(&self, queries: &[&I], k: usize) -> Vec<Vec<(usize, U)>> {
queries.par_iter().map(|query| self.tuned_knn(query, k)).collect()
}
pub fn tuned_knn(&self, query: &I, k: usize) -> Vec<(usize, U)> {
self.knn_search(query, k, self.best_knn.unwrap_or_default())
}
pub fn batch_linear_knn(&self, queries: &[&I], k: usize) -> Vec<Vec<(usize, U)>> {
queries.par_iter().map(|query| self.linear_knn(query, k)).collect()
}
pub fn linear_knn(&self, query: &I, k: usize) -> Vec<(usize, U)> {
self.knn_search(query, k, knn::Algorithm::Linear)
}
}