Skip to main content

competitive/tree/
centroid_decomposition.rs

1use super::{Graph, UndirectedSparseGraph};
2use std::mem::swap;
3
4#[derive(Debug, Clone)]
5struct RootedTree {
6    parents: Vec<usize>,
7    vs: Vec<usize>,
8}
9
10impl RootedTree {
11    fn len(&self) -> usize {
12        self.vs.len()
13    }
14
15    fn split_centroid(self) -> CentroidSplit {
16        let n = self.len();
17        assert!(n > 2);
18        let parents = &self.parents;
19        let vs = &self.vs;
20        let mut size = vec![1; n];
21        let mut c = usize::MAX;
22        for i in (0..n).rev() {
23            if size[i] >= n.div_ceil(2) {
24                c = i;
25                break;
26            }
27            size[parents[i]] += size[i];
28        }
29        let mut side = vec![u8::MAX; n];
30        let mut order = vec![usize::MAX; n];
31        order[c] = 0;
32        let mut count = 1usize;
33        let mut taken = 0usize;
34        for u in 1..n {
35            if parents[u] == c && taken + size[u] <= (n - 1) / 2 {
36                taken += size[u];
37                side[u] = 0;
38                order[u] = count;
39                count += 1;
40            }
41        }
42        for u in 1..n {
43            if side[parents[u]] == 0 {
44                side[u] = 0;
45                order[u] = count;
46                count += 1;
47            }
48        }
49        let lsize = count - 1;
50        {
51            let mut u = parents[c];
52            while u != usize::MAX {
53                side[u] = 1;
54                order[u] = count;
55                count += 1;
56                u = parents[u];
57            }
58        }
59        for u in 0..n {
60            if u != c && side[u] == u8::MAX {
61                side[u] = 1;
62                order[u] = count;
63                count += 1;
64            }
65        }
66        assert_eq!(count, n);
67        let mut whole_parents = vec![usize::MAX; n];
68        let mut whole_vs = vec![usize::MAX; n];
69        for u in 0..n {
70            whole_vs[order[u]] = vs[u];
71        }
72        for u in 1..n {
73            let mut x = order[u];
74            let mut y = order[parents[u]];
75            if x > y {
76                swap(&mut x, &mut y);
77            }
78            whole_parents[y] = x;
79        }
80        let left = RootedTree {
81            parents: whole_parents[..=lsize].to_vec(),
82            vs: whole_vs[..=lsize].to_vec(),
83        };
84        let right = RootedTree {
85            parents: std::iter::once(usize::MAX)
86                .chain(
87                    whole_parents[lsize + 1..]
88                        .iter()
89                        .map(|&p| if p == 0 { 0 } else { p - lsize }),
90                )
91                .collect(),
92            vs: std::iter::once(whole_vs[0])
93                .chain(whole_vs[lsize + 1..].iter().copied())
94                .collect(),
95        };
96        CentroidSplit {
97            whole: RootedTree {
98                parents: whole_parents,
99                vs: whole_vs,
100            },
101            left,
102            right,
103            lsize,
104        }
105    }
106
107    fn centroid_decomposition(self, f: &mut impl FnMut(&[usize], &[usize], usize, usize)) {
108        if self.len() <= 2 {
109            return;
110        }
111        let split = self.split_centroid();
112        f(
113            &split.whole.parents,
114            &split.whole.vs,
115            split.lsize,
116            split.rsize(),
117        );
118        split.left.centroid_decomposition(f);
119        split.right.centroid_decomposition(f);
120    }
121}
122
123impl From<&UndirectedSparseGraph> for RootedTree {
124    fn from(graph: &UndirectedSparseGraph) -> Self {
125        let n = graph.vertices_size();
126        let mut vs = Vec::with_capacity(n);
127        let mut parent = vec![usize::MAX; n];
128        vs.push(0usize);
129        for i in 0..n {
130            let u = vs[i];
131            for a in graph.neighbors(u) {
132                if a.to != parent[u] {
133                    vs.push(a.to);
134                    parent[a.to] = u;
135                }
136            }
137        }
138        let mut new_idx = vec![0; n];
139        for (i, &v) in vs.iter().enumerate() {
140            new_idx[v] = i;
141        }
142        let mut parents = vec![usize::MAX; n];
143        for v in 1..n {
144            parents[new_idx[v]] = new_idx[parent[v]];
145        }
146        Self { parents, vs }
147    }
148}
149
150#[derive(Debug)]
151struct CentroidSplit {
152    whole: RootedTree,
153    left: RootedTree,
154    right: RootedTree,
155    lsize: usize,
156}
157
158impl CentroidSplit {
159    fn rsize(&self) -> usize {
160        self.whole.len() - self.lsize - 1
161    }
162}
163
164#[derive(Debug, Clone, Copy)]
165struct ContourInfo {
166    comp: u32,
167    dep: u32,
168}
169
170#[derive(Debug, Clone)]
171pub struct ContourQueryRange {
172    comp_range: Vec<usize>,
173    info_indptr: Vec<usize>,
174    infos: Vec<ContourInfo>,
175    local_info: Vec<(usize, usize)>,
176    local_offsets: Vec<usize>,
177    local_masks: Vec<u32>,
178}
179
180impl ContourQueryRange {
181    pub fn len(&self) -> usize {
182        self.comp_range.last().copied().unwrap_or_default()
183    }
184
185    pub fn is_empty(&self) -> bool {
186        self.len() == 0
187    }
188
189    pub fn component_sizes(&self) -> impl ExactSizeIterator<Item = usize> + '_ {
190        self.comp_range.windows(2).map(|range| range[1] - range[0])
191    }
192
193    /// Calls `f(component, index)` for each position representing `v`.
194    pub fn for_each_index(&self, v: usize, mut f: impl FnMut(usize, usize)) {
195        for info in &self.infos[self.info_indptr[v]..self.info_indptr[v + 1]] {
196            f(info.comp as usize, info.dep as usize);
197        }
198        let (comp, index) = self.local_info[v];
199        if comp != usize::MAX {
200            f(
201                self.comp_range.len() - 1 - self.local_offsets.len() + comp,
202                index,
203            );
204        }
205    }
206
207    /// Calls `f(component, start, end)` for disjoint ranges at distances in `l..r` from `v`.
208    /// The ranges exclude `v` itself, even when `l == 0`.
209    pub fn for_each_contour_range(
210        &self,
211        v: usize,
212        l: usize,
213        r: usize,
214        mut f: impl FnMut(usize, usize, usize),
215    ) {
216        for info in &self.infos[self.info_indptr[v]..self.info_indptr[v + 1]] {
217            let comp = (info.comp ^ 1) as usize;
218            let start = self.comp_range[comp];
219            let len = self.comp_range[comp + 1] - start;
220            let lo = l.saturating_sub(info.dep as usize).min(len);
221            let hi = r.saturating_sub(info.dep as usize).min(len);
222            if lo < hi {
223                f(comp, lo, hi);
224            }
225        }
226        let (local, index) = self.local_info[v];
227        if local != usize::MAX {
228            let comp = self.comp_range.len() - 1 - self.local_offsets.len() + local;
229            let len = self.comp_range[comp + 1] - self.comp_range[comp];
230            let lo = l.max(1).min(len);
231            let hi = r.min(len);
232            if lo < hi {
233                let offset = self.local_offsets[local] + index * (len + 1);
234                let mut mask = self.local_masks[offset + hi] ^ self.local_masks[offset + lo];
235                while mask != 0 {
236                    let start = mask.trailing_zeros();
237                    let end = start + (mask >> start).trailing_ones();
238                    f(comp, start as usize, end as usize);
239                    mask &= mask.wrapping_add(1 << start);
240                }
241            }
242        }
243    }
244}
245
246impl UndirectedSparseGraph {
247    /// 1/3 centroid decomposition
248    ///
249    /// - f: (parents: &[usize], vs: &[usize], lsize: usize, rsize: usize)
250    /// - 0: root, 1..=lsize: left subtree, lsize+1..=lsize+rsize: right subtree
251    pub fn centroid_decomposition(&self, mut f: impl FnMut(&[usize], &[usize], usize, usize)) {
252        if self.vertices_size() <= 1 {
253            return;
254        }
255        RootedTree::from(self).centroid_decomposition(&mut f);
256    }
257
258    pub fn contour_query_range(&self) -> ContourQueryRange {
259        let n = self.vertices_size();
260        assert!(n <= u32::MAX as usize / 2);
261        if n <= 1 {
262            return ContourQueryRange {
263                comp_range: vec![0],
264                info_indptr: vec![0; n + 1],
265                infos: vec![],
266                local_info: vec![(usize::MAX, 0); n],
267                local_offsets: vec![],
268                local_masks: vec![],
269            };
270        }
271        let (vertices, graph) = {
272            let (vertices, parents) = self.tree_order(0);
273            let mut indices = vec![0; n];
274            for (i, &v) in vertices.iter().enumerate() {
275                indices[v] = i;
276            }
277            let edges = vertices
278                .iter()
279                .enumerate()
280                .skip(1)
281                .map(|(i, &v)| (i, indices[parents[v]]))
282                .collect();
283            let graph = UndirectedSparseGraph::from_edges(n, edges);
284            (vertices, graph)
285        };
286        let mut comp_range = vec![0usize];
287        let mut vertex_info = Vec::with_capacity(n * (n.ilog2() as usize + 1));
288        let mut info_indptr = vec![0usize; n + 1];
289        let mut local_info = vec![(usize::MAX, 0); n];
290        let mut local_offsets = Vec::new();
291        let mut local_masks = Vec::new();
292        let mut distances = Vec::new();
293        let mut local_sizes = Vec::new();
294        let mut removed = vec![false; n];
295        let mut parents = vec![usize::MAX; n];
296        let mut sizes = vec![0usize; n];
297        let mut tasks = vec![0];
298        let mut order = Vec::with_capacity(n);
299        let mut entries = Vec::with_capacity(n);
300        let mut boundaries = Vec::new();
301        let mut groups = Vec::new();
302        while let Some(root) = tasks.pop() {
303            order.clear();
304            order.push(root);
305            parents[root] = usize::MAX;
306            let mut i = 0;
307            while i < order.len() {
308                let v = order[i];
309                sizes[v] = 1;
310                for edge in graph.neighbors(v) {
311                    if !removed[edge.to] && edge.to != parents[v] {
312                        parents[edge.to] = v;
313                        order.push(edge.to);
314                    }
315                }
316                i += 1;
317            }
318            if order.len() <= 32 {
319                let len = order.len();
320                if len > 1 {
321                    let comp = local_offsets.len();
322                    let offset = local_masks.len();
323                    local_offsets.push(offset);
324                    local_sizes.push(len);
325                    local_masks.resize(offset + len * (len + 1), 0u32);
326                    distances.clear();
327                    distances.resize(len * len, 0u8);
328                    for (i, &v) in order.iter().enumerate() {
329                        local_info[vertices[v]] = (comp, i);
330                        sizes[v] = i;
331                        local_masks[offset + i * (len + 1) + 1] = 1 << i;
332                    }
333                    for (i, &v) in order.iter().enumerate().skip(1) {
334                        let parent = sizes[parents[v]];
335                        for j in 0..i {
336                            let distance = distances[parent * len + j] + 1;
337                            distances[i * len + j] = distance;
338                            distances[j * len + i] = distance;
339                            local_masks[offset + i * (len + 1) + distance as usize + 1] |= 1 << j;
340                            local_masks[offset + j * (len + 1) + distance as usize + 1] |= 1 << i;
341                        }
342                    }
343                    for row in local_masks[offset..].chunks_exact_mut(len + 1) {
344                        for d in 1..=len {
345                            row[d] |= row[d - 1];
346                        }
347                    }
348                }
349                continue;
350            }
351            let mut centroid = root;
352            for &v in order.iter().rev() {
353                if sizes[v] >= order.len().div_ceil(2) {
354                    centroid = v;
355                    break;
356                }
357                sizes[parents[v]] += sizes[v];
358            }
359            removed[centroid] = true;
360            entries.clear();
361            entries.push((centroid, 0));
362            boundaries.clear();
363            boundaries.extend([0, 1]);
364            for edge in graph.neighbors(centroid) {
365                let v = edge.to;
366                if removed[v] {
367                    continue;
368                }
369                tasks.push(v);
370                parents[v] = centroid;
371                let mut i = entries.len();
372                entries.push((v, 1));
373                while i < entries.len() {
374                    let (v, distance) = entries[i];
375                    for edge in graph.neighbors(v) {
376                        if !removed[edge.to] && edge.to != parents[v] {
377                            parents[edge.to] = v;
378                            entries.push((edge.to, distance + 1));
379                        }
380                    }
381                    i += 1;
382                }
383                boundaries.push(entries.len());
384            }
385            groups.push((0, boundaries.len() - 1));
386            while let Some((first, last)) = groups.pop() {
387                if last - first < 2 {
388                    continue;
389                }
390                let weight = boundaries[last] - boundaries[first];
391                let target = boundaries[first] + weight.div_ceil(2);
392                let mut middle =
393                    first + 1 + boundaries[first + 1..last].partition_point(|&p| p < target);
394                middle = middle.min(last - 1);
395                if middle > first + 1 {
396                    let x = boundaries[middle] - boundaries[first];
397                    let y = boundaries[middle - 1] - boundaries[first];
398                    if y.max(weight - y) < x.max(weight - x) {
399                        middle -= 1;
400                    }
401                }
402                for (l, r) in [(first, middle), (middle, last)] {
403                    let comp = comp_range.len() - 1;
404                    let mut max_distance = 0;
405                    for &(v, dep) in &entries[boundaries[l]..boundaries[r]] {
406                        vertex_info.push((
407                            vertices[v] as u32,
408                            ContourInfo {
409                                comp: comp as u32,
410                                dep: dep as u32,
411                            },
412                        ));
413                        info_indptr[vertices[v] + 1] += 1;
414                        max_distance = max_distance.max(dep);
415                    }
416                    comp_range.push(comp_range[comp] + max_distance + 1);
417                }
418                groups.extend([(middle, last), (first, middle)]);
419            }
420        }
421        for len in local_sizes {
422            comp_range.push(comp_range.last().unwrap() + len);
423        }
424        for v in 1..=n {
425            info_indptr[v] += info_indptr[v - 1];
426        }
427        let mut infos = vec![ContourInfo { comp: 0, dep: 0 }; vertex_info.len()];
428        let mut positions = info_indptr.clone();
429        for (v, info) in vertex_info {
430            let v = v as usize;
431            infos[positions[v]] = info;
432            positions[v] += 1;
433        }
434        ContourQueryRange {
435            comp_range,
436            info_indptr,
437            infos,
438            local_info,
439            local_offsets,
440            local_masks,
441        }
442    }
443}
444
445#[cfg(test)]
446mod tests {
447    use crate::{
448        graph::UndirectedSparseGraph,
449        tools::{Xorshift, testutil::exhaustive_sequences},
450        tree::{MixedTree, PathTree, StarTree},
451    };
452
453    #[test]
454    fn test_contour_query_range() {
455        let mut rng = Xorshift::default();
456        let mut graphs = vec![UndirectedSparseGraph::from_edges(0, vec![])];
457        for n in 1..=5 {
458            for parents in exhaustive_sequences(0..n, n - 1..=n - 1) {
459                if parents.iter().enumerate().all(|(i, &p)| p <= i) {
460                    graphs.push(UndirectedSparseGraph::from_edges(
461                        n,
462                        parents
463                            .into_iter()
464                            .enumerate()
465                            .map(|(i, p)| (p, i + 1))
466                            .collect(),
467                    ));
468                }
469            }
470        }
471        for n in 1..=80 {
472            graphs.extend([rng.random(PathTree(n)), rng.random(StarTree(n))]);
473        }
474        graphs.extend((0..200).map(|_| rng.random(MixedTree(1usize..80))));
475        for graph in graphs {
476            let n = graph.vertices_size();
477            let query = graph.contour_query_range();
478            let mut values = vec![0i64; n];
479            let mut data: Vec<_> = query.component_sizes().map(|n| vec![0i64; n]).collect();
480            assert_eq!(query.len(), data.iter().map(Vec::len).sum());
481            assert_eq!(query.is_empty(), n <= 1);
482            let updates: Vec<_> = if n <= 5 {
483                (0..n)
484                    .flat_map(|u| (-1..=1).map(move |delta| (u, delta)))
485                    .collect()
486            } else {
487                (0..200)
488                    .map(|_| (rng.random(0..n), rng.random(-100..=100i64)))
489                    .collect()
490            };
491            for (u, delta) in updates {
492                values[u] += delta;
493                query.for_each_index(u, |c, i| data[c][i] += delta);
494                let ranges: Vec<_> = if n <= 5 {
495                    (0..n)
496                        .flat_map(|v| {
497                            (0..=n).flat_map(move |l| (l..=n + 1).map(move |r| (v, l, r)))
498                        })
499                        .collect()
500                } else {
501                    let v = rng.random(0..n);
502                    let l = rng.random(0..=n);
503                    vec![(v, l, rng.random(l..=n + 1))]
504                };
505                for (v, l, r) in ranges {
506                    let distances = graph.tree_depth(v);
507                    let expected: i64 = (0..n)
508                        .filter(|&u| {
509                            u != v && l <= distances[u] as usize && (distances[u] as usize) < r
510                        })
511                        .map(|u| values[u])
512                        .sum();
513                    let mut actual = 0;
514                    query.for_each_contour_range(v, l, r, |c, start, end| {
515                        actual += data[c][start..end].iter().sum::<i64>()
516                    });
517                    assert_eq!(actual, expected);
518                }
519            }
520        }
521    }
522}