competitive_library/graph/
lowest_common_ancestor_doubling.rs1use 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 #[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}