Skip to main content

competitive/data_structure/
kdtree.rs

1use std::ops::Range;
2
3pub struct Static2DTree<T, U, V>
4where
5    T: Ord,
6    U: Ord,
7{
8    data: Vec<(T, U, V)>,
9}
10impl<T, U, V> Static2DTree<T, U, V>
11where
12    T: Ord,
13    U: Ord,
14{
15    pub fn new<I>(data: I) -> Self
16    where
17        I: IntoIterator<Item = (T, U, V)>,
18    {
19        let mut data: Vec<_> = data.into_iter().collect();
20        let n = data.len();
21        Self::build(&mut data, 0, n, 0);
22        Self { data }
23    }
24    fn build(data: &mut [(T, U, V)], l: usize, r: usize, depth: usize) {
25        if r - l <= 1 {
26            return;
27        }
28        let m = l.midpoint(r);
29        if depth.is_multiple_of(2) {
30            data[l..r].select_nth_unstable_by(m - l, |p, q| p.0.cmp(&q.0));
31        } else {
32            data[l..r].select_nth_unstable_by(m - l, |p, q| p.1.cmp(&q.1));
33        }
34        Self::build(data, l, m, depth + 1);
35        Self::build(data, m + 1, r, depth + 1);
36    }
37    /// Returns the values in the half-open rectangle. Their order is unspecified.
38    pub fn range(&self, range1: Range<T>, range2: Range<U>) -> Vec<&V> {
39        let mut res = vec![];
40        self.range_inner(&range1, &range2, 0, self.data.len(), 0, &mut res);
41        res
42    }
43    fn range_inner<'a>(
44        &'a self,
45        range1: &Range<T>,
46        range2: &Range<U>,
47        l: usize,
48        r: usize,
49        depth: usize,
50        res: &mut Vec<&'a V>,
51    ) {
52        if l < r {
53            let m = l.midpoint(r);
54            let (t, u, v) = &self.data[m];
55            if range1.contains(t) && range2.contains(u) {
56                res.push(v);
57            }
58            if if depth.is_multiple_of(2) {
59                &range1.start <= t
60            } else {
61                &range2.start <= u
62            } {
63                self.range_inner(range1, range2, l, m, depth + 1, res);
64            }
65            if if depth.is_multiple_of(2) {
66                t < &range1.end
67            } else {
68                u < &range2.end
69            } {
70                self.range_inner(range1, range2, m + 1, r, depth + 1, res);
71            }
72        }
73    }
74}
75
76#[cfg(test)]
77mod tests {
78    use super::*;
79    use crate::tools::Xorshift;
80
81    #[test]
82    fn test_static_2d_tree() {
83        let mut rng = Xorshift::default();
84        for _ in 0..64 {
85            let len = rng.rand(300) as usize;
86            let radius = rng.rand(64) as i32 + 1;
87            let values: Vec<_> = (0..len)
88                .map(|index| {
89                    (
90                        rng.rand((radius * 2 + 1) as u64) as i32 - radius,
91                        rng.rand((radius * 2 + 1) as u64) as i32 - radius,
92                        index,
93                    )
94                })
95                .collect();
96            let tree = Static2DTree::new(values.iter().copied());
97            for _ in 0..200 {
98                let mut x = [
99                    rng.rand((radius * 2 + 3) as u64) as i32 - radius - 1,
100                    rng.rand((radius * 2 + 3) as u64) as i32 - radius - 1,
101                ];
102                let mut y = [
103                    rng.rand((radius * 2 + 3) as u64) as i32 - radius - 1,
104                    rng.rand((radius * 2 + 3) as u64) as i32 - radius - 1,
105                ];
106                x.sort_unstable();
107                y.sort_unstable();
108
109                let mut actual: Vec<_> = tree
110                    .range(x[0]..x[1], y[0]..y[1])
111                    .into_iter()
112                    .copied()
113                    .collect();
114                let mut expected: Vec<_> = values
115                    .iter()
116                    .filter(|(px, py, _)| x[0] <= *px && *px < x[1] && y[0] <= *py && *py < y[1])
117                    .map(|(_, _, value)| *value)
118                    .collect();
119                actual.sort_unstable();
120                expected.sort_unstable();
121                assert_eq!(actual, expected);
122            }
123        }
124    }
125}