Skip to main content

competitive/tree/
rerooting.rs

1//! dynamic programming on all-rooted trees
2
3use crate::algebra::{AbelianGroup, Monoid};
4use crate::graph::{Graph, Neighbor, UndirectedSparseGraph};
5
6#[codesnip::entry("ReRooting", include("algebra", "tree_order"))]
7/// dynamic programming on all-rooted trees
8///
9/// Neighbors are merged in adjacency order before applying `rooting`.
10#[derive(Clone, Debug)]
11pub struct ReRooting<'a, M: Monoid, F: Fn(&M::T, usize, Option<usize>) -> M::T> {
12    graph: &'a UndirectedSparseGraph,
13    /// dp\[v\]: result of v-rooted tree
14    pub dp: Vec<M::T>,
15    /// ep\[e\]: result of e-subtree; e >= m denotes the reverse direction.
16    pub ep: Vec<M::T>,
17    /// rooting(data, vid, (Optional)eid): add root node(vid), result subtree is edge(eid)
18    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}