const NONE: usize = usize::MAX;
struct Solver<'a> {
cost: &'a [f64],
nr: usize,
nc: usize,
u: Vec<f64>,
v: Vec<f64>,
shortest: Vec<f64>,
path: Vec<usize>,
col4row: Vec<usize>,
row4col: Vec<usize>,
sr: Vec<bool>,
sc: Vec<bool>,
remaining: Vec<usize>,
}
impl<'a> Solver<'a> {
fn new(cost: &'a [f64], nr: usize, nc: usize) -> Self {
Solver {
cost,
nr,
nc,
u: vec![0.0; nr],
v: vec![0.0; nc],
shortest: vec![f64::INFINITY; nc],
path: vec![NONE; nc],
col4row: vec![NONE; nr],
row4col: vec![NONE; nc],
sr: vec![false; nr],
sc: vec![false; nc],
remaining: vec![0; nc],
}
}
fn augmenting_path(&mut self, start: usize) -> (usize, f64) {
let nc = self.nc;
let mut min_val = 0.0f64;
let mut num_remaining = nc;
for it in 0..nc {
self.remaining[it] = nc - it - 1;
}
self.sr.fill(false);
self.sc.fill(false);
self.shortest.fill(f64::INFINITY);
let mut i = start;
let mut sink = NONE;
while sink == NONE {
let mut index = NONE;
let mut lowest = f64::INFINITY;
self.sr[i] = true;
for it in 0..num_remaining {
let j = self.remaining[it];
let r = min_val + self.cost[i * nc + j] - self.u[i] - self.v[j];
if r < self.shortest[j] {
self.path[j] = i;
self.shortest[j] = r;
}
if self.shortest[j] < lowest
|| (self.shortest[j] == lowest && self.row4col[j] == NONE)
{
lowest = self.shortest[j];
index = it;
}
}
min_val = lowest;
if min_val.is_infinite() {
return (NONE, min_val);
}
let j = self.remaining[index];
if self.row4col[j] == NONE {
sink = j;
} else {
i = self.row4col[j];
}
self.sc[j] = true;
num_remaining -= 1;
self.remaining[index] = self.remaining[num_remaining];
self.remaining[num_remaining] = j;
}
(sink, min_val)
}
fn solve(&mut self) {
for cur_row in 0..self.nr {
let (sink, min_val) = self.augmenting_path(cur_row);
assert!(
sink != NONE,
"lsap: cost matrix is infeasible — it likely contains NaN or infinite entries"
);
self.u[cur_row] += min_val;
for i in 0..self.nr {
if self.sr[i] && i != cur_row {
self.u[i] += min_val - self.shortest[self.col4row[i]];
}
}
for j in 0..self.nc {
if self.sc[j] {
self.v[j] -= min_val - self.shortest[j];
}
}
let mut j = sink;
loop {
let i = self.path[j];
self.row4col[j] = i;
std::mem::swap(&mut self.col4row[i], &mut j);
if i == cur_row {
break;
}
}
}
}
}
pub fn lsap(cost: &[f64], nr: usize, nc: usize, maximize: bool) -> (Vec<usize>, Vec<usize>) {
if nr == 0 || nc == 0 {
return (Vec::new(), Vec::new());
}
assert_eq!(
cost.len(),
nr * nc,
"lsap: cost must be nr*nc row-major ({nr}x{nc})"
);
let transpose = nc < nr;
let (rn, cn) = if transpose { (nc, nr) } else { (nr, nc) };
let owned: Option<Vec<f64>> = if transpose {
let mut c = vec![0.0f64; rn * cn];
for i in 0..nr {
for j in 0..nc {
let v = cost[i * nc + j];
c[j * nr + i] = if maximize { -v } else { v };
}
}
Some(c)
} else if maximize {
Some(cost.iter().map(|x| -x).collect())
} else {
None
};
let c: &[f64] = owned.as_deref().unwrap_or(cost);
let mut solver = Solver::new(c, rn, cn);
solver.solve();
let col4row = solver.col4row;
let mut row_ind = Vec::with_capacity(rn);
let mut col_ind = Vec::with_capacity(rn);
if transpose {
let mut order: Vec<usize> = (0..rn).collect();
order.sort_by_key(|&k| col4row[k]);
for k in order {
row_ind.push(col4row[k]);
col_ind.push(k);
}
} else {
for (i, &c4r) in col4row.iter().enumerate() {
row_ind.push(i);
col_ind.push(c4r);
}
}
(row_ind, col_ind)
}
#[cfg(test)]
mod tests {
use super::*;
use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
fn cost_of(cost: &[f64], nc: usize, r: &[usize], c: &[usize]) -> f64 {
r.iter().zip(c).map(|(&i, &j)| cost[i * nc + j]).sum()
}
fn brute_min(cost: &[f64], nr: usize, nc: usize) -> f64 {
if nr > nc {
let mut t = vec![0.0; nr * nc];
for i in 0..nr {
for j in 0..nc {
t[j * nr + i] = cost[i * nc + j];
}
}
return brute_min(&t, nc, nr);
}
fn rec(cost: &[f64], nr: usize, nc: usize, row: usize, used: &mut Vec<bool>) -> f64 {
if row == nr {
return 0.0;
}
let mut best = f64::INFINITY;
for j in 0..nc {
if !used[j] {
used[j] = true;
let sub = cost[row * nc + j] + rec(cost, nr, nc, row + 1, used);
used[j] = false;
best = best.min(sub);
}
}
best
}
rec(cost, nr, nc, 0, &mut vec![false; nc])
}
#[test]
fn known_square_assignment() {
let cost = [4.0, 1.0, 3.0, 2.0, 0.0, 5.0, 3.0, 2.0, 2.0];
let (r, c) = lsap(&cost, 3, 3, false);
assert_eq!(r, vec![0, 1, 2]);
assert!((cost_of(&cost, 3, &r, &c) - brute_min(&cost, 3, 3)).abs() < 1e-12);
}
#[test]
fn wide_matrix_assigns_all_rows() {
let cost = [1.0, 2.0, 3.0, 4.0, 1.0, 5.0];
let (r, c) = lsap(&cost, 2, 3, false);
assert_eq!(r, vec![0, 1]);
assert_eq!(c.len(), 2);
assert!((cost_of(&cost, 3, &r, &c) - brute_min(&cost, 2, 3)).abs() < 1e-12);
}
#[test]
fn tall_matrix_triggers_transpose_and_keeps_row_ind_sorted() {
let cost = [1.0, 4.0, 2.0, 1.0, 5.0, 3.0];
let (r, c) = lsap(&cost, 3, 2, false);
assert_eq!(r.len(), 2);
assert!(r.windows(2).all(|w| w[0] < w[1]), "row_ind ascending");
assert!(r.iter().all(|&i| i < 3) && c.iter().all(|&j| j < 2));
assert!((cost_of(&cost, 2, &r, &c) - brute_min(&cost, 3, 2)).abs() < 1e-12);
}
#[test]
fn maximize_picks_largest() {
let cost = [1.0, 2.0, 3.0, 4.0];
let (r, c) = lsap(&cost, 2, 2, true);
assert_eq!((r, c), (vec![0, 1], vec![1, 0]));
}
#[test]
fn trivial_sizes() {
assert_eq!(lsap(&[], 0, 0, false), (vec![], vec![]));
assert_eq!(lsap(&[], 0, 3, false), (vec![], vec![]));
assert_eq!(lsap(&[7.0], 1, 1, false), (vec![0], vec![0]));
}
#[test]
fn matches_scipy_assignment_vectors() {
let data = include_str!("testdata/lsap_scipy.json");
let cases: serde_json::Value = serde_json::from_str(data).expect("parse fixture");
let u = |x: &serde_json::Value| x.as_u64().expect("u64") as usize;
let uv = |x: &serde_json::Value| {
x.as_array()
.expect("array")
.iter()
.map(u)
.collect::<Vec<_>>()
};
for (idx, case) in cases.as_array().expect("array").iter().enumerate() {
let (nr, nc) = (u(&case["nr"]), u(&case["nc"]));
let maximize = case["maximize"].as_bool().expect("bool");
let cost: Vec<f64> = case["cost"]
.as_array()
.expect("array")
.iter()
.map(|x| x.as_f64().expect("f64"))
.collect();
let (r, c) = lsap(&cost, nr, nc, maximize);
let (exp_r, exp_c) = (uv(&case["row_ind"]), uv(&case["col_ind"]));
assert_eq!(r, exp_r, "row_ind case {idx} ({nr}x{nc} max={maximize})");
assert_eq!(c, exp_c, "col_ind case {idx} ({nr}x{nc} max={maximize})");
}
}
#[test]
fn brute_force_optimality_random() {
let mut rng = StdRng::seed_from_u64(0xC0C0);
for _ in 0..2000 {
let nr = rng.random_range(1..=5);
let nc = rng.random_range(1..=5);
let integer = rng.random_bool(0.5);
let cost: Vec<f64> = (0..nr * nc)
.map(|_| {
if integer {
rng.random_range(0..4) as f64
} else {
rng.random_range(0.0..10.0)
}
})
.collect();
let maximize = rng.random_bool(0.5);
let (r, c) = lsap(&cost, nr, nc, maximize);
assert_eq!(r.len(), nr.min(nc));
let mut rows = r.clone();
rows.sort_unstable();
rows.dedup();
assert_eq!(rows.len(), r.len());
let mut cols = c.clone();
cols.sort_unstable();
cols.dedup();
assert_eq!(cols.len(), c.len());
let got = cost_of(&cost, nc, &r, &c);
if maximize {
let neg: Vec<f64> = cost.iter().map(|x| -x).collect();
let opt = -brute_min(&neg, nr, nc);
assert!((got - opt).abs() < 1e-9, "maximize not optimal");
} else {
let opt = brute_min(&cost, nr, nc);
assert!((got - opt).abs() < 1e-9, "minimize not optimal");
}
}
}
#[test]
#[should_panic(expected = "infeasible")]
fn all_non_finite_row_reports_infeasible_not_index_oob() {
lsap(&[1.0, 2.0, f64::NAN, f64::NAN], 2, 2, false);
}
#[test]
#[should_panic(expected = "row-major")]
fn wrong_length_cost_is_rejected() {
lsap(&[1.0, 2.0, 3.0], 2, 2, false);
}
}