Skip to main content

competitive/data_structure/
union_find.rs

1use super::{Group, Monoid};
2use std::{
3    collections::HashMap,
4    fmt::{self, Debug},
5    marker::PhantomData,
6    mem::{MaybeUninit, needs_drop, swap},
7};
8
9pub struct UnionFindBase<U, F, M, P, H>
10where
11    U: UnionStrategy,
12    F: FindStrategy,
13    M: UfMergeSpec,
14    P: Monoid,
15    H: UndoStrategy<UfCell<M::Data, P>>,
16{
17    cells: Vec<UfCell<M::Data, P>>,
18    merger: M,
19    history: H::History,
20    _marker: PhantomData<fn() -> (U, F)>,
21}
22
23impl<U, F, M, P, H> Clone for UnionFindBase<U, F, M, P, H>
24where
25    U: UnionStrategy,
26    F: FindStrategy,
27    M: UfMergeSpec<Data: Clone> + Clone,
28    P: Monoid,
29    H: UndoStrategy<UfCell<M::Data, P>, History: Clone>,
30{
31    fn clone(&self) -> Self {
32        Self {
33            cells: self.cells.clone(),
34            merger: self.merger.clone(),
35            history: self.history.clone(),
36            _marker: self._marker,
37        }
38    }
39}
40
41impl<U, F, M, P, H> Debug for UnionFindBase<U, F, M, P, H>
42where
43    U: UnionStrategy,
44    F: FindStrategy,
45    M: UfMergeSpec<Data: Debug>,
46    P: Monoid<T: Debug>,
47    H: UndoStrategy<UfCell<M::Data, P>, History: Debug>,
48{
49    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
50        f.debug_struct("UnionFindBase")
51            .field("cells", &self.cells)
52            .field("history", &self.history)
53            .finish()
54    }
55}
56
57pub struct UfCell<D, P>
58where
59    P: Monoid,
60{
61    // Roots store `!info` and initialized data; children store a parent and no data.
62    parent_or_info: i32,
63    data: MaybeUninit<D>,
64    potential: P::T,
65}
66
67impl<D, P> Clone for UfCell<D, P>
68where
69    D: Clone,
70    P: Monoid,
71{
72    fn clone(&self) -> Self {
73        let is_root = self.is_root();
74        let potential = if is_root {
75            P::unit()
76        } else {
77            self.potential.clone()
78        };
79        Self {
80            parent_or_info: self.parent_or_info,
81            data: if is_root {
82                MaybeUninit::new(self.data().clone())
83            } else {
84                MaybeUninit::uninit()
85            },
86            potential,
87        }
88    }
89}
90
91impl<D, P> Debug for UfCell<D, P>
92where
93    D: Debug,
94    P: Monoid<T: Debug>,
95{
96    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
97        if let Some(info) = self.root_info() {
98            f.debug_tuple("Root").field(&(info, self.data())).finish()
99        } else {
100            f.debug_tuple("Child")
101                .field(&(self.parent().unwrap(), &self.potential))
102                .finish()
103        }
104    }
105}
106
107impl<D, P> UfCell<D, P>
108where
109    P: Monoid,
110{
111    fn root(info: u32, data: D) -> Self {
112        Self {
113            parent_or_info: !(info as i32),
114            data: MaybeUninit::new(data),
115            potential: P::unit(),
116        }
117    }
118
119    fn root_info(&self) -> Option<u32> {
120        (self.parent_or_info < 0).then_some((!self.parent_or_info) as u32)
121    }
122
123    fn set_root_info(&mut self, info: u32) {
124        self.parent_or_info = !(info as i32);
125    }
126
127    fn parent(&self) -> Option<usize> {
128        (self.parent_or_info >= 0).then_some(self.parent_or_info as usize)
129    }
130
131    fn set_child(&mut self, parent: usize, potential: P::T) {
132        if needs_drop::<D>() && self.is_root() {
133            unsafe { self.data.assume_init_drop() };
134        }
135        self.parent_or_info = parent as i32;
136        self.potential = potential;
137    }
138
139    fn is_root(&self) -> bool {
140        self.parent_or_info < 0
141    }
142
143    fn data(&self) -> &D {
144        unsafe { self.data.assume_init_ref() }
145    }
146
147    fn data_mut(&mut self) -> &mut D {
148        unsafe { self.data.assume_init_mut() }
149    }
150}
151
152impl<D, P> Drop for UfCell<D, P>
153where
154    P: Monoid,
155{
156    fn drop(&mut self) {
157        if needs_drop::<D>() && self.is_root() {
158            unsafe { self.data.assume_init_drop() };
159        }
160    }
161}
162
163pub trait FindStrategy {
164    const CHENGE_ROOT: bool;
165}
166
167pub enum PathCompression {}
168
169impl FindStrategy for PathCompression {
170    const CHENGE_ROOT: bool = true;
171}
172
173impl FindStrategy for () {
174    const CHENGE_ROOT: bool = false;
175}
176
177pub trait UnionStrategy {
178    fn single_info() -> u32;
179    fn check_directoin(parent: &u32, child: &u32) -> bool;
180    fn unite(parent: &u32, child: &u32) -> u32;
181}
182
183pub enum UnionBySize {}
184
185impl UnionStrategy for UnionBySize {
186    fn single_info() -> u32 {
187        1
188    }
189
190    fn check_directoin(parent: &u32, child: &u32) -> bool {
191        parent >= child
192    }
193
194    fn unite(parent: &u32, child: &u32) -> u32 {
195        parent + child
196    }
197}
198
199pub enum UnionByRank {}
200
201impl UnionStrategy for UnionByRank {
202    fn single_info() -> u32 {
203        0
204    }
205
206    fn check_directoin(parent: &u32, child: &u32) -> bool {
207        parent >= child
208    }
209
210    fn unite(parent: &u32, child: &u32) -> u32 {
211        parent + (parent == child) as u32
212    }
213}
214
215impl UnionStrategy for () {
216    fn single_info() -> u32 {
217        0
218    }
219
220    fn check_directoin(_parent: &u32, _child: &u32) -> bool {
221        false
222    }
223
224    fn unite(_parent: &u32, _child: &u32) -> u32 {
225        0
226    }
227}
228
229pub trait UfMergeSpec {
230    type Data;
231    fn merge(&mut self, to: &mut Self::Data, from: &mut Self::Data);
232}
233
234#[derive(Debug, Clone)]
235pub struct FnMerger<T, F> {
236    f: F,
237    _marker: PhantomData<fn() -> T>,
238}
239
240impl<T, F> UfMergeSpec for FnMerger<T, F>
241where
242    F: FnMut(&mut T, &mut T),
243{
244    type Data = T;
245
246    fn merge(&mut self, to: &mut Self::Data, from: &mut Self::Data) {
247        (self.f)(to, from)
248    }
249}
250
251impl UfMergeSpec for () {
252    type Data = ();
253
254    fn merge(&mut self, _to: &mut Self::Data, _from: &mut Self::Data) {}
255}
256
257pub trait UndoStrategy<T> {
258    const UNDOABLE: bool;
259
260    type History: Default;
261
262    fn unite(history: &mut Self::History, x: usize, y: usize, cells: &[T]);
263
264    fn undo_unite(history: &mut Self::History, cells: &mut [T]);
265}
266
267pub enum Undoable {}
268
269impl<T> UndoStrategy<T> for Undoable
270where
271    T: Clone,
272{
273    const UNDOABLE: bool = true;
274
275    type History = Vec<[(usize, T); 2]>;
276
277    fn unite(history: &mut Self::History, x: usize, y: usize, cells: &[T]) {
278        let cx = cells[x].clone();
279        let cy = cells[y].clone();
280        history.push([(x, cx), (y, cy)]);
281    }
282
283    fn undo_unite(history: &mut Self::History, cells: &mut [T]) {
284        if let Some([(x, cx), (y, cy)]) = history.pop() {
285            cells[x] = cx;
286            cells[y] = cy;
287        }
288    }
289}
290
291impl<T> UndoStrategy<T> for () {
292    const UNDOABLE: bool = false;
293
294    type History = ();
295
296    fn unite(_history: &mut Self::History, _x: usize, _y: usize, _cells: &[T]) {}
297
298    fn undo_unite(_history: &mut Self::History, _cells: &mut [T]) {}
299}
300
301impl<U, F, P, H> UnionFindBase<U, F, (), P, H>
302where
303    U: UnionStrategy,
304    F: FindStrategy,
305    P: Monoid,
306    H: UndoStrategy<UfCell<(), P>>,
307{
308    pub fn new(n: usize) -> Self {
309        let cells: Vec<_> = (0..n).map(|_| UfCell::root(U::single_info(), ())).collect();
310        Self {
311            cells,
312            merger: (),
313            history: Default::default(),
314            _marker: PhantomData,
315        }
316    }
317    pub fn push(&mut self) {
318        self.cells.push(UfCell::root(U::single_info(), ()));
319    }
320}
321
322impl<U, F, T, Merge, P, H> UnionFindBase<U, F, FnMerger<T, Merge>, P, H>
323where
324    U: UnionStrategy,
325    F: FindStrategy,
326    Merge: FnMut(&mut T, &mut T),
327    P: Monoid,
328    H: UndoStrategy<UfCell<T, P>>,
329{
330    pub fn new_with_merger(n: usize, mut init: impl FnMut(usize) -> T, merge: Merge) -> Self {
331        let cells: Vec<_> = (0..n)
332            .map(|i| UfCell::root(U::single_info(), init(i)))
333            .collect();
334        Self {
335            cells,
336            merger: FnMerger {
337                f: merge,
338                _marker: PhantomData,
339            },
340            history: Default::default(),
341            _marker: PhantomData,
342        }
343    }
344}
345
346impl<F, M, P, H> UnionFindBase<UnionBySize, F, M, P, H>
347where
348    F: FindStrategy,
349    M: UfMergeSpec,
350    P: Monoid,
351    H: UndoStrategy<UfCell<M::Data, P>>,
352{
353    pub fn size(&mut self, x: usize) -> usize {
354        let root = self.find_root(x);
355        self.root_info(root).unwrap() as usize
356    }
357}
358
359impl<U, F, M, P, H> UnionFindBase<U, F, M, P, H>
360where
361    U: UnionStrategy,
362    F: FindStrategy,
363    M: UfMergeSpec,
364    P: Monoid,
365    H: UndoStrategy<UfCell<M::Data, P>>,
366{
367    fn root_info(&self, x: usize) -> Option<u32> {
368        self.cells[x].root_info()
369    }
370
371    fn set_root_info(&mut self, x: usize, info: u32) {
372        self.cells[x].set_root_info(info);
373    }
374
375    pub fn same(&mut self, x: usize, y: usize) -> bool {
376        self.find_root(x) == self.find_root(y)
377    }
378
379    pub fn merge_data(&mut self, x: usize) -> &M::Data {
380        let root = self.find_root(x);
381        self.cells[root].data()
382    }
383
384    pub fn merge_data_mut(&mut self, x: usize) -> &mut M::Data {
385        let root = self.find_root(x);
386        self.cells[root].data_mut()
387    }
388
389    pub fn roots(&self) -> impl Iterator<Item = usize> + '_ {
390        (0..self.cells.len()).filter(|&x| self.cells[x].is_root())
391    }
392
393    pub fn all_group_members(&mut self) -> HashMap<usize, Vec<usize>> {
394        let mut groups_map = HashMap::new();
395        for x in 0..self.cells.len() {
396            let r = self.find_root(x);
397            groups_map.entry(r).or_insert_with(Vec::new).push(x);
398        }
399        groups_map
400    }
401
402    pub fn find(&mut self, x: usize) -> (usize, P::T) {
403        let mut current = x;
404        let mut potential = P::unit();
405        while let Some(parent) = self.cells[current].parent() {
406            let current_potential = self.cells[current].potential.clone();
407            potential = P::operate(&current_potential, &potential);
408            if F::CHENGE_ROOT
409                && let Some(parent_parent) = self.cells[parent].parent()
410            {
411                let potential = P::operate(&self.cells[parent].potential, &current_potential);
412                self.cells[current].set_child(parent_parent, potential);
413            }
414            current = parent;
415        }
416        (current, potential)
417    }
418
419    pub fn find_root(&mut self, x: usize) -> usize {
420        let mut current = x;
421        while let Some(parent) = self.cells[current].parent() {
422            if F::CHENGE_ROOT
423                && let Some(parent_parent) = self.cells[parent].parent()
424            {
425                let potential = P::operate(
426                    &self.cells[parent].potential,
427                    &self.cells[current].potential,
428                );
429                self.cells[current].set_child(parent_parent, potential);
430            }
431            current = parent;
432        }
433        current
434    }
435
436    pub fn unite_noninv(&mut self, x: usize, y: usize, potential: P::T) -> bool {
437        let (rx, potx) = self.find(x);
438        let ry = self.find_root(y);
439        if rx == ry || y != ry {
440            return false;
441        }
442        H::unite(&mut self.history, rx, ry, &self.cells);
443        {
444            let ptr = self.cells.as_mut_ptr();
445            let (cx, cy) = unsafe { (&mut *ptr.add(rx), &mut *ptr.add(ry)) };
446            self.merger.merge(cx.data_mut(), cy.data_mut());
447        }
448        let info = U::unite(&self.root_info(rx).unwrap(), &self.root_info(ry).unwrap());
449        self.set_root_info(rx, info);
450        self.cells[ry].set_child(rx, P::operate(&potx, &potential));
451        true
452    }
453}
454
455impl<U, F, M, P, H> UnionFindBase<U, F, M, P, H>
456where
457    U: UnionStrategy,
458    F: FindStrategy,
459    M: UfMergeSpec,
460    P: Group,
461    H: UndoStrategy<UfCell<M::Data, P>>,
462{
463    pub fn difference(&mut self, x: usize, y: usize) -> Option<P::T> {
464        let (rx, potx) = self.find(x);
465        let (ry, poty) = self.find(y);
466        if rx == ry {
467            Some(P::operate(&P::inverse(&potx), &poty))
468        } else {
469            None
470        }
471    }
472
473    pub fn unite_with(&mut self, x: usize, y: usize, potential: P::T) -> bool {
474        let (mut rx, potx) = self.find(x);
475        let (mut ry, poty) = self.find(y);
476        if rx == ry {
477            return false;
478        }
479        let mut xinfo = self.root_info(rx).unwrap();
480        let mut yinfo = self.root_info(ry).unwrap();
481        let inverse = !U::check_directoin(&xinfo, &yinfo);
482        let potential = if inverse {
483            P::rinv_operate(&poty, &P::operate(&potx, &potential))
484        } else {
485            P::operate(&potx, &P::rinv_operate(&potential, &poty))
486        };
487        if inverse {
488            swap(&mut rx, &mut ry);
489            swap(&mut xinfo, &mut yinfo);
490        }
491        H::unite(&mut self.history, rx, ry, &self.cells);
492        {
493            let ptr = self.cells.as_mut_ptr();
494            let (cx, cy) = unsafe { (&mut *ptr.add(rx), &mut *ptr.add(ry)) };
495            self.merger.merge(cx.data_mut(), cy.data_mut());
496        }
497        self.set_root_info(rx, U::unite(&xinfo, &yinfo));
498        self.cells[ry].set_child(rx, potential);
499        true
500    }
501
502    pub fn unite(&mut self, x: usize, y: usize) -> bool {
503        self.unite_with(x, y, P::unit())
504    }
505}
506
507impl<U, M, P, H> UnionFindBase<U, (), M, P, H>
508where
509    U: UnionStrategy,
510    M: UfMergeSpec,
511    P: Monoid,
512    H: UndoStrategy<UfCell<M::Data, P>>,
513{
514    pub fn undo(&mut self) {
515        H::undo_unite(&mut self.history, &mut self.cells);
516    }
517}
518
519pub type UnionFind = UnionFindBase<UnionBySize, PathCompression, (), (), ()>;
520pub type MergingUnionFind<T, M> =
521    UnionFindBase<UnionBySize, PathCompression, FnMerger<T, M>, (), ()>;
522pub type PotentializedUnionFind<P> = UnionFindBase<UnionBySize, PathCompression, (), P, ()>;
523pub type UndoableUnionFind = UnionFindBase<UnionBySize, (), (), (), Undoable>;
524
525#[cfg(test)]
526mod tests {
527    use super::*;
528    use crate::{
529        algebra::{Invertible, LinearOperation, Magma, Unital},
530        graph::{Graph, UndirectedSparseGraph},
531        num::mint_basic::MInt998244353 as M,
532        rand,
533        tools::Xorshift,
534        tree::MixedTree,
535    };
536    use std::collections::HashSet;
537
538    fn distinct_edges(rng: &mut Xorshift, n: usize, m: usize) -> Vec<(usize, usize)> {
539        let mut edges = vec![];
540        for x in 0..n {
541            for y in 0..n {
542                edges.push((x, y));
543            }
544        }
545        rng.shuffle(&mut edges);
546        edges.truncate(m);
547        edges
548    }
549
550    fn dfs(
551        g: &UndirectedSparseGraph,
552        u: usize,
553        vis: &mut [bool],
554        f: &mut impl FnMut(usize),
555        f2: &mut impl FnMut(usize, usize, usize),
556    ) {
557        vis[u] = true;
558        f(u);
559        for a in g.neighbors(u) {
560            if !vis[a.to] {
561                f2(u, a.to, a.label);
562                dfs(g, a.to, vis, f, f2);
563            }
564        }
565    }
566
567    #[test]
568    fn test_union_find() {
569        const N: usize = 20;
570        let mut rng = Xorshift::default();
571        for _ in 0..1000 {
572            rand!(rng, n: 1..=N, m: 1..=n * n);
573            let edges = distinct_edges(&mut rng, n, m);
574
575            macro_rules! test_uf {
576                ($union:ty, $find:ty) => {{
577                    let mut uf = UnionFindBase::<$union, $find, FnMerger<Vec<usize>, _>, (), ()>::new_with_merger(n, |i| vec![i], |x, y| x.append(y));
578                    for &(x, y) in &edges {
579                        uf.unite(x, y);
580                    }
581                    let g = UndirectedSparseGraph::from_edges(n, edges.to_vec());
582                    let mut id = vec![!0; n];
583                    {
584                        let mut vis = vec![false; n];
585                        for x in 0..n {
586                            if vis[x] {
587                                continue;
588                            }
589                            let mut set = HashSet::new();
590                            dfs(
591                                &g,
592                                x,
593                                &mut vis,
594                                &mut |x| {
595                                    set.insert(x);
596                                },
597                                &mut |_, _, _| {},
598                            );
599                            for s in set {
600                                id[s] = x;
601                            }
602                        }
603                    }
604                    for x in 0..n {
605                        for y in 0..n {
606                            assert_eq!(id[x] == id[y], uf.same(x, y));
607                        }
608                        assert_eq!(
609                            (0..n).filter(|&y| id[x] == id[y]).collect::<HashSet<_>>(),
610                            uf.merge_data(x).iter().cloned().collect()
611                        );
612                    }
613                }};
614            }
615            test_uf!(UnionBySize, PathCompression);
616            test_uf!(UnionByRank, PathCompression);
617            test_uf!((), PathCompression);
618            test_uf!(UnionBySize, ());
619            test_uf!(UnionByRank, ());
620            test_uf!((), ());
621        }
622    }
623
624    #[test]
625    fn test_potential_union_find() {
626        const N: usize = 20;
627        let mut rng = Xorshift::default();
628        type G = LinearOperation<M>;
629        for _ in 0..1000 {
630            rand!(rng, n: 1..=N, g: MixedTree(n), p: [(.., ..); n - 1], k: 0..n);
631
632            macro_rules! test_uf {
633                ($union:ty, $find:ty) => {{
634                    let mut uf = UnionFindBase::<$union, $find, (), G, ()>::new(n);
635                    for (i, &(u, v)) in g.edges.iter().enumerate().take(k) {
636                        uf.unite_with(u, v, p[i]);
637                    }
638                    for x in 0..n {
639                        let mut vis = vec![false; n];
640                        let mut dp = vec![None; n];
641                        dp[x] = Some(G::unit());
642                        dfs(&g, x, &mut vis, &mut |_| {}, &mut |u, to, id| {
643                            let p = if g.edges[id] == (u, to) {
644                                p[id]
645                            } else {
646                                G::inverse(&p[id])
647                            };
648                            if id < k {
649                                if let Some(d) = dp[u] {
650                                    dp[to] = Some(G::operate(&d, &p));
651                                }
652                            }
653                        });
654                        for (y, d) in dp.into_iter().enumerate() {
655                            assert_eq!(d, uf.difference(x, y));
656                        }
657                    }
658                }};
659            }
660            test_uf!(UnionBySize, PathCompression);
661            test_uf!(UnionByRank, PathCompression);
662            test_uf!((), PathCompression);
663            test_uf!(UnionBySize, ());
664            test_uf!(UnionByRank, ());
665            test_uf!((), ());
666        }
667    }
668
669    #[test]
670    fn test_undoable_union_find() {
671        const N: usize = 10;
672        const M: usize = 200;
673        let mut rng = Xorshift::default();
674        for _ in 0..10 {
675            rand!(rng, n: 1..=N, m: 1..=M, g: MixedTree(m), p: [(0..n, 0..n); m]);
676
677            macro_rules! test_uf {
678                ($union:ty, $find:ty) => {{
679                    let uf = UnionFind::new(n);
680                    let mut uf2 = UnionFindBase::<$union, $find, (), (), Undoable>::new(n);
681                    fn dfs(
682                        n: usize,
683                        g: &UndirectedSparseGraph,
684                        u: usize,
685                        vis: &mut [bool],
686                        mut uf: UnionFindBase<UnionBySize, PathCompression, (), (), ()>,
687                        uf2: &mut UnionFindBase<$union, $find, (), (), Undoable>,
688                        p: &[(usize, usize)],
689                    ) {
690                        vis[u] = true;
691                        for x in 0..n {
692                            for y in 0..n {
693                                assert_eq!(uf.same(x, y), uf2.same(x, y));
694                            }
695                        }
696                        for a in g.neighbors(u) {
697                            if !vis[a.to] {
698                                let (x, y) = p[a.label];
699                                let mut uf = uf.clone();
700                                uf.unite(x, y);
701                                let merged = uf2.unite(x, y);
702                                dfs(n, g, a.to, vis, uf, uf2, p);
703                                if merged {
704                                    uf2.undo();
705                                }
706                            }
707                        }
708                    }
709                    for u in 0..m {
710                        dfs(n, &g, u, &mut vec![false; m], uf.clone(), &mut uf2, &p);
711                    }
712                }};
713            }
714            test_uf!(UnionBySize, ());
715            test_uf!(UnionByRank, ());
716            test_uf!((), ());
717        }
718    }
719}