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 #[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 #[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 #[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 #[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 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 #[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}