use std::collections::HashSet;
pub struct Branching {
pub parent: Vec<Option<usize>>,
pub tree: Vec<usize>,
pub roots: Vec<usize>,
}
struct Arc {
u: usize,
v: usize,
w: f64,
orig: usize,
landed: usize,
}
pub fn max_branching(n: usize, arcs: &[(usize, usize, f32)], root_affinity: &[f32]) -> Branching {
assert_eq!(root_affinity.len(), n, "root_affinity must have length n");
if n == 0 {
return Branching {
parent: Vec::new(),
tree: Vec::new(),
roots: Vec::new(),
};
}
let sroot = n;
let mut all: Vec<Arc> = Vec::with_capacity(arcs.len() + n);
for (i, &(u, v, w)) in arcs.iter().enumerate() {
if u == v {
continue; }
all.push(Arc {
u,
v,
w: w as f64,
orig: i,
landed: v,
});
}
let sroot_orig_base = arcs.len();
for (v, &aff) in root_affinity.iter().enumerate() {
all.push(Arc {
u: sroot,
v,
w: aff as f64,
orig: sroot_orig_base + v,
landed: v,
});
}
let used = solve(n + 1, sroot, &all);
let mut parent: Vec<Option<usize>> = vec![None; n];
for &orig in &used {
let (u, v) = if orig < sroot_orig_base {
(arcs[orig].0, arcs[orig].1)
} else {
(sroot, orig - sroot_orig_base)
};
if v < n {
parent[v] = if u == sroot { None } else { Some(u) };
}
}
let roots: Vec<usize> = (0..n).filter(|&v| parent[v].is_none()).collect();
let mut root_of_comp = vec![usize::MAX; n];
for (c, &r) in roots.iter().enumerate() {
root_of_comp[r] = c;
}
let mut tree = vec![usize::MAX; n];
for v in 0..n {
let mut path = Vec::new();
let mut x = v;
while tree[x] == usize::MAX && root_of_comp[x] == usize::MAX {
path.push(x);
match parent[x] {
Some(p) => x = p,
None => break,
}
}
let comp = if tree[x] != usize::MAX {
tree[x]
} else {
root_of_comp[x]
};
for &p in &path {
tree[p] = comp;
}
tree[v] = comp;
}
Branching {
parent,
tree,
roots,
}
}
fn solve(n: usize, root: usize, arcs: &[Arc]) -> Vec<usize> {
let mut best: Vec<Option<usize>> = vec![None; n];
for (i, a) in arcs.iter().enumerate() {
if a.v == root {
continue;
}
match best[a.v] {
None => best[a.v] = Some(i),
Some(j) if a.w > arcs[j].w => best[a.v] = Some(i),
_ => {}
}
}
let par = |v: usize| best[v].map(|i| arcs[i].u);
let mut color = vec![0u8; n]; let mut cycle: Option<Vec<usize>> = None;
for s in 0..n {
if color[s] != 0 || s == root {
continue;
}
let mut stack: Vec<usize> = Vec::new();
let mut v = s;
loop {
if v == root || color[v] == 2 {
break;
}
if color[v] == 1 {
let start = stack.iter().position(|&x| x == v).unwrap();
cycle = Some(stack[start..].to_vec());
break;
}
color[v] = 1;
stack.push(v);
match par(v) {
Some(p) => v = p,
None => break,
}
}
for &x in &stack {
color[x] = 2;
}
if cycle.is_some() {
break;
}
}
let Some(cyc) = cycle else {
let mut used = Vec::new();
for (v, slot) in best.iter().enumerate() {
if v == root {
continue;
}
if let Some(i) = slot {
used.push(arcs[*i].orig);
}
}
return used;
};
let in_cycle: HashSet<usize> = cyc.iter().copied().collect();
let mut map = vec![usize::MAX; n];
let mut next = 0;
for (v, m) in map.iter_mut().enumerate() {
if !in_cycle.contains(&v) {
*m = next;
next += 1;
}
}
let cnode = next;
next += 1;
for &v in &cyc {
map[v] = cnode;
}
let new_n = next;
let mut new_arcs: Vec<Arc> = Vec::with_capacity(arcs.len());
for a in arcs {
let nu = map[a.u];
let nv = map[a.v];
if nu == nv {
continue; }
let w = if in_cycle.contains(&a.v) {
a.w - arcs[best[a.v].unwrap()].w
} else {
a.w
};
new_arcs.push(Arc {
u: nu,
v: nv,
w,
orig: a.orig,
landed: a.v,
});
}
let sub_used = solve(new_n, map[root], &new_arcs);
let sub_set: HashSet<usize> = sub_used.iter().copied().collect();
let mut v_enter = None;
for a in &new_arcs {
if a.v == cnode && sub_set.contains(&a.orig) {
v_enter = Some(a.landed);
break;
}
}
let v_enter = v_enter.expect("contracted cycle must have an entering arc");
let mut used = sub_used;
for &x in &cyc {
if x != v_enter {
used.push(arcs[best[x].unwrap()].orig);
}
}
used
}
#[cfg(test)]
mod tests;