use std::{cmp::Ordering, collections::BinaryHeap};
use rand::Rng;
use num_traits::Float;
use crossbeam::utils::CachePadded;
#[derive(Clone, Debug)]
pub(crate) struct Node<T: Float> {
index: usize,
threshold: T,
left: Option<Box<Node<T>>>,
right: Option<Box<Node<T>>>,
}
impl<T: Float> Default for Node<T> {
fn default() -> Self {
Node {
index: 0,
threshold: T::zero(),
left: None,
right: None,
}
}
}
struct HeapItem<T: Float> {
index: usize,
distance: T,
}
impl<T: Float> PartialOrd for HeapItem<T> {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
impl<T: Float> PartialEq for HeapItem<T> {
fn eq(&self, other: &Self) -> bool {
self.distance == other.distance
}
}
impl<T: Float> Eq for HeapItem<T> {}
impl<T: Float> Ord for HeapItem<T> {
fn cmp(&self, other: &Self) -> Ordering {
self.distance.partial_cmp(&other.distance).unwrap()
}
}
struct VPTreeBuilder<'a, T: Float> {
root: &'a mut Option<Box<Node<T>>>,
lower: usize,
upper: usize,
}
impl<'a, T: Float> VPTreeBuilder<'a, T> {
fn new(root: &'a mut Option<Box<Node<T>>>, lower: usize, upper: usize) -> Self {
Self { root, lower, upper }
}
}
pub(crate) struct VPTree<'a, T: Float + Send + Sync, U> {
items: Vec<(usize, &'a U)>,
pub(crate) root: Option<Box<Node<T>>>,
}
impl<'a, T: Float + Send + Sync, U> VPTree<'a, T, U> {
pub fn new<F>(items: &'a [U], metric_f: F) -> Self
where
F: Fn(&U, &U) -> T,
{
let mut tree = VPTree {
items: items.iter().enumerate().collect(),
root: None,
};
let n_samples = tree.items.len(); tree.build_from_points(0, n_samples, &metric_f);
tree
}
fn build_from_points<F>(&mut self, lower: usize, upper: usize, metric_f: F)
where
F: Fn(&U, &U) -> T,
{
let mut stack = vec![VPTreeBuilder::new(&mut self.root, lower, upper)];
let mut thread_rng = super::make_rng();
while let Some(builder) = stack.pop() {
let VPTreeBuilder { root, lower, upper } = builder;
if upper != lower {
*root = Some(Box::new(Node::default()));
let node = root.as_deref_mut().unwrap();
node.index = lower;
if upper - lower > 1 {
let i = thread_rng.random_range(lower..upper);
self.items.swap(lower, i);
let (_, to_cmp) = self.items[lower];
let median: usize = (upper + lower) / 2;
self.items[lower + 1..upper].select_nth_unstable_by(
median - lower - 1,
&mut |(_, a): &(usize, &U), (_, b): &(usize, &U)| {
if metric_f(to_cmp, a) < metric_f(to_cmp, b) {
Ordering::Less
} else if metric_f(to_cmp, a) == metric_f(to_cmp, b) {
Ordering::Equal
} else {
Ordering::Greater
}
},
);
node.threshold = metric_f(self.items[lower].1, self.items[median].1);
stack.push(VPTreeBuilder::new(&mut node.left, lower + 1, median));
stack.push(VPTreeBuilder::new(&mut node.right, median, upper));
}
}
}
}
fn look_up<F>(
&self,
tau: &mut T, target: &U,
k: usize,
heap: &mut BinaryHeap<HeapItem<T>>,
metric_f: F,
) where
F: Fn(&U, &U) -> T,
{
let mut stack: Vec<&Option<Box<Node<T>>>> = vec![&self.root];
while let Some(next_in_stack) = stack.pop() {
if let Some(node) = next_in_stack {
let (original_position, point) = &self.items[node.index];
let distance: T = metric_f(point, target);
if distance < *tau {
if heap.len() == k {
heap.pop();
}
heap.push(HeapItem {
index: *original_position,
distance,
});
if heap.len() == k {
*tau = heap.peek().unwrap().distance;
}
}
match (node.left.as_ref(), node.right.as_ref()) {
(None, None) => continue,
(_, _) => {
if distance < node.threshold {
if distance - *tau <= node.threshold {
stack.push(&node.left)
}
if distance + *tau >= node.threshold {
stack.push(&node.right)
}
} else {
if distance + *tau >= node.threshold {
stack.push(&node.right)
}
if distance - *tau <= node.threshold {
stack.push(&node.left)
}
}
}
}
}
}
}
pub fn search<F>(
&self,
target: &U,
target_index: usize,
k: usize,
neighbors_indices: &mut [CachePadded<usize>],
distances: &mut [CachePadded<T>],
metric_f: F,
) where
F: Fn(&U, &U) -> T,
{
debug_assert_eq!(neighbors_indices.len(), distances.len());
let mut heap: BinaryHeap<HeapItem<T>> = BinaryHeap::with_capacity(k);
self.look_up(&mut T::max_value(), target, k, &mut heap, metric_f);
let results = heap.into_sorted_vec();
neighbors_indices
.iter_mut()
.zip(distances.iter_mut())
.zip(results.iter().filter(|result| {
let HeapItem { index, distance: _ } = result;
target_index != *index
}))
.for_each(|((idx, d), result)| {
let HeapItem { index, distance } = result;
**idx = *index;
**d = *distance;
});
}
}