Skip to main content

competitive/graph/
bipartite_matching.rs

1use std::{collections::VecDeque, mem::swap};
2
3#[derive(Debug, Clone)]
4pub struct BipartiteMatching {
5    left_size: usize,
6    right_size: usize,
7    left_graph: Vec<Vec<usize>>,
8    right_graph: Vec<Vec<usize>>,
9    left_match: Vec<Option<usize>>,
10    right_match: Vec<Option<usize>>,
11    matching_size: usize,
12}
13
14impl BipartiteMatching {
15    pub fn new(left_size: usize, right_size: usize) -> Self {
16        Self {
17            left_size,
18            right_size,
19            left_graph: vec![vec![]; left_size],
20            right_graph: vec![vec![]; right_size],
21            left_match: vec![None; left_size],
22            right_match: vec![None; right_size],
23            matching_size: 0,
24        }
25    }
26    pub fn add_edge(&mut self, l: usize, r: usize) {
27        assert!(l < self.left_size);
28        assert!(r < self.right_size);
29        self.left_graph[l].push(r);
30        self.right_graph[r].push(l);
31        self.matching_size = !0;
32    }
33    pub fn from_edges(left_size: usize, right_size: usize, lr: &[(usize, usize)]) -> Self {
34        let mut left_deg = vec![0usize; left_size];
35        let mut right_deg = vec![0usize; right_size];
36        for &(l, r) in lr {
37            left_deg[l] += 1;
38            right_deg[r] += 1;
39        }
40        let mut left_graph: Vec<_> = left_deg.into_iter().map(Vec::with_capacity).collect();
41        let mut right_graph: Vec<_> = right_deg.into_iter().map(Vec::with_capacity).collect();
42        for &(l, r) in lr {
43            assert!(l < left_size);
44            assert!(r < right_size);
45            left_graph[l].push(r);
46            right_graph[r].push(l);
47        }
48        Self {
49            left_size,
50            right_size,
51            left_graph,
52            right_graph,
53            left_match: vec![None; left_size],
54            right_match: vec![None; right_size],
55            matching_size: !0,
56        }
57    }
58    pub fn hopcroft_karp(&mut self) {
59        fn bfs(bm: &BipartiteMatching, deq: &mut VecDeque<usize>, level: &mut [usize]) {
60            deq.clear();
61            for (i, r) in bm.left_match.iter().enumerate() {
62                if r.is_none() {
63                    deq.push_back(i);
64                    level[i] = 0;
65                }
66            }
67            while let Some(l) = deq.pop_front() {
68                for &r in &bm.left_graph[l] {
69                    if let Some(nl) = bm.right_match[r]
70                        && level[nl] == !0
71                    {
72                        deq.push_back(nl);
73                        level[nl] = level[l] + 1;
74                    }
75                }
76            }
77        }
78        fn dfs(
79            bm: &mut BipartiteMatching,
80            l: usize,
81            level: &mut [usize],
82            used: &mut [bool],
83        ) -> bool {
84            used[l] = true;
85            for i in 0..bm.left_graph[l].len() {
86                let r = bm.left_graph[l][i];
87                if let Some(nl) = bm.right_match[r]
88                    && (used[nl] || level[l] + 1 != level[nl] || !dfs(bm, nl, level, used))
89                {
90                    continue;
91                }
92                bm.right_match[r] = Some(l);
93                bm.left_match[l] = Some(r);
94                return true;
95            }
96            false
97        }
98        if self.matching_size != !0 {
99            return;
100        }
101        self.matching_size = self.left_match.iter().filter(|r| r.is_some()).count();
102        let mut used = vec![false; self.left_size];
103        let mut level = vec![!0; self.left_size];
104        let mut deq = VecDeque::with_capacity(self.left_size);
105        loop {
106            bfs(self, &mut deq, &mut level);
107            let prev_size = self.matching_size;
108            for l in 0..self.left_size {
109                if self.left_match[l].is_none() {
110                    self.matching_size += dfs(self, l, &mut level, &mut used) as usize;
111                }
112            }
113            if self.matching_size == prev_size {
114                break;
115            }
116            for item in level.iter_mut() {
117                *item = !0;
118            }
119            for item in used.iter_mut() {
120                *item = false;
121            }
122        }
123    }
124    pub fn kuhn_multi_start_bfs(&mut self) {
125        if self.matching_size != !0 {
126            return;
127        }
128        self.matching_size = self.left_match.iter().filter(|r| r.is_some()).count();
129        let mut deq = VecDeque::with_capacity(self.left_size);
130        let mut prev = vec![!0usize; self.left_size];
131        let mut root = vec![!0usize; self.left_size];
132        loop {
133            let mut changed = false;
134            for (l, r) in self.left_match.iter().enumerate() {
135                if r.is_none() {
136                    root[l] = l;
137                    deq.push_back(l);
138                }
139            }
140            while let Some(mut l) = deq.pop_front() {
141                if self.left_match[root[l]].is_some() {
142                    continue;
143                }
144                for mut r in self.left_graph[l].iter().cloned() {
145                    if let Some(nl) = self.right_match[r] {
146                        if prev[nl] == !0 {
147                            prev[nl] = l;
148                            root[nl] = root[l];
149                            deq.push_back(nl);
150                        }
151                    } else {
152                        loop {
153                            self.right_match[r] = Some(l);
154                            if let Some(nr) = &mut self.left_match[l] {
155                                swap(nr, &mut r);
156                                l = prev[l];
157                            } else {
158                                self.left_match[l] = Some(r);
159                                break;
160                            }
161                        }
162                        changed = true;
163                        self.matching_size += 1;
164                        break;
165                    }
166                }
167            }
168            if !changed {
169                break;
170            }
171            for item in prev.iter_mut() {
172                *item = !0;
173            }
174            for item in root.iter_mut() {
175                *item = !0;
176            }
177        }
178    }
179    pub fn push_relabel(&mut self) {
180        if self.matching_size != !0 {
181            return;
182        }
183        let size = self.left_size + self.right_size;
184        let mut level_left = vec![size; self.left_size];
185        let mut level_right = vec![size; self.right_size];
186        let mut bfs = VecDeque::with_capacity(self.left_size);
187        let mut queue: VecDeque<_> = self
188            .right_match
189            .iter()
190            .enumerate()
191            .filter_map(|(right, left)| left.is_none().then_some(right))
192            .collect();
193        let mut iteration = 0;
194        while let Some(right) = queue.pop_front() {
195            if iteration == 0 {
196                level_left.fill(size);
197                level_right.fill(size);
198                bfs.clear();
199                for (left, right) in self.left_match.iter().enumerate() {
200                    if right.is_none() {
201                        level_left[left] = 0;
202                        bfs.push_back(left);
203                    }
204                }
205                while let Some(left) = bfs.pop_front() {
206                    for &right in &self.left_graph[left] {
207                        if level_right[right] > level_left[left] + 1 {
208                            level_right[right] = level_left[left] + 1;
209                            if let Some(next_left) = self.right_match[right] {
210                                level_left[next_left] = level_right[right] + 1;
211                                bfs.push_back(next_left);
212                            }
213                        }
214                    }
215                }
216            }
217
218            let mut selected = !0;
219            let mut selected_level = size;
220            for &left in &self.right_graph[right] {
221                if level_left[left] < selected_level {
222                    selected = left;
223                    selected_level = level_left[left];
224                }
225            }
226            if selected != !0 {
227                level_right[right] = selected_level + 1;
228                if let Some(previous_right) = self.left_match[selected].take() {
229                    self.right_match[previous_right] = None;
230                    queue.push_back(previous_right);
231                }
232                self.left_match[selected] = Some(right);
233                self.right_match[right] = Some(selected);
234                level_left[selected] += 2;
235            }
236
237            iteration += 1;
238            if iteration == size {
239                iteration = 0;
240            }
241        }
242        self.matching_size = self
243            .left_match
244            .iter()
245            .filter(|right| right.is_some())
246            .count();
247    }
248    pub fn maximum_matching(&mut self) -> Vec<(usize, usize)> {
249        self.push_relabel();
250        self.left_match
251            .iter()
252            .enumerate()
253            .filter_map(|(l, r)| r.map(|r| (l, r)))
254            .collect()
255    }
256    pub fn minimum_edge_cover(&mut self) -> Vec<(usize, usize)> {
257        self.push_relabel();
258        let mut res = Vec::with_capacity(self.left_size + self.right_size - self.matching_size);
259        let mut left_used: Vec<_> = self.left_match.iter().map(Option::is_some).collect();
260        let mut right_used: Vec<_> = self.right_match.iter().map(Option::is_some).collect();
261        for (l, lg) in self.left_graph.iter().enumerate() {
262            if let Some(r) = self.left_match[l] {
263                res.push((l, r));
264            }
265            for &r in lg {
266                if !left_used[l] || !right_used[r] {
267                    left_used[l] = true;
268                    right_used[r] = true;
269                    res.push((l, r));
270                }
271            }
272        }
273        res
274    }
275    fn reachable(&mut self) -> (Vec<bool>, Vec<bool>) {
276        #[derive(Clone, Copy)]
277        enum Either {
278            Left(usize),
279            Right(usize),
280        }
281        self.push_relabel();
282        let mut left_used = vec![false; self.left_size];
283        let mut right_used = vec![false; self.right_size];
284        let mut deq = VecDeque::new();
285        for (l, r) in self.left_match.iter().enumerate() {
286            if r.is_none() {
287                left_used[l] = true;
288                deq.push_back(Either::Left(l));
289            }
290        }
291        loop {
292            match deq.pop_front() {
293                Some(Either::Left(l)) => {
294                    for &r in &self.left_graph[l] {
295                        if self.left_match[l] != Some(r) && !right_used[r] {
296                            right_used[r] = true;
297                            deq.push_back(Either::Right(r));
298                        }
299                    }
300                }
301                Some(Either::Right(r)) => {
302                    if let Some(l) = self.right_match[r]
303                        && !left_used[l]
304                    {
305                        left_used[l] = true;
306                        deq.push_back(Either::Left(l));
307                    }
308                }
309                None => break,
310            }
311        }
312        (left_used, right_used)
313    }
314    pub fn minimum_vertex_cover(&mut self) -> (Vec<usize>, Vec<usize>) {
315        let (left_used, right_used) = self.reachable();
316        (
317            left_used
318                .into_iter()
319                .enumerate()
320                .filter_map(|(l, b)| if !b { Some(l) } else { None })
321                .collect(),
322            right_used
323                .into_iter()
324                .enumerate()
325                .filter_map(|(r, b)| if b { Some(r) } else { None })
326                .collect(),
327        )
328    }
329    pub fn maximum_independent_set(&mut self) -> (Vec<usize>, Vec<usize>) {
330        let (left_used, right_used) = self.reachable();
331        (
332            left_used
333                .into_iter()
334                .enumerate()
335                .filter_map(|(l, b)| if b { Some(l) } else { None })
336                .collect(),
337            right_used
338                .into_iter()
339                .enumerate()
340                .filter_map(|(r, b)| if !b { Some(r) } else { None })
341                .collect(),
342        )
343    }
344}
345
346#[cfg(test)]
347mod tests {
348    use super::*;
349    use crate::{chmax, chmin, data_structure::UnionFind, rand, tools::Xorshift};
350
351    fn gen_graph(n: usize, m: usize, rng: &mut Xorshift) -> Vec<(usize, usize)> {
352        let mut uf = UnionFind::new(n + m);
353        let mut lr = vec![];
354        while uf.size(0) < n + m {
355            rand!(rng, l: 0..n, r: 0..m);
356            uf.unite(l, r + n);
357            lr.push((l, r));
358        }
359        lr.sort_unstable();
360        lr.dedup();
361        lr
362    }
363
364    const Q: usize = 100;
365    const N: usize = 8;
366    const M: usize = 8;
367
368    #[test]
369    fn test_maximum_matching() {
370        let mut rng = Xorshift::default();
371        for _ in 0..Q {
372            let n = rng.rand((N + 1) as u64) as usize;
373            let m = rng.rand((M + 1) as u64) as usize;
374            let mut lr = vec![];
375            for l in 0..n {
376                for r in 0..m {
377                    if rng.rand(3) == 0 {
378                        lr.push((l, r));
379                    }
380                }
381            }
382
383            let mut reachable = vec![false; 1 << m];
384            reachable[0] = true;
385            for l in 0..n {
386                let mut next = reachable.clone();
387                for (bits, &can_match) in reachable.iter().enumerate() {
388                    if can_match {
389                        for &(_, r) in lr.iter().filter(|&&(left, _)| left == l) {
390                            next[bits | 1 << r] = true;
391                        }
392                    }
393                }
394                reachable = next;
395            }
396            let expected = reachable
397                .iter()
398                .enumerate()
399                .filter_map(|(bits, &reachable)| reachable.then_some(bits.count_ones() as usize))
400                .max()
401                .unwrap();
402
403            for incremental in [false, true] {
404                let mut bm = if incremental {
405                    let mut bm = BipartiteMatching::new(n, m);
406                    for &(l, r) in &lr[..lr.len() / 2] {
407                        bm.add_edge(l, r);
408                    }
409                    bm.maximum_matching();
410                    for &(l, r) in &lr[lr.len() / 2..] {
411                        bm.add_edge(l, r);
412                    }
413                    bm
414                } else {
415                    BipartiteMatching::from_edges(n, m, &lr)
416                };
417                let matching = bm.maximum_matching();
418                assert_eq!(matching.len(), expected);
419                let mut left_used = vec![false; n];
420                let mut right_used = vec![false; m];
421                for (l, r) in matching {
422                    assert!(lr.contains(&(l, r)));
423                    assert!(!left_used[l]);
424                    assert!(!right_used[r]);
425                    left_used[l] = true;
426                    right_used[r] = true;
427                }
428            }
429        }
430    }
431
432    #[test]
433    fn test_minimum_edge_cover() {
434        let mut rng = Xorshift::default();
435        for _ in 0..Q {
436            rand!(rng, n: 4..=N, m: 4..=M);
437            let lr = gen_graph(n, m, &mut rng);
438            let mut dp = vec![vec![!0usize; 1 << m]; 1 << n];
439            dp[0][0] = 0;
440            for bitl in 0usize..1 << n {
441                for bitr in 0usize..1 << m {
442                    if dp[bitl][bitr] == !0 {
443                        continue;
444                    }
445                    for &(l, r) in &lr {
446                        chmin!(dp[bitl | (1 << l)][bitr | (1 << r)], dp[bitl][bitr] + 1);
447                    }
448                }
449            }
450            let cover = BipartiteMatching::from_edges(n, m, &lr).minimum_edge_cover();
451            assert_eq!(*dp.last().unwrap().last().unwrap(), cover.len());
452            let mut left_used = vec![false; n];
453            let mut right_used = vec![false; m];
454            for (l, r) in cover {
455                left_used[l] = true;
456                right_used[r] = true;
457            }
458            assert!(left_used.iter().all(|&b| b));
459            assert!(right_used.iter().all(|&b| b));
460
461            let mut bm = BipartiteMatching::new(n, m);
462            for &(l, r) in &lr {
463                bm.add_edge(l, r);
464            }
465            bm.hopcroft_karp();
466            let cover = bm.minimum_edge_cover();
467            assert_eq!(*dp.last().unwrap().last().unwrap(), cover.len());
468            let mut left_used = vec![false; n];
469            let mut right_used = vec![false; m];
470            for (l, r) in cover {
471                left_used[l] = true;
472                right_used[r] = true;
473            }
474            assert!(left_used.iter().all(|&b| b));
475            assert!(right_used.iter().all(|&b| b));
476        }
477    }
478
479    #[test]
480    fn test_minimum_vertex_cover() {
481        let mut rng = Xorshift::default();
482        for _ in 0..Q {
483            rand!(rng, n: 4..=N, m: 4..=M);
484            let lr = gen_graph(n, m, &mut rng);
485            let mut ans = !0usize;
486            for bitl in 0usize..1 << n {
487                for bitr in 0usize..1 << m {
488                    if lr
489                        .iter()
490                        .all(|&(l, r)| bitl & (1 << l) != 0 || bitr & (1 << r) != 0)
491                    {
492                        chmin!(ans, (bitl.count_ones() + bitr.count_ones()) as usize);
493                    }
494                }
495            }
496            let set = BipartiteMatching::from_edges(n, m, &lr).minimum_vertex_cover();
497            assert_eq!(ans, set.0.len() + set.1.len());
498            for &(l, r) in &lr {
499                assert!(set.0.contains(&l) || set.1.contains(&r));
500            }
501
502            let mut bm = BipartiteMatching::new(n, m);
503            for &(l, r) in &lr {
504                bm.add_edge(l, r);
505            }
506            bm.hopcroft_karp();
507            let set = bm.minimum_vertex_cover();
508            assert_eq!(ans, set.0.len() + set.1.len());
509            for &(l, r) in &lr {
510                assert!(set.0.contains(&l) || set.1.contains(&r));
511            }
512        }
513    }
514
515    #[test]
516    fn test_maximum_independent_set() {
517        let mut rng = Xorshift::default();
518        for _ in 0..Q {
519            rand!(rng, n: 4..=N, m: 4..=M);
520            let lr = gen_graph(n, m, &mut rng);
521            let mut ans = 0usize;
522            for bitl in 0usize..1 << n {
523                for bitr in 0usize..1 << m {
524                    if lr
525                        .iter()
526                        .all(|&(l, r)| bitl & (1 << l) == 0 || bitr & (1 << r) == 0)
527                    {
528                        chmax!(ans, (bitl.count_ones() + bitr.count_ones()) as usize);
529                    }
530                }
531            }
532            let set = BipartiteMatching::from_edges(n, m, &lr).maximum_independent_set();
533            assert_eq!(ans, set.0.len() + set.1.len());
534            for &(l, r) in &lr {
535                assert!(!set.0.contains(&l) || !set.1.contains(&r));
536            }
537
538            let mut bm = BipartiteMatching::new(n, m);
539            for &(l, r) in &lr {
540                bm.add_edge(l, r);
541            }
542            bm.hopcroft_karp();
543            let set = bm.maximum_independent_set();
544            assert_eq!(ans, set.0.len() + set.1.len());
545            for &(l, r) in &lr {
546                assert!(!set.0.contains(&l) || !set.1.contains(&r));
547            }
548        }
549    }
550}