use std::cmp::Ordering;
use std::collections::BinaryHeap;
const KNN: usize = 8;
const SUPPORT_REL: f64 = 1e-3;
#[derive(Debug, Clone)]
pub struct RecoveredMatching {
pub assignment: Vec<u32>,
pub total_cost: f64,
pub support_edges: usize,
}
#[inline]
fn dist(a: [f64; 2], b: [f64; 2]) -> f64 {
((a[0] - b[0]).powi(2) + (a[1] - b[1]).powi(2)).sqrt()
}
#[inline]
fn orient(a: [f64; 2], b: [f64; 2], c: [f64; 2]) -> f64 {
(b[0] - a[0]) * (c[1] - a[1]) - (b[1] - a[1]) * (c[0] - a[0])
}
pub fn segments_cross(a: [f64; 2], b: [f64; 2], c: [f64; 2], d: [f64; 2]) -> bool {
let d1 = orient(a, b, c);
let d2 = orient(a, b, d);
let d3 = orient(c, d, a);
let d4 = orient(c, d, b);
(d1 * d2 < 0.0) && (d3 * d4 < 0.0)
}
pub const FLOOR_MAX_SWEEPS: usize = 100;
#[derive(Debug, Clone, Copy)]
pub struct CertifiedFloor {
pub value: f64,
pub sweeps: usize,
}
pub fn certified_floor_ascent(
y: &[f64],
riders: &[[f64; 2]],
cabs: &[[f64; 2]],
rider_mass: f64,
cab_cap: f64,
max_sweeps: usize,
) -> CertifiedFloor {
let ns = riders.len();
let nt = cabs.len();
assert_eq!(y.len(), ns + nt, "dual length must be ns + nt");
let mut dmat = vec![0.0f64; ns * nt];
for (i, r) in riders.iter().enumerate() {
let row = &mut dmat[i * nt..(i + 1) * nt];
for (j, k) in cabs.iter().enumerate() {
row[j] = dist(*r, *k);
}
}
let update_u = |dmat: &[f64], v: &[f64], u: &mut [f64]| {
for (i, u_i) in u.iter_mut().enumerate() {
let row = &dmat[i * nt..(i + 1) * nt];
let mut m = f64::INFINITY;
for j in 0..nt {
let cand = row[j] - v[j];
if cand < m {
m = cand;
}
}
*u_i = m;
}
};
let update_v = |dmat: &[f64], u: &[f64], v: &mut [f64]| {
for vj in v.iter_mut() {
*vj = f64::INFINITY;
}
for (i, &u_i) in u.iter().enumerate() {
let row = &dmat[i * nt..(i + 1) * nt];
for (j, vj) in v.iter_mut().enumerate() {
let cand = row[j] - u_i;
if cand < *vj {
*vj = cand;
}
}
}
for vj in v.iter_mut() {
if *vj > 0.0 {
*vj = 0.0;
}
}
};
let nt_f = nt as f64;
let coord_floor = |u: &[f64], v: &[f64]| -> f64 {
(rider_mass * u.iter().sum::<f64>() + cab_cap * v.iter().sum::<f64>()) * nt_f
};
let mut v: Vec<f64> = (0..nt).map(|j| (-y[ns + j]).min(0.0)).collect();
let mut u = vec![0.0f64; ns];
update_u(&dmat, &v, &mut u);
let mut floor = coord_floor(&u, &v);
let mut sweeps = 0;
for _ in 0..max_sweeps {
update_v(&dmat, &u, &mut v);
update_u(&dmat, &v, &mut u);
sweeps += 1;
let new_floor = coord_floor(&u, &v);
let gained = new_floor - floor;
floor = new_floor;
if gained < 1e-9 * (1.0 + new_floor.abs()) {
break;
}
}
CertifiedFloor {
value: floor,
sweeps,
}
}
pub fn certified_floor(
y: &[f64],
riders: &[[f64; 2]],
cabs: &[[f64; 2]],
rider_mass: f64,
cab_cap: f64,
) -> f64 {
certified_floor_ascent(y, riders, cabs, rider_mass, cab_cap, FLOOR_MAX_SWEEPS).value
}
#[derive(PartialEq)]
struct Key(f64);
impl Eq for Key {}
impl PartialOrd for Key {
fn partial_cmp(&self, o: &Self) -> Option<Ordering> {
Some(self.cmp(o))
}
}
impl Ord for Key {
fn cmp(&self, o: &Self) -> Ordering {
self.0.total_cmp(&o.0)
}
}
struct Mcmf {
to: Vec<usize>,
cap: Vec<i64>,
cost: Vec<f64>,
adj: Vec<Vec<usize>>,
}
impl Mcmf {
fn new(nodes: usize) -> Self {
Mcmf {
to: Vec::new(),
cap: Vec::new(),
cost: Vec::new(),
adj: vec![Vec::new(); nodes],
}
}
fn add(&mut self, u: usize, v: usize, cap: i64, cost: f64) -> usize {
let e = self.to.len();
self.to.push(v);
self.cap.push(cap);
self.cost.push(cost);
self.adj[u].push(e);
self.to.push(u);
self.cap.push(0);
self.cost.push(-cost);
self.adj[v].push(e + 1);
e
}
}
fn min_cost_matching(riders: &[[f64; 2]], cabs: &[[f64; 2]], cand: &[Vec<usize>]) -> Vec<u32> {
let ns = riders.len();
let nt = cabs.len();
let s = 0;
let t = ns + nt + 1;
let nodes = ns + nt + 2;
let mut g = Mcmf::new(nodes);
for i in 0..ns {
g.add(s, 1 + i, 1, 0.0);
}
for j in 0..nt {
g.add(1 + ns + j, t, 1, 0.0);
}
let mut rider_edges: Vec<Vec<(usize, usize)>> = vec![Vec::new(); ns]; for i in 0..ns {
for &j in &cand[i] {
let e = g.add(1 + i, 1 + ns + j, 1, dist(riders[i], cabs[j]));
rider_edges[i].push((j, e));
}
}
let mut h = vec![0.0f64; nodes]; let mut d = vec![f64::INFINITY; nodes];
let mut prev_edge = vec![usize::MAX; nodes];
for _ in 0..ns {
for x in d.iter_mut() {
*x = f64::INFINITY;
}
d[s] = 0.0;
let mut heap: BinaryHeap<std::cmp::Reverse<(Key, usize)>> = BinaryHeap::new();
heap.push(std::cmp::Reverse((Key(0.0), s)));
while let Some(std::cmp::Reverse((Key(du), u))) = heap.pop() {
if du > d[u] {
continue;
}
for &e in &g.adj[u] {
if g.cap[e] <= 0 {
continue;
}
let v = g.to[e];
let nd = du + g.cost[e] + h[u] - h[v];
if nd + 1e-15 < d[v] {
d[v] = nd;
prev_edge[v] = e;
heap.push(std::cmp::Reverse((Key(nd), v)));
}
}
}
assert!(
d[t].is_finite(),
"candidate graph has no rider-perfect matching (should be impossible: kNN ∪ support)"
);
for node in 0..nodes {
if d[node].is_finite() {
h[node] += d[node];
}
}
let mut node = t;
while node != s {
let e = prev_edge[node];
g.cap[e] -= 1;
g.cap[e ^ 1] += 1;
node = g.to[e ^ 1]; }
}
let mut assignment = vec![u32::MAX; ns];
for i in 0..ns {
for &(j, e) in &rider_edges[i] {
if g.cap[e] == 0 {
assignment[i] = j as u32;
break;
}
}
assert_ne!(assignment[i], u32::MAX, "rider {i} left unmatched");
}
assignment
}
fn uncross(assignment: &mut [u32], riders: &[[f64; 2]], cabs: &[[f64; 2]]) {
let n = assignment.len();
let max_passes = 8 * n + 64;
let mut converged = false;
for _ in 0..max_passes {
let mut swapped = false;
for a in 0..n {
for b in (a + 1)..n {
let ra = riders[a];
let rb = riders[b];
let pa = cabs[assignment[a] as usize];
let pb = cabs[assignment[b] as usize];
if segments_cross(ra, pa, rb, pb) {
assignment.swap(a, b);
swapped = true;
}
}
}
if !swapped {
converged = true;
break;
}
}
assert!(
converged,
"uncrossing sweep did not converge within {max_passes} passes"
);
for a in 0..n {
for b in (a + 1)..n {
let pa = cabs[assignment[a] as usize];
let pb = cabs[assignment[b] as usize];
debug_assert!(
!segments_cross(riders[a], pa, riders[b], pb),
"crossing survived the sweep"
);
}
}
}
pub fn recover_matching(x: &[f64], riders: &[[f64; 2]], cabs: &[[f64; 2]]) -> RecoveredMatching {
let ns = riders.len();
let nt = cabs.len();
assert_eq!(x.len(), ns * nt, "plan length must be ns·nt");
assert!(nt >= ns, "need cabs ≥ riders");
let mut support_edges = 0usize;
let mut cand: Vec<Vec<usize>> = Vec::with_capacity(ns);
for i in 0..ns {
let row = &x[i * nt..(i + 1) * nt];
let row_mass: f64 = row.iter().copied().map(|v| v.max(0.0)).sum();
let thresh = SUPPORT_REL * row_mass;
let mut in_set = vec![false; nt];
let mut edges: Vec<usize> = Vec::new();
for (j, &v) in row.iter().enumerate() {
if v > thresh && v > 0.0 {
support_edges += 1;
if !in_set[j] {
in_set[j] = true;
edges.push(j);
}
}
}
let mut order: Vec<usize> = (0..nt).collect();
order.sort_by(|&p, &q| {
let dp = dist(riders[i], cabs[p]);
let dq = dist(riders[i], cabs[q]);
dp.total_cmp(&dq).then(p.cmp(&q))
});
for &j in order.iter().take(KNN.min(nt)) {
if !in_set[j] {
in_set[j] = true;
edges.push(j);
}
}
edges.sort_unstable(); cand.push(edges);
}
let mut assignment = min_cost_matching(riders, cabs, &cand);
uncross(&mut assignment, riders, cabs);
let total_cost: f64 = assignment
.iter()
.enumerate()
.map(|(i, &j)| dist(riders[i], cabs[j as usize]))
.sum();
RecoveredMatching {
assignment,
total_cost,
support_edges,
}
}