use crate::graph::{SearchOutputBuffer, search_output_buffer};
mod queue;
pub use queue::{NeighborPriorityQueue, NeighborPriorityQueueIdType, NeighborQueue};
#[cfg(feature = "experimental_diversity_search")]
mod diverse_priority_queue;
#[cfg(feature = "experimental_diversity_search")]
pub use diverse_priority_queue::{
Attribute, AttributeValueProvider, DiverseNeighborQueue, VectorIdWithAttribute,
};
#[derive(Debug, Default, Clone, Copy)]
pub struct Neighbor<I, D = f32> {
id: I,
distance: D,
}
impl<I, D> Neighbor<I, D> {
#[inline]
pub fn new(id: I, distance: D) -> Self {
Self { id, distance }
}
#[inline]
pub fn as_tuple(self) -> (I, D) {
(self.id, self.distance)
}
#[inline]
pub fn distance(&self) -> &D {
&self.distance
}
#[inline]
pub fn id(&self) -> &I {
&self.id
}
#[cfg(test)]
pub(crate) fn from_tuple((id, distance): (I, D)) -> Self {
Self::new(id, distance)
}
}
#[cfg(test)]
impl<I, D> crate::test::cmp::VerboseEq for Neighbor<I, D>
where
I: crate::test::cmp::VerboseEq,
D: crate::test::cmp::VerboseEq,
{
#[inline(never)]
#[track_caller]
fn verbose_eq(&self, other: &Self) -> crate::ANNResult<()> {
if let Err(err) = (self.id).verbose_eq(&other.id) {
return Err(err.context(crate::test::cmp::Field("id")));
}
if let Err(err) = (self.distance).verbose_eq(&other.distance) {
return Err(err.context(crate::test::cmp::Field("distance")));
}
Ok(())
}
}
pub mod ord {
use super::Neighbor;
pub fn fast_distance<I>(x: &Neighbor<I>, y: &Neighbor<I>) -> std::cmp::Ordering {
x.distance()
.partial_cmp(y.distance())
.unwrap_or(std::cmp::Ordering::Equal)
}
pub fn fast_distance_total<I>(x: &Neighbor<I>, y: &Neighbor<I>) -> std::cmp::Ordering
where
I: Ord,
{
fast_distance(x, y).then(x.id().cmp(y.id()))
}
pub fn reverse<F, I>(f: F) -> impl Fn(&Neighbor<I>, &Neighbor<I>) -> std::cmp::Ordering
where
F: Fn(&Neighbor<I>, &Neighbor<I>) -> std::cmp::Ordering,
{
move |x: &Neighbor<I>, y: &Neighbor<I>| f(x, y).reverse()
}
}
#[derive(Debug)]
pub struct BackInserter<'a, I, D = f32> {
buffer: &'a mut [Neighbor<I, D>],
position: usize,
}
impl<'a, I, D> BackInserter<'a, I, D> {
pub fn new(buffer: &'a mut [Neighbor<I, D>]) -> Self {
Self {
buffer,
position: 0,
}
}
pub fn capacity(&self) -> usize {
self.buffer.len()
}
}
impl<I, D> SearchOutputBuffer<I, D> for BackInserter<'_, I, D> {
fn size_hint(&self) -> Option<usize> {
Some(self.buffer.len() - self.position)
}
fn push(&mut self, neighbor: Neighbor<I, D>) -> search_output_buffer::BufferState {
if self.position == self.buffer.len() {
return search_output_buffer::BufferState::Full;
}
self.buffer[self.position] = neighbor;
self.position += 1;
if self.position == self.buffer.len() {
search_output_buffer::BufferState::Full
} else {
search_output_buffer::BufferState::Available
}
}
fn current_len(&self) -> usize {
self.position
}
fn extend<Itr>(&mut self, itr: Itr) -> usize
where
Itr: IntoIterator<Item = Neighbor<I, D>>,
{
let mut i = 0;
std::iter::zip(self.buffer.iter_mut().skip(self.position), itr).for_each(|(dst, src)| {
i += 1;
*dst = src;
});
self.position += i;
i
}
}
impl<I, D> SearchOutputBuffer<I, D> for Vec<Neighbor<I, D>> {
fn size_hint(&self) -> Option<usize> {
None
}
fn push(&mut self, neighbor: Neighbor<I, D>) -> search_output_buffer::BufferState {
self.push(neighbor);
search_output_buffer::BufferState::Available
}
fn current_len(&self) -> usize {
self.len()
}
fn extend<Itr>(&mut self, itr: Itr) -> usize
where
Itr: IntoIterator<Item = Neighbor<I, D>>,
{
let before = self.len();
Extend::extend(self, itr);
self.len() - before
}
}
#[cfg(test)]
mod neighbor_test {
use super::*;
use crate::test::cmp::assert_eq_verbose;
#[test]
fn fast_distance() {
let n1 = Neighbor::new(1, 1.0);
let n2 = Neighbor::new(2, 2.0);
assert!(ord::fast_distance(&n1, &n2).is_lt());
assert!(ord::fast_distance(&n2, &n1).is_gt());
assert!(ord::reverse(ord::fast_distance)(&n1, &n2).is_gt());
assert!(ord::reverse(ord::fast_distance)(&n2, &n1).is_lt());
assert!(ord::fast_distance(&n1, &n1).is_eq());
assert!(ord::fast_distance(&n2, &n2).is_eq());
assert!(ord::reverse(ord::fast_distance)(&n1, &n1).is_eq());
assert!(ord::reverse(ord::fast_distance)(&n2, &n2).is_eq());
let nan = Neighbor::new(3, f32::NAN);
assert!(ord::fast_distance(&n1, &nan).is_eq());
assert!(ord::fast_distance(&nan, &n1).is_eq());
assert!(ord::fast_distance(&nan, &nan).is_eq());
assert!(ord::reverse(ord::fast_distance)(&n1, &nan).is_eq());
assert!(ord::reverse(ord::fast_distance)(&nan, &n1).is_eq());
assert!(ord::reverse(ord::fast_distance)(&nan, &nan).is_eq());
}
#[test]
fn fast_distance_total() {
let n1 = Neighbor::new(1, 1.0);
let n2 = Neighbor::new(2, 2.0);
let n3 = Neighbor::new(3, 2.0);
let n4 = Neighbor::new(4, 3.0);
assert!(ord::fast_distance_total(&n1, &n1).is_eq());
assert!(ord::fast_distance_total(&n1, &n2).is_lt());
assert!(ord::fast_distance_total(&n1, &n3).is_lt());
assert!(ord::fast_distance_total(&n1, &n4).is_lt());
assert!(ord::fast_distance_total(&n2, &n1).is_gt());
assert!(ord::fast_distance_total(&n2, &n2).is_eq());
assert!(ord::fast_distance_total(&n2, &n3).is_lt());
assert!(ord::fast_distance_total(&n2, &n4).is_lt());
assert!(ord::fast_distance_total(&n3, &n1).is_gt());
assert!(ord::fast_distance_total(&n3, &n2).is_gt());
assert!(ord::fast_distance_total(&n3, &n3).is_eq());
assert!(ord::fast_distance_total(&n3, &n4).is_lt());
assert!(ord::fast_distance_total(&n4, &n1).is_gt());
assert!(ord::fast_distance_total(&n4, &n2).is_gt());
assert!(ord::fast_distance_total(&n4, &n3).is_gt());
assert!(ord::fast_distance_total(&n4, &n4).is_eq());
}
#[test]
fn test_search_output_buffer() {
const MAX_LENGTH: usize = 5;
fn f(i: usize) -> Neighbor<u32> {
Neighbor::new(i as u32, i as f32)
}
{
let mut buffer = [Neighbor::<u32>::default(); MAX_LENGTH];
let mut inserter = BackInserter::new(&mut buffer);
assert_eq!(inserter.capacity(), MAX_LENGTH);
assert_eq!(inserter.size_hint(), Some(MAX_LENGTH));
assert_eq!(inserter.current_len(), 0);
assert!(inserter.push(Neighbor::new(1, 1.0)).is_available());
assert_eq!(inserter.current_len(), 1);
assert_eq!(inserter.size_hint(), Some(MAX_LENGTH - 1));
assert!(inserter.push(Neighbor::new(2, 2.0)).is_available());
assert_eq!(inserter.current_len(), 2);
assert_eq!(inserter.size_hint(), Some(MAX_LENGTH - 2));
assert!(inserter.push(Neighbor::new(3, 3.0)).is_available());
assert_eq!(inserter.current_len(), 3);
assert_eq!(inserter.size_hint(), Some(MAX_LENGTH - 3));
assert!(inserter.push(Neighbor::new(4, 4.0)).is_available());
assert_eq!(inserter.current_len(), 4);
assert_eq!(inserter.size_hint(), Some(MAX_LENGTH - 4));
assert!(inserter.push(Neighbor::new(5, 5.0)).is_full());
assert_eq!(inserter.current_len(), 5);
assert_eq!(inserter.size_hint(), Some(0));
assert!(inserter.push(Neighbor::new(6, 6.0)).is_full());
assert_eq!(inserter.current_len(), 5);
assert_eq!(inserter.size_hint(), Some(0));
assert_eq_verbose!(buffer, [f(1), f(2), f(3), f(4), f(5)]);
}
{
let mut buffer = [Neighbor::<u32>::default(); MAX_LENGTH];
let mut inserter = BackInserter::new(&mut buffer);
assert_eq!(inserter.capacity(), MAX_LENGTH);
assert_eq!(inserter.size_hint(), Some(MAX_LENGTH));
assert_eq!(inserter.current_len(), 0);
let set = inserter.extend(
[(1, 1.0), (2, 2.0), (3, 3.0), (4, 4.0), (5, 5.0), (6, 6.0)]
.map(Neighbor::from_tuple),
);
assert_eq!(set, MAX_LENGTH);
assert_eq!(inserter.current_len(), MAX_LENGTH);
assert_eq!(inserter.size_hint(), Some(0));
assert!(inserter.push(Neighbor::new(7, 7.0)).is_full());
let set = inserter.extend([(10, 10.0), (20, 20.0)].map(Neighbor::from_tuple));
assert_eq!(set, 0, "no more items can be added");
assert_eq_verbose!(buffer, [f(1), f(2), f(3), f(4), f(5)]);
}
{
let mut buffer = [Neighbor::<u32>::default(); MAX_LENGTH];
let mut inserter = BackInserter::new(&mut buffer);
assert!(inserter.push(Neighbor::new(1, 1.0)).is_available());
let set = inserter.extend([(2, 2.0), (3, 3.0)].map(Neighbor::from_tuple));
assert_eq!(set, 2, "only two items were pushed");
assert_eq!(inserter.current_len(), 3);
assert_eq!(inserter.size_hint(), Some(2));
assert!(inserter.push(Neighbor::new(4, 4.0)).is_available());
assert_eq!(inserter.current_len(), 4);
assert_eq!(inserter.size_hint(), Some(1));
let set = inserter.extend([(5, 5.0), (6, 6.0)].map(Neighbor::from_tuple));
assert_eq!(
set, 1,
"there should only be room for one more item in the buffer"
);
assert_eq!(inserter.current_len(), 5);
assert_eq!(inserter.size_hint(), Some(0));
assert_eq_verbose!(buffer, [f(1), f(2), f(3), f(4), f(5)]);
}
}
#[test]
fn test_vec_neighbor_search_output_buffer() {
use crate::graph::search_output_buffer::SearchOutputBuffer;
let mut buf: Vec<Neighbor<u32>> = Vec::new();
assert_eq!(SearchOutputBuffer::<u32>::size_hint(&buf), None);
assert_eq!(SearchOutputBuffer::<u32>::current_len(&buf), 0);
assert!(SearchOutputBuffer::push(&mut buf, Neighbor::new(1, 0.5)).is_available());
assert!(SearchOutputBuffer::push(&mut buf, Neighbor::new(2, 1.0)).is_available());
assert_eq!(SearchOutputBuffer::<u32>::current_len(&buf), 2);
assert_eq_verbose!(buf[0], Neighbor::new(1, 0.5));
assert_eq_verbose!(buf[1], Neighbor::new(2, 1.0));
let count = SearchOutputBuffer::extend(
&mut buf,
[(3u32, 1.5), (4, 2.0), (5, 2.5)].map(Neighbor::from_tuple),
);
assert_eq!(count, 3);
assert_eq!(SearchOutputBuffer::<u32>::current_len(&buf), 5);
assert_eq_verbose!(buf[4], Neighbor::new(5, 2.5));
}
}