competitive_library/graph/
lowest_common_ancestor_doubling.rs

1//! ダブリングを使用してLCA を求める
2use std::collections::VecDeque;
3
4struct Node {
5    pub parent: Option<usize>,
6    pub number: usize,
7    pub depth: i64,
8}
9
10impl Node {
11    #[inline]
12    pub fn new(parent: Option<usize>, number: usize, depth: i64) -> Self {
13        Node {
14            parent,
15            number,
16            depth,
17        }
18    }
19}
20
21pub struct LowestCommonAncestor {
22    max_log_v: usize,
23    root: usize,
24    depths: Vec<i64>,
25    ancestors: Vec<Vec<Option<usize>>>,
26}
27
28impl LowestCommonAncestor {
29    // 隣接リストで受け取る
30    #[inline]
31    pub fn new(edges: &[Vec<i64>], root: usize) -> Self {
32        let max_v = edges.len();
33        let max_log_v = ((max_v as f64).ln() / 2.0_f64.ln()) as usize + 1;
34        let mut ancestors = vec![vec![None; max_v]; max_log_v + 1];
35        let mut depths = vec![0; max_v];
36
37        let mut q = VecDeque::new();
38        q.push_back(Node::new(None, root, 0));
39        while let Some(node) = q.pop_front() {
40            ancestors[0][node.number] = node.parent;
41
42            depths[node.number] = node.depth;
43
44            edges[node.number]
45                .iter()
46                .filter(|&&v| node.parent.filter(|&x| x == v as usize).is_none())
47                .for_each(|&v| {
48                    q.push_back(Node::new(Some(node.number), v as usize, node.depth + 1))
49                });
50        }
51
52        (0..max_log_v).for_each(|i| {
53            (0..max_v).for_each(|j| {
54                if let Some(ancetor) = ancestors[i][j] {
55                    ancestors[i + 1][j] = ancestors[i][ancetor];
56                }
57            })
58        });
59
60        LowestCommonAncestor {
61            max_log_v,
62            root,
63            depths,
64            ancestors,
65        }
66    }
67    #[inline]
68    pub fn get_lca(&self, u: usize, v: usize) -> Option<usize> {
69        let (mut u, mut v) = if self.depths[u] > self.depths[v] {
70            (v, u)
71        } else {
72            (u, v)
73        };
74
75        for k in 0..self.max_log_v {
76            if (((self.depths[v] - self.depths[u]) >> k) & 1) == 1 {
77                v = self.ancestors[k][v].unwrap();
78            }
79        }
80
81        if u == v {
82            return Some(u);
83        }
84
85        for k in (0..self.max_log_v).rev() {
86            if self.ancestors[k][u].is_none()
87                || self.ancestors[k][v].is_none()
88                || self.ancestors[k][u] == self.ancestors[k][v]
89            {
90                continue;
91            }
92
93            u = self.ancestors[k][u].unwrap();
94            v = self.ancestors[k][v].unwrap();
95        }
96        self.ancestors[0][u]
97    }
98    #[inline]
99    pub fn get_distance(&self, u: usize, v: usize) -> i64 {
100        let lca = self.get_lca(u, v).unwrap_or(self.root);
101        self.depths[u] + self.depths[v] - self.depths[lca] * 2
102    }
103}
104
105#[cfg(test)]
106mod tests {
107
108    use super::*;
109    #[test]
110    fn test_lca() {
111        let n = 5;
112        let mut e = vec![vec![]; n];
113        for (i, &v) in [0, 0, 2, 2].iter().enumerate() {
114            e[v].push(i as i64 + 1);
115        }
116
117        let lca = LowestCommonAncestor::new(&e, 0);
118        for &(u, v, ans) in [
119            (0, 0, 0),
120            (0, 1, 0),
121            (0, 2, 0),
122            (0, 3, 0),
123            (0, 4, 0),
124            (1, 1, 1),
125            (1, 2, 0),
126            (1, 3, 0),
127            (1, 4, 0),
128            (2, 2, 2),
129            (2, 3, 2),
130            (2, 4, 2),
131            (3, 3, 3),
132            (3, 4, 2),
133            (4, 4, 4),
134        ]
135        .iter()
136        {
137            assert_eq!(lca.get_lca(u, v).unwrap_or(0), ans);
138        }
139    }
140}