Skip to main content

competitive/graph/
assignment.rs

1/// Returns the minimum cost and a column assigned to each row of a square cost matrix.
2pub fn minimum_assignment<T>(cost: &[Vec<T>]) -> (i64, Vec<usize>)
3where
4    T: Copy + Into<i64>,
5{
6    let n = cost.len();
7    assert!(cost.iter().all(|row| row.len() == n));
8    if n == 0 {
9        return (0, Vec::new());
10    }
11    let value = |row: usize, column: usize| cost[row][column].into();
12
13    let mut monge = true;
14    let mut anti_monge = true;
15    'outer: for row in 0..n - 1 {
16        for column in 0..n - 1 {
17            let straight = i128::from(value(row, column)) + i128::from(value(row + 1, column + 1));
18            let crossed = i128::from(value(row, column + 1)) + i128::from(value(row + 1, column));
19            monge &= straight <= crossed;
20            anti_monge &= straight >= crossed;
21            if !monge && !anti_monge {
22                break 'outer;
23            }
24        }
25    }
26    if monge || anti_monge {
27        let assignment: Vec<_> = if monge {
28            (0..n).collect()
29        } else {
30            (0..n).rev().collect()
31        };
32        let total = (0..n).map(|row| value(row, assignment[row])).sum();
33        return (total, assignment);
34    }
35
36    let none = !0;
37    let mut row_mate = vec![none; n];
38    let mut column_mate = vec![none; n];
39    let mut potential = vec![0; n];
40    let mut transferable = vec![false; n];
41    for column in 0..n {
42        let mut row = 0;
43        for next in 1..n {
44            if value(next, column) < value(row, column) {
45                row = next;
46            }
47        }
48        potential[column] = value(row, column);
49        if row_mate[row] == none {
50            row_mate[row] = column;
51            column_mate[column] = row;
52            transferable[row] = true;
53        } else {
54            transferable[row] = false;
55        }
56    }
57    for row in 0..n {
58        if transferable[row] {
59            let column = row_mate[row];
60            let mut best = i64::MAX;
61            for (next, &next_potential) in potential.iter().enumerate() {
62                if next != column {
63                    best = best.min(value(row, next) - next_potential);
64                }
65            }
66            if best != i64::MAX {
67                potential[column] -= best;
68            }
69        }
70    }
71    for _ in 0..2 {
72        for row in 0..n {
73            if row_mate[row] != none {
74                continue;
75            }
76            let mut best = value(row, 0) - potential[0];
77            let mut second = i64::MAX;
78            let mut column = 0;
79            for (next, &next_potential) in potential.iter().enumerate().skip(1) {
80                let reduced = value(row, next) - next_potential;
81                if reduced < best || reduced == best && column_mate[column] != none {
82                    second = best;
83                    best = reduced;
84                    column = next;
85                } else {
86                    second = second.min(reduced);
87                }
88            }
89            if best < second {
90                potential[column] -= second - best;
91            }
92            let replaced = column_mate[column];
93            if replaced != none {
94                row_mate[replaced] = none;
95            }
96            row_mate[row] = column;
97            column_mate[column] = row;
98        }
99    }
100
101    let mut columns: Vec<_> = (0..n).collect();
102    let mut distance = vec![0; n];
103    let mut predecessor = vec![0; n];
104    for start in 0..n {
105        if row_mate[start] != none {
106            continue;
107        }
108        for column in 0..n {
109            distance[column] = value(start, column) - potential[column];
110            predecessor[column] = start;
111        }
112        let mut scanned = 0;
113        let mut labeled = 0;
114        let mut last = 0;
115        let free_column = loop {
116            if scanned == labeled {
117                last = scanned;
118                let mut best = distance[columns[scanned]];
119                for next in scanned..n {
120                    let column = columns[next];
121                    if distance[column] <= best {
122                        if distance[column] < best {
123                            best = distance[column];
124                            labeled = scanned;
125                        }
126                        columns.swap(next, labeled);
127                        labeled += 1;
128                    }
129                }
130                if let Some(column) = columns[scanned..labeled]
131                    .iter()
132                    .copied()
133                    .find(|&column| column_mate[column] == none)
134                {
135                    break column;
136                }
137            }
138            let column = columns[scanned];
139            scanned += 1;
140            let row = column_mate[column];
141            let base = value(row, column) - potential[column];
142            let mut next = labeled;
143            let mut free_column = none;
144            while next < n {
145                let other = columns[next];
146                let edge = value(row, other) - potential[other] - base;
147                let candidate = distance[column] + edge;
148                if candidate < distance[other] {
149                    distance[other] = candidate;
150                    predecessor[other] = row;
151                    if edge == 0 {
152                        if column_mate[other] == none {
153                            free_column = other;
154                            break;
155                        }
156                        columns.swap(next, labeled);
157                        labeled += 1;
158                    }
159                }
160                next += 1;
161            }
162            if free_column != none {
163                break free_column;
164            }
165        };
166        for &column in &columns[..last] {
167            potential[column] += distance[column] - distance[free_column];
168        }
169        let mut column = free_column;
170        loop {
171            let row = predecessor[column];
172            column_mate[column] = row;
173            let next = row_mate[row];
174            row_mate[row] = column;
175            if next == none {
176                break;
177            }
178            column = next;
179        }
180    }
181
182    let total = (0..n).map(|row| value(row, row_mate[row])).sum();
183    (total, row_mate)
184}
185
186#[cfg(test)]
187mod tests {
188    use crate::{algorithm::SliceCombinationsExt, graph::minimum_assignment, tools::Xorshift};
189
190    #[test]
191    fn test_minimum_assignment() {
192        let mut rng = Xorshift::default();
193        for case in 0..100 {
194            let n = rng.random(0..=8);
195            let mut cost: Vec<Vec<i64>> = (0..n)
196                .map(|_| rng.random_iter(-100..=100).take(n).collect())
197                .collect();
198            if case % 3 != 0 {
199                let sign = if case % 3 == 1 { 1 } else { -1 };
200                let row_bias: Vec<i64> = rng.random_iter(-100..=100).take(n).collect();
201                let column_bias: Vec<i64> = rng.random_iter(-100..=100).take(n).collect();
202                for row in 0..n {
203                    for column in 0..n {
204                        cost[row][column] = row_bias[row]
205                            + column_bias[column]
206                            + sign * (row as i64 - column as i64).pow(2);
207                    }
208                }
209            }
210            let mut permutation: Vec<_> = (0..n).collect();
211            let mut expected = i64::MAX;
212            loop {
213                expected = expected.min(
214                    cost.iter()
215                        .zip(&permutation)
216                        .map(|(row, &column)| row[column])
217                        .sum(),
218                );
219                if !permutation.next_permutation() {
220                    break;
221                }
222            }
223            let (actual, assignment) = minimum_assignment(&cost);
224            assert_eq!(expected, actual);
225            assert_eq!(
226                actual,
227                cost.iter()
228                    .zip(&assignment)
229                    .map(|(row, &column)| row[column])
230                    .sum()
231            );
232            let mut sorted = assignment;
233            sorted.sort_unstable();
234            assert!(sorted.into_iter().eq(0..n));
235        }
236    }
237}