1use crate::algebra::{AbelianGroup, Monoid};
4use crate::graph::{Graph, Neighbor, UndirectedSparseGraph};
5
6#[codesnip::entry("ReRooting", include("algebra", "tree_order"))]
7#[derive(Clone, Debug)]
11pub struct ReRooting<'a, M: Monoid, F: Fn(&M::T, usize, Option<usize>) -> M::T> {
12 graph: &'a UndirectedSparseGraph,
13 pub dp: Vec<M::T>,
15 pub ep: Vec<M::T>,
17 rooting: F,
19}
20#[codesnip::entry("ReRooting")]
21impl<'a, M, F> ReRooting<'a, M, F>
22where
23 M: Monoid,
24 F: Fn(&M::T, usize, Option<usize>) -> M::T,
25{
26 pub fn new(graph: &'a UndirectedSparseGraph, rooting: F) -> Self {
27 Self::build(graph, rooting, None::<fn(&M::T, &M::T) -> M::T>)
28 }
29
30 pub fn new_with_inverse(graph: &'a UndirectedSparseGraph, rooting: F) -> Self
31 where
32 M: AbelianGroup,
33 {
34 Self::build(graph, rooting, Some(M::rinv_operate))
35 }
36
37 fn build<I>(graph: &'a UndirectedSparseGraph, rooting: F, inverse: Option<I>) -> Self
38 where
39 I: Fn(&M::T, &M::T) -> M::T,
40 {
41 let dp = vec![M::unit(); graph.vertices_size()];
42 let ep = vec![M::unit(); graph.vertices_size() * 2];
43 let mut self_ = Self {
44 graph,
45 dp,
46 ep,
47 rooting,
48 };
49 self_.rerooting(inverse);
50 self_
51 }
52 #[inline]
53 fn eidx(&self, u: usize, a: Neighbor<usize, usize>) -> usize {
54 a.label + self.graph.edges_size() * (u > a.to) as usize
55 }
56 #[inline]
57 fn reidx(&self, u: usize, a: Neighbor<usize, usize>) -> usize {
58 a.label + self.graph.edges_size() * (u < a.to) as usize
59 }
60 #[inline]
61 fn merge(&self, x: &M::T, y: &M::T) -> M::T {
62 M::operate(x, y)
63 }
64 #[inline]
65 fn add_subroot(&self, x: &M::T, vid: usize, eid: usize) -> M::T {
66 (self.rooting)(x, vid, Some(eid))
67 }
68 #[inline]
69 fn add_root(&self, x: &M::T, vid: usize) -> M::T {
70 (self.rooting)(x, vid, None)
71 }
72 fn rerooting<I: Fn(&M::T, &M::T) -> M::T>(&mut self, inverse: Option<I>) {
73 let (order, parents) = self.graph.tree_order(0);
74 for &u in order.iter().skip(1).rev() {
75 let mut sum = M::unit();
76 let mut parent = None;
77 for a in self.graph.neighbors(u) {
78 if a.to == parents[u] {
79 parent = Some(a);
80 } else {
81 sum = self.merge(&sum, &self.ep[self.eidx(u, a)]);
82 }
83 }
84 let a = parent.unwrap();
85 let i = self.reidx(u, a);
86 self.ep[i] = self.add_subroot(&sum, u, a.label);
87 if inverse.is_some() {
88 self.dp[u] = sum;
89 }
90 }
91 if let Some(inverse) = inverse {
92 for u in order {
93 let sum = if u == 0 {
94 self.graph.neighbors(u).fold(M::unit(), |sum, a| {
95 self.merge(&sum, &self.ep[self.eidx(u, a)])
96 })
97 } else {
98 let a = self
99 .graph
100 .neighbors(u)
101 .find(|a| a.to == parents[u])
102 .unwrap();
103 self.merge(&self.dp[u], &self.ep[self.eidx(u, a)])
104 };
105 self.dp[u] = self.add_root(&sum, u);
106 for a in self.graph.neighbors(u) {
107 if a.to != parents[u] {
108 let value = inverse(&sum, &self.ep[self.eidx(u, a)]);
109 let i = self.reidx(u, a);
110 self.ep[i] = self.add_subroot(&value, u, a.label);
111 }
112 }
113 }
114 return;
115 }
116 let mut prefix = Vec::new();
117 for u in order {
118 prefix.clear();
119 prefix.push(M::unit());
120 for a in self.graph.neighbors(u) {
121 prefix.push(self.merge(prefix.last().unwrap(), &self.ep[self.eidx(u, a)]));
122 }
123 self.dp[u] = self.add_root(prefix.last().unwrap(), u);
124 let mut suffix = M::unit();
125 for (k, a) in self.graph.neighbors(u).enumerate().rev() {
126 if a.to != parents[u] {
127 let i = self.reidx(u, a);
128 self.ep[i] = self.add_subroot(&self.merge(&prefix[k], &suffix), u, a.label);
129 }
130 suffix = self.merge(&self.ep[self.eidx(u, a)], &suffix);
131 }
132 }
133 }
134}
135
136#[cfg(test)]
137mod tests {
138 use super::*;
139 use crate::{
140 algebra::{AdditiveOperation, ConcatenateOperation},
141 tools::{Xorshift, testutil::exhaustive_sequences},
142 tree::{MixedTree, PathTree, StarTree},
143 };
144
145 #[test]
146 fn test_rerooting() {
147 let mut rng = Xorshift::default();
148 let mut graphs = Vec::new();
149 for n in 1..=5 {
150 for parents in exhaustive_sequences(0..n, n - 1..=n - 1) {
151 if parents.iter().enumerate().all(|(i, &p)| p <= i) {
152 graphs.push(UndirectedSparseGraph::from_edges(
153 n,
154 parents
155 .into_iter()
156 .enumerate()
157 .map(|(i, p)| (p, i + 1))
158 .collect(),
159 ));
160 }
161 }
162 }
163 for n in 1..=32 {
164 graphs.extend([
165 rng.random(PathTree(n)),
166 rng.random(StarTree(n)),
167 rng.random(MixedTree(n)),
168 ]);
169 }
170 for graph in graphs {
171 let n = graph.vertices_size();
172 let dp = ReRooting::<ConcatenateOperation<_>, _>::new(&graph, |xs, v, edge| {
173 let mut xs = xs.clone();
174 xs.push((v, edge));
175 xs
176 });
177 let sums = ReRooting::<AdditiveOperation<i64>, _>::new_with_inverse(
178 &graph,
179 |&sum, v, edge| sum + v as i64 + edge.map_or(0, |e| e as i64),
180 );
181 for root in 0..n {
182 for parent in std::iter::once(None).chain(graph.neighbors(root).map(Some)) {
183 let mut expected = Vec::new();
184 let mut stack = vec![(
185 root,
186 parent.map_or(n, |a| a.to),
187 parent.map(|a| a.label),
188 false,
189 )];
190 while let Some((u, p, edge, visited)) = stack.pop() {
191 if visited {
192 expected.push((u, edge));
193 } else {
194 stack.push((u, p, edge, true));
195 stack.extend(
196 graph
197 .neighbors(u)
198 .rev()
199 .filter(|a| a.to != p)
200 .map(|a| (a.to, u, Some(a.label), false)),
201 );
202 }
203 }
204 let actual = if let Some(parent) = parent {
205 &dp.ep[dp.reidx(root, parent)]
206 } else {
207 &dp.dp[root]
208 };
209 assert_eq!(*actual, expected);
210 let sum = if let Some(parent) = parent {
211 sums.ep[sums.reidx(root, parent)]
212 } else {
213 sums.dp[root]
214 };
215 assert_eq!(
216 sum,
217 expected
218 .iter()
219 .map(|&(v, e)| v as i64 + e.map_or(0, |e| e as i64))
220 .sum()
221 );
222 }
223 }
224 }
225 }
226}