Skip to main content

rstar/algorithm/
nearest_neighbor.rs

1use crate::point::min_inline;
2use crate::{
3    node::{ParentNode, RTreeNode},
4    object::Distance,
5};
6use crate::{Envelope, PointDistance, RTreeObject};
7
8#[cfg(doc)]
9use crate::RTree;
10
11use alloc::collections::BinaryHeap;
12#[cfg(not(test))]
13use alloc::{vec, vec::Vec};
14use core::mem::replace;
15use heapless::binary_heap as static_heap;
16use num_traits::Bounded;
17
18struct RTreeNodeDistanceWrapper<'a, T>
19where
20    T: PointDistance + 'a,
21{
22    node: &'a RTreeNode<T>,
23    distance: Distance<T>,
24}
25
26impl<T> PartialEq for RTreeNodeDistanceWrapper<'_, T>
27where
28    T: PointDistance,
29{
30    fn eq(&self, other: &Self) -> bool {
31        self.distance == other.distance
32    }
33}
34
35impl<T> PartialOrd for RTreeNodeDistanceWrapper<'_, T>
36where
37    T: PointDistance,
38{
39    fn partial_cmp(&self, other: &Self) -> Option<::core::cmp::Ordering> {
40        Some(self.cmp(other))
41    }
42}
43
44impl<T> Eq for RTreeNodeDistanceWrapper<'_, T> where T: PointDistance {}
45
46impl<T> Ord for RTreeNodeDistanceWrapper<'_, T>
47where
48    T: PointDistance,
49{
50    fn cmp(&self, other: &Self) -> ::core::cmp::Ordering {
51        // Inverse comparison creates a min heap
52        other.distance.partial_cmp(&self.distance).unwrap()
53    }
54}
55
56impl<'a, T> NearestNeighborDistance2Iterator<'a, T>
57where
58    T: PointDistance,
59{
60    pub(crate) fn new(
61        root: &'a ParentNode<T>,
62        query_point: <T::Envelope as Envelope>::Point,
63    ) -> Self {
64        let mut result = NearestNeighborDistance2Iterator {
65            nodes: SmallHeap::new(),
66            query_point,
67        };
68        result.extend_heap(&root.children);
69        result
70    }
71
72    fn extend_heap(&mut self, children: &'a [RTreeNode<T>]) {
73        let &mut NearestNeighborDistance2Iterator {
74            ref mut nodes,
75            ref query_point,
76        } = self;
77        nodes.extend(children.iter().map(|child: &RTreeNode<T>| {
78            let distance = match child {
79                RTreeNode::Parent(ref data) => data.envelope.distance_2(query_point),
80                RTreeNode::Leaf(ref t) => t.distance_2(query_point),
81            };
82
83            RTreeNodeDistanceWrapper {
84                node: child,
85                distance,
86            }
87        }));
88    }
89}
90
91impl<'a, T> Iterator for NearestNeighborDistance2Iterator<'a, T>
92where
93    T: PointDistance,
94{
95    type Item = (&'a T, Distance<T>);
96
97    fn next(&mut self) -> Option<Self::Item> {
98        while let Some(current) = self.nodes.pop() {
99            match current {
100                RTreeNodeDistanceWrapper {
101                    node: RTreeNode::Parent(ref data),
102                    ..
103                } => {
104                    self.extend_heap(&data.children);
105                }
106                RTreeNodeDistanceWrapper {
107                    node: RTreeNode::Leaf(ref t),
108                    distance,
109                } => {
110                    return Some((t, distance));
111                }
112            }
113        }
114        None
115    }
116}
117
118/// Iterator returned by [`RTree::nearest_neighbor_iter_with_distance_2`].
119pub struct NearestNeighborDistance2Iterator<'a, T>
120where
121    T: PointDistance + 'a,
122{
123    nodes: SmallHeap<RTreeNodeDistanceWrapper<'a, T>>,
124    query_point: <T::Envelope as Envelope>::Point,
125}
126
127impl<'a, T> NearestNeighborIterator<'a, T>
128where
129    T: PointDistance,
130{
131    pub(crate) fn new(
132        root: &'a ParentNode<T>,
133        query_point: <T::Envelope as Envelope>::Point,
134    ) -> Self {
135        NearestNeighborIterator {
136            iter: NearestNeighborDistance2Iterator::new(root, query_point),
137        }
138    }
139}
140
141impl<'a, T> Iterator for NearestNeighborIterator<'a, T>
142where
143    T: PointDistance,
144{
145    type Item = &'a T;
146
147    fn next(&mut self) -> Option<Self::Item> {
148        self.iter.next().map(|(t, _distance)| t)
149    }
150}
151
152/// Iterator returned by [`RTree::nearest_neighbor_iter`].
153pub struct NearestNeighborIterator<'a, T>
154where
155    T: PointDistance + 'a,
156{
157    iter: NearestNeighborDistance2Iterator<'a, T>,
158}
159
160enum SmallHeap<T: Ord> {
161    Stack(static_heap::BinaryHeap<T, static_heap::Max, 32>),
162    Heap(BinaryHeap<T>),
163}
164
165impl<T: Ord> SmallHeap<T> {
166    pub fn new() -> Self {
167        Self::Stack(static_heap::BinaryHeap::new())
168    }
169
170    pub fn pop(&mut self) -> Option<T> {
171        match self {
172            SmallHeap::Stack(heap) => heap.pop(),
173            SmallHeap::Heap(heap) => heap.pop(),
174        }
175    }
176
177    pub fn push(&mut self, item: T) {
178        match self {
179            SmallHeap::Stack(heap) => {
180                if let Err(item) = heap.push(item) {
181                    let capacity = heap.len() + 1;
182                    let new_heap = self.spill(capacity);
183                    new_heap.push(item);
184                }
185            }
186            SmallHeap::Heap(heap) => heap.push(item),
187        }
188    }
189
190    pub fn extend<I>(&mut self, iter: I)
191    where
192        I: ExactSizeIterator<Item = T>,
193    {
194        match self {
195            SmallHeap::Stack(heap) => {
196                if heap.capacity() >= heap.len() + iter.len() {
197                    for item in iter {
198                        if heap.push(item).is_err() {
199                            unreachable!();
200                        }
201                    }
202                } else {
203                    let capacity = heap.len() + iter.len();
204                    let new_heap = self.spill(capacity);
205                    new_heap.extend(iter);
206                }
207            }
208            SmallHeap::Heap(heap) => heap.extend(iter),
209        }
210    }
211
212    #[cold]
213    fn spill(&mut self, capacity: usize) -> &mut BinaryHeap<T> {
214        let new_heap = BinaryHeap::with_capacity(capacity);
215        let old_heap = replace(self, SmallHeap::Heap(new_heap));
216
217        let new_heap = match self {
218            SmallHeap::Heap(new_heap) => new_heap,
219            SmallHeap::Stack(_) => unreachable!(),
220        };
221        let old_heap = match old_heap {
222            SmallHeap::Stack(old_heap) => old_heap,
223            SmallHeap::Heap(_) => unreachable!(),
224        };
225
226        new_heap.extend(old_heap.into_vec());
227
228        new_heap
229    }
230}
231
232pub fn nearest_neighbor_with_distance_2<T>(
233    node: &ParentNode<T>,
234    query_point: <T::Envelope as Envelope>::Point,
235) -> Option<(&T, Distance<T>)>
236where
237    T: PointDistance,
238{
239    fn extend_heap<'a, T>(
240        nodes: &mut SmallHeap<RTreeNodeDistanceWrapper<'a, T>>,
241        node: &'a ParentNode<T>,
242        query_point: &<T::Envelope as Envelope>::Point,
243        min_max_distance: &mut Distance<T>,
244    ) where
245        T: PointDistance + 'a,
246    {
247        for child in &node.children {
248            let distance_if_less_or_equal = match child {
249                RTreeNode::Parent(ref data) => {
250                    let distance = data.envelope.distance_2(query_point);
251                    if distance <= *min_max_distance {
252                        Some(distance)
253                    } else {
254                        None
255                    }
256                }
257                RTreeNode::Leaf(ref t) => {
258                    t.distance_2_if_less_or_equal(query_point, *min_max_distance)
259                }
260            };
261            if let Some(distance) = distance_if_less_or_equal {
262                *min_max_distance = min_inline(
263                    *min_max_distance,
264                    child.envelope().min_max_dist_2(query_point),
265                );
266                nodes.push(RTreeNodeDistanceWrapper {
267                    node: child,
268                    distance,
269                });
270            }
271        }
272    }
273
274    // Calculate smallest minmax-distance
275    let mut smallest_min_max: Distance<T> = Bounded::max_value();
276    let mut nodes = SmallHeap::new();
277    extend_heap(&mut nodes, node, &query_point, &mut smallest_min_max);
278    while let Some(current) = nodes.pop() {
279        match current {
280            RTreeNodeDistanceWrapper {
281                node: RTreeNode::Parent(ref data),
282                ..
283            } => {
284                extend_heap(&mut nodes, data, &query_point, &mut smallest_min_max);
285            }
286            RTreeNodeDistanceWrapper {
287                node: RTreeNode::Leaf(ref t),
288                distance,
289            } => {
290                return Some((t, distance));
291            }
292        }
293    }
294    None
295}
296
297pub fn nearest_neighbors_with_distance_2<T>(
298    node: &ParentNode<T>,
299    query_point: <T::Envelope as Envelope>::Point,
300) -> Option<(Vec<&T>, Distance<T>)>
301where
302    T: PointDistance,
303{
304    let mut nearest_neighbors = NearestNeighborDistance2Iterator::new(node, query_point);
305
306    let (first, first_distance_2) = nearest_neighbors.next()?;
307
308    // The result will at least contain the first nearest neighbor.
309    let mut result = vec![first];
310
311    // Use the distance to the first nearest neighbor
312    // to filter out the rest of the nearest neighbors
313    // that are farther than this first neighbor.
314    result.extend(
315        nearest_neighbors
316            .take_while(|(_, next_distance_2)| next_distance_2 == &first_distance_2)
317            .map(|(next, _)| next),
318    );
319
320    Some((result, first_distance_2))
321}
322
323#[cfg(test)]
324mod test {
325    use crate::object::PointDistance;
326    use crate::rtree::RTree;
327    use crate::test_utilities::*;
328
329    #[test]
330    fn test_nearest_neighbor_empty() {
331        let tree: RTree<[f32; 2]> = RTree::new();
332        assert!(tree.nearest_neighbor([0.0, 213.0]).is_none());
333    }
334
335    #[test]
336    fn test_nearest_neighbor() {
337        let points = create_random_points(1000, SEED_1);
338        let tree = RTree::bulk_load(points.clone());
339
340        let sample_points = create_random_points(100, SEED_2);
341        for sample_point in sample_points {
342            let mut nearest = None;
343            let mut closest_dist = f64::INFINITY;
344            for point in &points {
345                let delta = [point[0] - sample_point[0], point[1] - sample_point[1]];
346                let new_dist = delta[0] * delta[0] + delta[1] * delta[1];
347                if new_dist < closest_dist {
348                    closest_dist = new_dist;
349                    nearest = Some(point);
350                }
351            }
352            assert_eq!(nearest, tree.nearest_neighbor(sample_point));
353        }
354    }
355
356    #[test]
357    fn test_nearest_neighbors_empty() {
358        let tree: RTree<[f32; 2]> = RTree::new();
359        assert!(tree.nearest_neighbors(&[0.0, 213.0]).is_empty());
360    }
361
362    #[test]
363    fn test_nearest_neighbors() {
364        let points = create_random_points(1000, SEED_1);
365        let tree = RTree::bulk_load(points);
366
367        let sample_points = create_random_points(50, SEED_2);
368        for sample_point in &sample_points {
369            let nearest_neighbors = tree.nearest_neighbors(sample_point);
370            let mut distance = -1.0;
371            for nn in &nearest_neighbors {
372                if distance < 0.0 {
373                    distance = sample_point.distance_2(nn);
374                } else {
375                    let new_distance = sample_point.distance_2(nn);
376                    assert_eq!(new_distance, distance);
377                }
378            }
379        }
380    }
381
382    #[test]
383    fn test_nearest_neighbor_iterator() {
384        let mut points = create_random_points(1000, SEED_1);
385        let tree = RTree::bulk_load(points.clone());
386
387        let sample_points = create_random_points(50, SEED_2);
388        for sample_point in sample_points {
389            points.sort_by(|r, l| {
390                r.distance_2(&sample_point)
391                    .partial_cmp(&l.distance_2(&sample_point))
392                    .unwrap()
393            });
394            let collected: Vec<_> = tree.nearest_neighbor_iter(sample_point).cloned().collect();
395            assert_eq!(points, collected);
396        }
397    }
398
399    #[test]
400    fn test_nearest_neighbor_iterator_with_distance_2() {
401        let points = create_random_points(1000, SEED_2);
402        let tree = RTree::bulk_load(points);
403
404        let sample_points = create_random_points(50, SEED_1);
405        for sample_point in sample_points {
406            let mut last_distance = 0.0;
407            for (point, distance) in tree.nearest_neighbor_iter_with_distance_2(sample_point) {
408                assert_eq!(point.distance_2(&sample_point), distance);
409                assert!(last_distance < distance);
410                last_distance = distance;
411            }
412        }
413    }
414}