use super::model::Hypergraph;
use crate::partition::common::{FmBalance, GainBuckets, Stall, commit_best_prefix, fm_balance};
pub(super) fn fm_refine_pass(hg: &Hypergraph, part: &mut [u8], max_imbalance: f64) -> bool {
let n = hg.num_vertices;
let Some(FmBalance {
mut weight,
min_part_weight,
max_part_weight,
}) = fm_balance(n, &hg.vertex_weights, part, max_imbalance)
else {
return false;
};
let mut pin_counts = hg.pin_counts(part);
let mut gain = vec![0i64; n];
let mut bq = [GainBuckets::new(n), GainBuckets::new(n)];
for v in 0..n {
let from = part[v] as usize;
let to = 1 - from;
let mut g = 0i64;
let mut on_boundary = false;
for &hei in hg.vertex_hyperedges(v) {
let hei = hei as usize;
let w = i64::from(hg.hyperedge_weights[hei]);
if pin_counts[hei][from] == 1 {
g += w;
}
if pin_counts[hei][to] == 0 {
g -= w;
}
if pin_counts[hei][0] > 0 && pin_counts[hei][1] > 0 {
on_boundary = true;
}
}
gain[v] = g;
if on_boundary {
bq[from].insert(v, g);
}
}
let mut locked = vec![false; n];
let mut moves: Vec<usize> = Vec::with_capacity(n);
let mut cumulative_gain: Vec<i64> = Vec::with_capacity(n);
let mut running_gain: i64 = 0;
let mut stall = Stall::new((n / 2).max(20));
for _ in 0..n {
let mut best_v: Option<usize> = None;
let mut best_gain = i64::MIN;
let mut best_from: usize = 0;
for side in 0..2 {
let candidate = bq[side].best_satisfying(|vertex| {
let to = 1 - side;
!locked[vertex]
&& weight[side] - hg.vertex_weights[vertex] >= min_part_weight
&& weight[to] + hg.vertex_weights[vertex] <= max_part_weight
});
if let Some(vertex) = candidate {
let g = gain[vertex];
if g > best_gain {
best_gain = g;
best_v = Some(vertex);
best_from = side;
}
}
}
let v = match best_v {
Some(v) => v,
None => break,
};
let from = best_from;
let to = 1 - from;
bq[from].remove(v);
weight[from] -= hg.vertex_weights[v];
weight[to] += hg.vertex_weights[v];
part[v] = to as u8;
locked[v] = true;
running_gain += best_gain;
moves.push(v);
cumulative_gain.push(running_gain);
if stall.record(running_gain) {
break;
}
for &hei in hg.vertex_hyperedges(v) {
let hei = hei as usize;
let old_from = pin_counts[hei][from];
let old_to = pin_counts[hei][to];
pin_counts[hei][from] -= 1;
pin_counts[hei][to] += 1;
let new_from = old_from - 1;
let _new_to = old_to + 1;
for &u in hg.charged_hyperedge_pins(hei) {
let u = u as usize;
if locked[u] {
continue;
}
let u_side = part[u] as usize;
let w = i64::from(hg.hyperedge_weights[hei]);
let mut delta = 0i64;
if u_side == from {
if old_from == 2 {
delta += w;
}
if old_to == 0 {
delta += w;
}
} else {
if old_to == 1 {
delta -= w;
}
if new_from == 0 {
delta -= w;
}
}
if delta != 0 {
gain[u] += delta;
}
let hyperedge_is_cut = pin_counts[hei][0] > 0 && pin_counts[hei][1] > 0;
let was_in_queue = bq[u_side].contains(u);
let was_cut = old_from > 0 && old_to > 0;
if hyperedge_is_cut != was_cut {
let on_boundary = if hyperedge_is_cut {
true
} else {
hg.vertex_hyperedges(u).iter().any(|&hej| {
let hej = hej as usize;
pin_counts[hej][0] > 0 && pin_counts[hej][1] > 0
})
};
if on_boundary {
if was_in_queue {
bq[u_side].update(u, gain[u]);
} else {
bq[u_side].insert(u, gain[u]);
}
} else if was_in_queue {
bq[u_side].remove(u);
}
} else if was_in_queue && delta != 0 {
bq[u_side].update(u, gain[u]);
}
}
}
}
commit_best_prefix(&moves, &cumulative_gain, part)
}
pub(super) fn localized_fm_pass(
hg: &Hypergraph,
part: &mut [u8],
seed: usize,
max_imbalance: f64,
) -> bool {
let n = hg.num_vertices;
let Some(FmBalance {
mut weight,
min_part_weight,
max_part_weight,
}) = fm_balance(n, &hg.vertex_weights, part, max_imbalance)
else {
return false;
};
let mut pin_counts = hg.pin_counts(part);
let max_region = (n / 4).max(20).min(n);
let mut in_region = vec![false; n];
let mut region_queue = std::collections::VecDeque::new();
in_region[seed] = true;
region_queue.push_back(seed);
let mut region_size = 1usize;
while let Some(v) = region_queue.pop_front() {
if region_size >= max_region {
break;
}
for &hei in hg.vertex_hyperedges(v) {
let hei_idx = hei as usize;
if pin_counts[hei_idx][0] == 0 || pin_counts[hei_idx][1] == 0 {
continue;
}
for &u in hg.charged_hyperedge_pins(hei_idx) {
let u = u as usize;
if !in_region[u] && region_size < max_region {
in_region[u] = true;
region_queue.push_back(u);
region_size += 1;
}
}
}
}
let mut gain = vec![0i64; n];
for v in 0..n {
if !in_region[v] {
continue;
}
let from = part[v] as usize;
let to = 1 - from;
for &hei in hg.vertex_hyperedges(v) {
let hei = hei as usize;
let w = i64::from(hg.hyperedge_weights[hei]);
if pin_counts[hei][from] == 1 {
gain[v] += w;
}
if pin_counts[hei][to] == 0 {
gain[v] -= w;
}
}
}
let mut locked = vec![false; n];
let mut moves: Vec<usize> = Vec::new();
let mut cumulative_gain: Vec<i64> = Vec::new();
let mut running_gain: i64 = 0;
let mut stall = Stall::new(region_size / 2);
let region_list: Vec<usize> = (0..n).filter(|&v| in_region[v]).collect();
for _ in 0..region_size {
let mut best_v = None;
let mut best_gain = i64::MIN;
for &v in ®ion_list {
if locked[v] {
continue;
}
let from = part[v] as usize;
let to = 1 - from;
let nfw = weight[from] - hg.vertex_weights[v];
let ntw = weight[to] + hg.vertex_weights[v];
if nfw < min_part_weight || ntw > max_part_weight {
continue;
}
if best_v.is_none() || gain[v] > best_gain {
best_gain = gain[v];
best_v = Some(v);
}
}
let Some(v) = best_v else {
break;
};
let from = part[v] as usize;
let to = 1 - from;
weight[from] -= hg.vertex_weights[v];
weight[to] += hg.vertex_weights[v];
part[v] = to as u8;
locked[v] = true;
running_gain += best_gain;
moves.push(v);
cumulative_gain.push(running_gain);
if stall.record(running_gain) {
break;
}
for &hei in hg.vertex_hyperedges(v) {
let hei = hei as usize;
let w = i64::from(hg.hyperedge_weights[hei]);
let old_from = pin_counts[hei][from];
let old_to = pin_counts[hei][to];
pin_counts[hei][from] -= 1;
pin_counts[hei][to] += 1;
let new_from = old_from - 1;
for &u in hg.charged_hyperedge_pins(hei) {
let u = u as usize;
if locked[u] || !in_region[u] {
continue;
}
let u_side = part[u] as usize;
let mut delta = 0i64;
if u_side == from {
if old_from == 2 {
delta += w;
}
if old_to == 0 {
delta += w;
}
} else {
if old_to == 1 {
delta -= w;
}
if new_from == 0 {
delta -= w;
}
}
if delta != 0 {
gain[u] += delta;
}
}
}
}
commit_best_prefix(&moves, &cumulative_gain, part)
}
pub(super) fn refine_level(hg: &Hypergraph, part: &mut [u8], imbalance: f64) {
let max_passes = 10;
for _ in 0..max_passes {
if !fm_refine_pass(hg, part, imbalance) {
break;
}
}
}