Skip to main content

competitive/tree/
heavy_light_decomposition.rs

1use super::{CartesianTree, Graph, Monoid, UndirectedSparseGraph};
2use std::ops::Range;
3
4#[derive(Clone, Debug)]
5pub struct HeavyLightDecomposition {
6    nodes: Vec<HeavyLightNode>,
7    order: Vec<usize>,
8}
9
10#[derive(Clone, Copy, Debug)]
11struct HeavyLightNode {
12    parent: u32,
13    size: u32,
14    head: u32,
15    index: u32,
16}
17
18impl UndirectedSparseGraph {
19    pub fn hld(&self, root: usize) -> HeavyLightDecomposition {
20        HeavyLightDecomposition::new(root, self)
21    }
22}
23
24impl HeavyLightDecomposition {
25    pub fn new(root: usize, graph: &UndirectedSparseGraph) -> Self {
26        let n = graph.vertices_size();
27        assert!(n <= u32::MAX as usize);
28        let mut self_ = Self {
29            nodes: vec![
30                HeavyLightNode {
31                    parent: n as u32,
32                    size: 1,
33                    head: n as u32,
34                    index: 0
35                };
36                n
37            ],
38            order: Vec::with_capacity(n),
39        };
40        self_.order.push(root);
41        for i in 0..n {
42            let u = self_.order[i];
43            for a in graph.neighbors(u) {
44                if a.to != self_.nodes[u].parent as usize {
45                    self_.nodes[a.to].parent = u as u32;
46                    self_.order.push(a.to);
47                }
48            }
49        }
50        for &u in self_.order.iter().skip(1).rev() {
51            let p = self_.nodes[u].parent as usize;
52            self_.nodes[p].size += self_.nodes[u].size;
53            let heavy = self_.nodes[p].head as usize;
54            if heavy == n || self_.nodes[heavy].size <= self_.nodes[u].size {
55                self_.nodes[p].head = u as u32;
56            }
57        }
58        self_.order.clear();
59        let mut stack = vec![root];
60        while let Some(head) = stack.pop() {
61            let chain_start = self_.order.len() as u32;
62            let chain_parent = self_.nodes[head].parent;
63            let mut u = head;
64            while u != n {
65                let heavy = self_.nodes[u].head as usize;
66                let parent = self_.nodes[u].parent as usize;
67                self_.nodes[u].head = chain_start;
68                self_.nodes[u].parent = chain_parent;
69                self_.nodes[u].index = self_.order.len() as u32;
70                self_.order.push(u);
71                for a in graph.neighbors(u).rev() {
72                    if a.to != parent && a.to != heavy {
73                        stack.push(a.to);
74                    }
75                }
76                u = heavy;
77            }
78        }
79        self_
80    }
81
82    #[inline]
83    pub fn len(&self) -> usize {
84        self.order.len()
85    }
86
87    #[inline]
88    pub fn is_empty(&self) -> bool {
89        self.order.is_empty()
90    }
91
92    #[inline]
93    pub fn root(&self) -> usize {
94        self.order[0]
95    }
96
97    #[inline]
98    pub fn parent(&self, v: usize) -> Option<usize> {
99        let index = self.nodes[v].index as usize;
100        if index == self.nodes[v].head as usize {
101            ((self.nodes[v].parent as usize) < self.len()).then_some(self.nodes[v].parent as usize)
102        } else {
103            Some(self.order[index - 1])
104        }
105    }
106
107    #[inline]
108    pub fn index(&self, v: usize) -> usize {
109        self.nodes[v].index as usize
110    }
111
112    #[inline]
113    pub fn vertex(&self, index: usize) -> usize {
114        self.order[index]
115    }
116
117    #[inline]
118    pub fn subtree_size(&self, v: usize) -> usize {
119        self.nodes[v].size as usize
120    }
121
122    #[inline]
123    pub fn subtree_range(&self, v: usize) -> Range<usize> {
124        self.nodes[v].index as usize..self.nodes[v].index as usize + self.nodes[v].size as usize
125    }
126
127    #[inline]
128    pub fn is_ancestor(&self, ancestor: usize, v: usize) -> bool {
129        self.subtree_range(ancestor)
130            .contains(&(self.nodes[v].index as usize))
131    }
132
133    #[inline]
134    pub fn kth_ancestor(&self, mut v: usize, mut k: usize) -> Option<usize> {
135        loop {
136            let head = self.nodes[v].head as usize;
137            let chain_len = self.nodes[v].index as usize - head;
138            if k <= chain_len {
139                return Some(self.order[self.nodes[v].index as usize - k]);
140            }
141            k -= chain_len + 1;
142            v = self.nodes[v].parent as usize;
143            if v == self.len() {
144                return None;
145            }
146        }
147    }
148
149    #[inline]
150    pub fn lca(&self, mut u: usize, mut v: usize) -> usize {
151        while self.nodes[u].head != self.nodes[v].head {
152            if self.nodes[u].index > self.nodes[v].index {
153                u = self.nodes[u].parent as usize;
154            } else {
155                v = self.nodes[v].parent as usize;
156            }
157        }
158        if self.nodes[u].index < self.nodes[v].index {
159            u
160        } else {
161            v
162        }
163    }
164
165    #[inline]
166    pub fn distance(&self, u: usize, v: usize) -> usize {
167        let (up, down) = self.path_lengths(u, v);
168        up + down
169    }
170
171    #[inline]
172    pub fn jump(&self, mut u: usize, mut v: usize, mut k: usize) -> Option<usize> {
173        let target = v;
174        let mut down = 0;
175        while self.nodes[u].head != self.nodes[v].head {
176            if self.nodes[u].index > self.nodes[v].index {
177                let up = self.nodes[u].index as usize - self.nodes[u].head as usize + 1;
178                if k < up {
179                    return Some(self.order[self.nodes[u].index as usize - k]);
180                }
181                k -= up;
182                u = self.nodes[u].parent as usize;
183            } else {
184                down += self.nodes[v].index as usize - self.nodes[v].head as usize + 1;
185                v = self.nodes[v].parent as usize;
186            }
187        }
188        if self.nodes[u].index >= self.nodes[v].index {
189            let up = self.nodes[u].index as usize - self.nodes[v].index as usize;
190            if k <= up {
191                return Some(self.order[self.nodes[u].index as usize - k]);
192            }
193            k -= up;
194        } else {
195            down += self.nodes[v].index as usize - self.nodes[u].index as usize;
196        }
197        down.checked_sub(k)
198            .and_then(|k| self.kth_ancestor(target, k))
199    }
200
201    #[inline]
202    fn path_lengths(&self, mut u: usize, mut v: usize) -> (usize, usize) {
203        let (mut up, mut down) = (0, 0);
204        while self.nodes[u].head != self.nodes[v].head {
205            if self.nodes[u].index > self.nodes[v].index {
206                up += self.nodes[u].index as usize - self.nodes[u].head as usize + 1;
207                u = self.nodes[u].parent as usize;
208            } else {
209                down += self.nodes[v].index as usize - self.nodes[v].head as usize + 1;
210                v = self.nodes[v].parent as usize;
211            }
212        }
213        if self.nodes[u].index > self.nodes[v].index {
214            up += self.nodes[u].index as usize - self.nodes[v].index as usize;
215        } else {
216            down += self.nodes[v].index as usize - self.nodes[u].index as usize;
217        }
218        (up, down)
219    }
220
221    /// Calls `f` once for each nonempty DFS-index range on the vertex path.
222    /// The callback order is unspecified.
223    #[inline]
224    pub fn path_vertices<F: FnMut(usize, usize)>(&self, u: usize, v: usize, f: F) {
225        self.path(u, v, false, f);
226    }
227
228    /// Calls `f` once for each nonempty DFS-index range on the edge path.
229    /// Each index represents the deeper endpoint of an edge. The callback order is unspecified.
230    #[inline]
231    pub fn path_edges<F: FnMut(usize, usize)>(&self, u: usize, v: usize, f: F) {
232        self.path(u, v, true, f);
233    }
234
235    #[inline]
236    fn path<F: FnMut(usize, usize)>(&self, mut u: usize, mut v: usize, is_edge: bool, mut f: F) {
237        loop {
238            if self.nodes[u].index > self.nodes[v].index {
239                std::mem::swap(&mut u, &mut v);
240            }
241            if self.nodes[u].head == self.nodes[v].head {
242                break;
243            }
244            f(
245                self.nodes[v].head as usize,
246                self.nodes[v].index as usize + 1,
247            );
248            v = self.nodes[v].parent as usize;
249        }
250        let l = self.nodes[u].index as usize + usize::from(is_edge);
251        let r = self.nodes[v].index as usize + 1;
252        if l < r {
253            f(l, r);
254        }
255    }
256
257    /// Folds a vertex path in `u`-to-`v` order.
258    /// `forward` folds a DFS-index range from left to right, and `reverse` folds it from right to
259    /// left.
260    #[inline]
261    pub fn fold_vertices<
262        M: Monoid,
263        F1: FnMut(usize, usize) -> M::T,
264        F2: FnMut(usize, usize) -> M::T,
265    >(
266        &self,
267        u: usize,
268        v: usize,
269        forward: F1,
270        reverse: F2,
271    ) -> M::T {
272        self.fold::<M, _, _>(u, v, false, forward, reverse)
273    }
274
275    /// Folds an edge path in `u`-to-`v` order.
276    /// Each index represents the deeper endpoint of an edge. `forward` folds a DFS-index range
277    /// from left to right, and `reverse` folds it from right to left.
278    #[inline]
279    pub fn fold_edges<
280        M: Monoid,
281        F1: FnMut(usize, usize) -> M::T,
282        F2: FnMut(usize, usize) -> M::T,
283    >(
284        &self,
285        u: usize,
286        v: usize,
287        forward: F1,
288        reverse: F2,
289    ) -> M::T {
290        self.fold::<M, _, _>(u, v, true, forward, reverse)
291    }
292
293    #[inline]
294    fn fold<M: Monoid, F1: FnMut(usize, usize) -> M::T, F2: FnMut(usize, usize) -> M::T>(
295        &self,
296        mut u: usize,
297        mut v: usize,
298        is_edge: bool,
299        mut forward: F1,
300        mut reverse: F2,
301    ) -> M::T {
302        let (mut left, mut right) = (M::unit(), M::unit());
303        while self.nodes[u].head != self.nodes[v].head {
304            if self.nodes[u].index > self.nodes[v].index {
305                left = M::operate(
306                    &left,
307                    &reverse(
308                        self.nodes[u].head as usize,
309                        self.nodes[u].index as usize + 1,
310                    ),
311                );
312                u = self.nodes[u].parent as usize;
313            } else {
314                right = M::operate(
315                    &forward(
316                        self.nodes[v].head as usize,
317                        self.nodes[v].index as usize + 1,
318                    ),
319                    &right,
320                );
321                v = self.nodes[v].parent as usize;
322            }
323        }
324        let middle = if self.nodes[u].index > self.nodes[v].index {
325            reverse(
326                self.nodes[v].index as usize + usize::from(is_edge),
327                self.nodes[u].index as usize + 1,
328            )
329        } else {
330            forward(
331                self.nodes[u].index as usize + usize::from(is_edge),
332                self.nodes[v].index as usize + 1,
333            )
334        };
335        M::operate(&M::operate(&left, &middle), &right)
336    }
337}
338
339pub struct HeavyLightPathFold<'a, M: Monoid> {
340    tree: &'a HeavyLightDecomposition,
341    nodes: Vec<PathFoldNode<M::T>>,
342}
343
344struct PathFoldNode<T> {
345    parent: u32,
346    children: [u32; 2],
347    priority: u32,
348    value: T,
349    aggregate: [T; 2],
350    prefix: [T; 2],
351}
352
353impl HeavyLightDecomposition {
354    /// `values` is indexed by vertex, not by DFS index.
355    pub fn build_fold<M: Monoid>(&self, values: &[M::T]) -> HeavyLightPathFold<'_, M> {
356        assert_eq!(values.len(), self.len());
357        let mut fold = HeavyLightPathFold {
358            tree: self,
359            nodes: self
360                .order
361                .iter()
362                .map(|&v| PathFoldNode {
363                    parent: u32::MAX,
364                    children: [u32::MAX; 2],
365                    priority: 0,
366                    value: values[v].clone(),
367                    aggregate: [values[v].clone(), values[v].clone()],
368                    prefix: [values[v].clone(), values[v].clone()],
369                })
370                .collect(),
371        };
372        let mut start = 0;
373        let mut priorities = Vec::new();
374        let mut stack = Vec::new();
375        while start < self.len() {
376            let mut end = start + 1;
377            while end < self.len() && self.nodes[self.order[end]].head as usize == start {
378                end += 1;
379            }
380            priorities.clear();
381            let mut sum = 0usize;
382            for i in start..end {
383                let weight = self.subtree_size(self.order[i])
384                    - if i + 1 < end {
385                        self.subtree_size(self.order[i + 1])
386                    } else {
387                        0
388                    };
389                let priority = (sum ^ (sum + weight)).ilog2();
390                sum += weight;
391                fold.nodes[i].priority = priority;
392                priorities.push(std::cmp::Reverse(priority));
393            }
394            let cartesian = CartesianTree::new(&priorities);
395            for i in start..end {
396                let parent = cartesian.parents[i - start];
397                fold.nodes[i].parent = if parent == usize::MAX {
398                    u32::MAX
399                } else {
400                    (parent + start) as u32
401                };
402                fold.nodes[i].children = cartesian.children[i - start].map(|v| {
403                    if v == usize::MAX {
404                        u32::MAX
405                    } else {
406                        (v + start) as u32
407                    }
408                });
409            }
410            stack.clear();
411            stack.push(cartesian.root + start);
412            let mut i = 0;
413            while i < stack.len() {
414                stack.extend(
415                    fold.nodes[stack[i]]
416                        .children
417                        .into_iter()
418                        .filter(|&v| v != u32::MAX)
419                        .map(|v| v as usize),
420                );
421                i += 1;
422            }
423            for &i in stack.iter().rev() {
424                fold.pull(i);
425            }
426            start = end;
427        }
428        fold
429    }
430}
431
432impl<M: Monoid> HeavyLightPathFold<'_, M> {
433    #[inline(always)]
434    fn pull(&mut self, i: usize) {
435        let [l, r] = self.nodes[i].children.map(|v| {
436            if v == u32::MAX {
437                usize::MAX
438            } else {
439                v as usize
440            }
441        });
442        self.nodes[i].prefix = if l == usize::MAX {
443            [self.nodes[i].value.clone(), self.nodes[i].value.clone()]
444        } else {
445            [
446                M::operate(&self.nodes[l].aggregate[0], &self.nodes[i].value),
447                M::operate(&self.nodes[i].value, &self.nodes[l].aggregate[1]),
448            ]
449        };
450        self.nodes[i].aggregate = if r == usize::MAX {
451            self.nodes[i].prefix.clone()
452        } else {
453            [
454                M::operate(&self.nodes[i].prefix[0], &self.nodes[r].aggregate[0]),
455                M::operate(&self.nodes[r].aggregate[1], &self.nodes[i].prefix[1]),
456            ]
457        };
458    }
459
460    pub fn set(&mut self, vertex: usize, value: M::T) {
461        let mut i = self.tree.index(vertex);
462        self.nodes[i].value = value;
463        while i != usize::MAX {
464            self.pull(i);
465            i = if self.nodes[i].parent == u32::MAX {
466                usize::MAX
467            } else {
468                self.nodes[i].parent as usize
469            };
470        }
471    }
472
473    fn fold_prefix<const REVERSE: bool>(&self, k: usize) -> M::T {
474        let mut result = M::unit();
475        let mut i = k;
476        while i != usize::MAX {
477            if i <= k {
478                result = if REVERSE {
479                    M::operate(&result, &self.nodes[i].prefix[1])
480                } else {
481                    M::operate(&self.nodes[i].prefix[0], &result)
482                };
483            }
484            i = if self.nodes[i].parent == u32::MAX {
485                usize::MAX
486            } else {
487                self.nodes[i].parent as usize
488            };
489        }
490        result
491    }
492
493    fn fold_range<const REVERSE: bool>(&self, l: usize, r: usize) -> M::T {
494        let (mut u, mut v) = (l, r);
495        let (mut left, mut right) = (M::unit(), M::unit());
496        while u != v {
497            if self.nodes[u].priority < self.nodes[v].priority {
498                if u >= l {
499                    let child = self.nodes[u].children[1];
500                    if REVERSE {
501                        left = M::operate(&self.nodes[u].value, &left);
502                        if child != u32::MAX {
503                            left = M::operate(&self.nodes[child as usize].aggregate[1], &left);
504                        }
505                    } else {
506                        left = M::operate(&left, &self.nodes[u].value);
507                        if child != u32::MAX {
508                            left = M::operate(&left, &self.nodes[child as usize].aggregate[0]);
509                        }
510                    }
511                }
512                u = self.nodes[u].parent as usize;
513            } else {
514                if v <= r {
515                    right = if REVERSE {
516                        M::operate(&right, &self.nodes[v].prefix[1])
517                    } else {
518                        M::operate(&self.nodes[v].prefix[0], &right)
519                    };
520                }
521                v = self.nodes[v].parent as usize;
522            }
523        }
524        if REVERSE {
525            M::operate(&M::operate(&right, &self.nodes[u].value), &left)
526        } else {
527            M::operate(&M::operate(&left, &self.nodes[u].value), &right)
528        }
529    }
530
531    /// Folds the vertex values in `u`-to-`v` order.
532    #[inline(always)]
533    pub fn fold_vertices(&self, mut u: usize, mut v: usize) -> M::T {
534        let (mut left, mut right) = (M::unit(), M::unit());
535        while self.tree.nodes[u].head != self.tree.nodes[v].head {
536            if self.tree.index(u) > self.tree.index(v) {
537                left = M::operate(&left, &self.fold_prefix::<true>(self.tree.index(u)));
538                u = self.tree.nodes[u].parent as usize;
539            } else {
540                right = M::operate(&self.fold_prefix::<false>(self.tree.index(v)), &right);
541                v = self.tree.nodes[v].parent as usize;
542            }
543        }
544        let middle = if self.tree.index(u) > self.tree.index(v) {
545            self.fold_range::<true>(self.tree.index(v), self.tree.index(u))
546        } else {
547            self.fold_range::<false>(self.tree.index(u), self.tree.index(v))
548        };
549        M::operate(&M::operate(&left, &middle), &right)
550    }
551}
552
553#[cfg(test)]
554mod tests {
555    use super::*;
556    use crate::{
557        algebra::ConcatenateOperation,
558        tools::{Xorshift, testutil::exhaustive_sequences},
559        tree::{MixedTree, PathTree, StarTree},
560    };
561
562    fn parent_and_depth(graph: &UndirectedSparseGraph, root: usize) -> (Vec<usize>, Vec<usize>) {
563        let n = graph.vertices_size();
564        let mut parent = vec![n; n];
565        let mut depth = vec![0; n];
566        let mut stack = vec![root];
567        while let Some(u) = stack.pop() {
568            for a in graph.neighbors(u) {
569                if a.to != parent[u] {
570                    parent[a.to] = u;
571                    depth[a.to] = depth[u] + 1;
572                    stack.push(a.to);
573                }
574            }
575        }
576        (parent, depth)
577    }
578
579    fn path(mut u: usize, mut v: usize, parent: &[usize], depth: &[usize]) -> Vec<usize> {
580        let mut left = vec![];
581        let mut right = vec![];
582        while depth[u] > depth[v] {
583            left.push(u);
584            u = parent[u];
585        }
586        while depth[v] > depth[u] {
587            right.push(v);
588            v = parent[v];
589        }
590        while u != v {
591            left.push(u);
592            right.push(v);
593            u = parent[u];
594            v = parent[v];
595        }
596        left.push(u);
597        left.extend(right.into_iter().rev());
598        left
599    }
600
601    fn verify(graph: UndirectedSparseGraph, root: usize) {
602        let n = graph.vertices_size();
603        let (parent, depth) = parent_and_depth(&graph, root);
604        let hld = graph.hld(root);
605        assert_eq!(hld.len(), n);
606        assert_eq!(hld.root(), root);
607
608        let mut vertex = vec![0; n];
609        for v in 0..n {
610            vertex[hld.index(v)] = v;
611        }
612
613        for v in 0..n {
614            assert_eq!(hld.vertex(hld.index(v)), v);
615            assert_eq!(hld.parent(v), (parent[v] < n).then_some(parent[v]));
616            for k in 0..=depth[v] + 1 {
617                let mut ancestor = Some(v);
618                for _ in 0..k {
619                    ancestor = ancestor.and_then(|u| (parent[u] < n).then_some(parent[u]));
620                }
621                assert_eq!(hld.kth_ancestor(v, k), ancestor);
622            }
623
624            let expected: Vec<_> = (0..n)
625                .filter(|&u| {
626                    let mut u = u;
627                    while depth[u] > depth[v] {
628                        u = parent[u];
629                    }
630                    u == v
631                })
632                .collect();
633            let range = hld.subtree_range(v);
634            let mut actual: Vec<_> = range.clone().map(|i| vertex[i]).collect();
635            actual.sort_unstable();
636            assert_eq!(actual, expected);
637            assert_eq!(hld.subtree_size(v), range.len());
638            for u in 0..n {
639                assert_eq!(hld.is_ancestor(v, u), expected.contains(&u));
640            }
641        }
642
643        let mut fold =
644            hld.build_fold::<ConcatenateOperation<_>>(&(0..n).map(|v| vec![v]).collect::<Vec<_>>());
645        for u in 0..n {
646            fold.set(u, vec![u + n]);
647            for v in 0..n {
648                let expected = path(u, v, &parent, &depth);
649                assert_eq!(
650                    fold.fold_vertices(u, v),
651                    expected
652                        .iter()
653                        .map(|&w| if w <= u { w + n } else { w })
654                        .collect::<Vec<_>>()
655                );
656                let lca = *expected.iter().min_by_key(|&&v| depth[v]).unwrap();
657                assert_eq!(hld.lca(u, v), lca);
658                assert_eq!(hld.distance(u, v), expected.len() - 1);
659                for k in 0..=expected.len() {
660                    assert_eq!(hld.jump(u, v, k), expected.get(k).copied());
661                }
662
663                let actual = hld.fold_vertices::<ConcatenateOperation<_>, _, _>(
664                    u,
665                    v,
666                    |l, r| (l..r).map(|i| vertex[i]).collect(),
667                    |l, r| (l..r).rev().map(|i| vertex[i]).collect(),
668                );
669                assert_eq!(actual, expected);
670
671                let mut actual = vec![];
672                hld.path_vertices(u, v, |l, r| {
673                    actual.extend((l..r).map(|i| vertex[i]));
674                });
675                actual.sort_unstable();
676                let mut expected_unordered = expected.clone();
677                expected_unordered.sort_unstable();
678                assert_eq!(actual, expected_unordered);
679
680                let expected_edges: Vec<_> =
681                    expected.iter().copied().filter(|&v| v != lca).collect();
682                let actual = hld.fold_edges::<ConcatenateOperation<_>, _, _>(
683                    u,
684                    v,
685                    |l, r| (l..r).map(|i| vertex[i]).collect(),
686                    |l, r| (l..r).rev().map(|i| vertex[i]).collect(),
687                );
688                assert_eq!(actual, expected_edges);
689            }
690        }
691    }
692
693    #[test]
694    fn heavy_light_decomposition_against_naive() {
695        let mut rng = Xorshift::default();
696        for n in 1..=5 {
697            for parents in exhaustive_sequences(0..n, n - 1..=n - 1) {
698                if parents.iter().enumerate().all(|(i, &p)| p <= i) {
699                    let edges: Vec<_> = parents
700                        .into_iter()
701                        .enumerate()
702                        .map(|(i, p)| (p, i + 1))
703                        .collect();
704                    for root in 0..n {
705                        verify(UndirectedSparseGraph::from_edges(n, edges.clone()), root);
706                    }
707                }
708            }
709        }
710        for n in 1..=20 {
711            verify(rng.random(PathTree(n)), rng.random(0..n));
712            verify(rng.random(StarTree(n)), rng.random(0..n));
713        }
714        for _ in 0..100 {
715            let n = rng.random(1..=40);
716            verify(rng.random(MixedTree(n)), rng.random(0..n));
717        }
718    }
719}