Skip to main content

competitive/graph/
general_matching.rs

1use std::mem::swap;
2
3#[derive(Debug, Clone)]
4pub struct GeneralMatching {
5    size: usize,
6    graph: Vec<Vec<usize>>,
7    mate: Vec<usize>,
8    matching_size: usize,
9    parent: Vec<usize>,
10    base: Vec<usize>,
11    state: Vec<u8>,
12    seen: Vec<usize>,
13    queue: Vec<usize>,
14    timestamp: usize,
15}
16
17impl GeneralMatching {
18    pub fn new(size: usize) -> Self {
19        Self {
20            size,
21            graph: vec![vec![]; size],
22            mate: vec![size; size],
23            matching_size: 0,
24            parent: vec![size; size + 1],
25            base: (0..=size).collect(),
26            state: vec![Self::UNSEEN; size],
27            seen: vec![0; size],
28            queue: Vec::with_capacity(size),
29            timestamp: 0,
30        }
31    }
32    pub fn add_edge(&mut self, u: usize, v: usize) {
33        assert!(u < self.size);
34        assert!(v < self.size);
35        if u == v {
36            return;
37        }
38        self.graph[u].push(v);
39        self.graph[v].push(u);
40        self.matching_size = !0;
41    }
42    pub fn from_edges(size: usize, edges: &[(usize, usize)]) -> Self {
43        let mut this = Self::new(size);
44        for &(u, v) in edges {
45            this.add_edge(u, v);
46        }
47        this
48    }
49    pub fn maximum_matching(&mut self) -> Vec<(usize, usize)> {
50        self.compute();
51        let mut res = Vec::with_capacity(self.matching_size);
52        for v in 0..self.size {
53            let u = self.mate[v];
54            if u != self.size && v < u {
55                res.push((v, u));
56            }
57        }
58        res
59    }
60    fn compute(&mut self) {
61        if self.matching_size != !0 {
62            return;
63        }
64        self.matching_size = self.mate.iter().filter(|&&mate| mate != self.size).count() / 2;
65
66        for v in 0..self.size {
67            if self.mate[v] == self.size && self.augment_from(v) {
68                self.matching_size += 1;
69            }
70        }
71    }
72
73    const UNSEEN: u8 = 0;
74    const OUTER: u8 = 1;
75    const INNER: u8 = 2;
76
77    fn find(&mut self, mut v: usize) -> usize {
78        while self.base[v] != v {
79            let parent = self.base[v];
80            self.base[v] = self.base[parent];
81            v = self.base[v];
82        }
83        v
84    }
85
86    fn lca(&mut self, mut u: usize, mut v: usize) -> usize {
87        self.timestamp += 1;
88        u = self.find(u);
89        v = self.find(v);
90        loop {
91            if u != self.size {
92                if self.seen[u] == self.timestamp {
93                    return u;
94                }
95                self.seen[u] = self.timestamp;
96                u = self.find(self.parent[self.mate[u]]);
97            }
98            swap(&mut u, &mut v);
99        }
100    }
101
102    fn contract(&mut self, mut v: usize, mut child: usize, ancestor: usize) {
103        while self.find(v) != ancestor {
104            self.parent[v] = child;
105            child = self.mate[v];
106            if self.state[child] == Self::INNER {
107                self.state[child] = Self::OUTER;
108                self.queue.push(child);
109            }
110            if self.base[v] == v {
111                self.base[v] = ancestor;
112            }
113            if self.base[child] == child {
114                self.base[child] = ancestor;
115            }
116            v = self.parent[child];
117        }
118    }
119
120    fn augment_from(&mut self, root: usize) -> bool {
121        for (v, base) in self.base.iter_mut().enumerate() {
122            *base = v;
123        }
124        self.state.fill(Self::UNSEEN);
125        self.queue.clear();
126        self.state[root] = Self::OUTER;
127        self.queue.push(root);
128        let mut head = 0;
129        while head < self.queue.len() {
130            let u = self.queue[head];
131            head += 1;
132            for edge in 0..self.graph[u].len() {
133                let v = self.graph[u][edge];
134                if self.state[v] == Self::UNSEEN {
135                    self.parent[v] = u;
136                    self.state[v] = Self::INNER;
137                    if self.mate[v] == self.size {
138                        let mut v = v;
139                        let mut u = u;
140                        while u != self.size {
141                            let next = self.mate[u];
142                            self.mate[u] = v;
143                            self.mate[v] = u;
144                            v = next;
145                            u = self.parent[v];
146                        }
147                        return true;
148                    }
149                    let v = self.mate[v];
150                    self.state[v] = Self::OUTER;
151                    self.queue.push(v);
152                } else if self.state[v] == Self::OUTER && self.find(u) != self.find(v) {
153                    let ancestor = self.lca(u, v);
154                    self.contract(u, v, ancestor);
155                    self.contract(v, u, ancestor);
156                }
157            }
158        }
159        false
160    }
161}
162
163#[cfg(test)]
164mod tests {
165    use super::*;
166    use crate::{rand, tools::Xorshift};
167
168    fn brute_maximum_matching(n: usize, edges: &[(usize, usize)]) -> usize {
169        let mut adj = vec![vec![false; n]; n];
170        for &(u, v) in edges {
171            adj[u][v] = true;
172            adj[v][u] = true;
173        }
174        let mut dp = vec![0usize; 1 << n];
175        for mask in 1usize..1 << n {
176            let i = mask.trailing_zeros() as usize;
177            let mask_without_i = mask & !(1 << i);
178            let mut best = dp[mask_without_i];
179            let mut m = mask_without_i;
180            while m != 0 {
181                let j = m.trailing_zeros() as usize;
182                if adj[i][j] {
183                    let val = 1 + dp[mask_without_i & !(1 << j)];
184                    if val > best {
185                        best = val;
186                    }
187                }
188                m &= m - 1;
189            }
190            dp[mask] = best;
191        }
192        dp[(1 << n) - 1]
193    }
194
195    #[test]
196    fn test_general_matching() {
197        const Q: usize = 200;
198        const N: usize = 10;
199        let mut rng = Xorshift::default();
200        for _ in 0..Q {
201            rand!(rng, n: 1..=N);
202            let mut edges = vec![];
203            for i in 0..n {
204                for j in i + 1..n {
205                    rand!(rng, b: 0..2usize);
206                    if b == 1 {
207                        edges.push((i, j));
208                    }
209                }
210            }
211            rand!(rng, split: 0..=edges.len());
212            let mut gm = GeneralMatching::from_edges(n, &edges[..split]);
213            assert_eq!(
214                gm.maximum_matching().len(),
215                brute_maximum_matching(n, &edges[..split])
216            );
217            for &(u, v) in &edges[split..] {
218                gm.add_edge(u, v);
219            }
220            let matching = gm.maximum_matching();
221            let mut used = vec![false; n];
222            let mut adj = vec![vec![false; n]; n];
223            for &(u, v) in &edges {
224                adj[u][v] = true;
225                adj[v][u] = true;
226            }
227            for &(u, v) in &matching {
228                assert!(u < v);
229                assert!(adj[u][v]);
230                assert!(!used[u]);
231                assert!(!used[v]);
232                used[u] = true;
233                used[v] = true;
234            }
235            assert_eq!(matching.len(), brute_maximum_matching(n, &edges));
236        }
237    }
238}