use super::network_simplex_value_type::{MulWithFloat, ToBigInt};
use core::convert::From;
use ebi_arithmetic::exact::MaybeExact;
use ebi_arithmetic::rand::rng;
use ebi_arithmetic::rand::seq::SliceRandom;
use ebi_arithmetic::{One, Signed, Zero, malachite::Integer};
use rayon::ThreadPool;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::{
cmp::{PartialEq, PartialOrd},
fmt::{Debug, Display},
iter::Sum,
ops::{AddAssign, MulAssign, Neg, SubAssign},
};
#[derive(Debug, PartialEq)]
pub enum ProblemType {
Optimal,
Infeasible,
Unbounded,
}
#[derive(Debug, PartialEq)]
pub enum SupplyType {
GEQ,
LEQ,
}
#[derive(Debug, PartialEq, Clone)]
pub enum ArcState<T> {
Upper(T),
Tree(T),
Lower(T),
}
impl<T> ArcState<T>
where
T: From<i32>, {
pub fn upper() -> Self {
ArcState::Upper(T::from(-1))
}
pub fn tree() -> Self {
ArcState::Tree(T::from(0))
}
pub fn lower() -> Self {
ArcState::Lower(T::from(1))
}
pub fn value(&self) -> &T {
match self {
ArcState::Upper(v) => v,
ArcState::Tree(v) => v,
ArcState::Lower(v) => v,
}
}
}
#[derive(Debug, PartialEq, Clone)]
pub enum ArcDirection<T> {
Down(T),
Up(T),
}
impl<T> ArcDirection<T>
where
T: From<i32>,
{
pub fn down() -> Self {
ArcDirection::Down(T::from(-1))
}
pub fn up() -> Self {
ArcDirection::Up(T::from(1))
}
pub fn value(&self) -> &T {
match self {
ArcDirection::Down(v) => v,
ArcDirection::Up(v) => v,
}
}
}
const EPSILON: f64 = 1e-15;
pub struct NetworkSimplex<T> {
node_num: usize,
all_node_num: usize,
arc_num: usize,
all_arc_num: usize,
search_arc_num: usize,
node_id: Vec<usize>,
source: Vec<usize>, target: Vec<usize>, cost: Vec<T>, supply: Vec<T>, sum_supply: T,
supply_type: SupplyType,
flow: Vec<T>, pi: Vec<T>,
parent: Vec<Option<usize>>, predecessor: Vec<Option<usize>>, thread: Vec<usize>, reverse_thread: Vec<usize>, successor_num: Vec<usize>, last_successor: Vec<usize>, predecessor_direction: Vec<ArcDirection<T>>, state: Vec<ArcState<T>>, dirty_revs: Vec<usize>, root: usize,
in_arc: usize,
join: usize,
u_in: usize,
v_in: usize,
u_out: usize,
v_out: usize,
delta: T, max: T,
block_size: usize,
next_arc: usize,
problem_type: Option<ProblemType>,
}
impl<T> NetworkSimplex<T>
where
T: Zero
+ One
+ MaybeExact
+ MulWithFloat
+ Clone
+ for<'a> AddAssign<&'a T>
+ for<'a> SubAssign<&'a T>
+ for<'a> MulAssign<&'a T>
+ Neg<Output = T>
+ Signed
+ PartialEq
+ PartialOrd
+ ?Sized
+ Display
+ Debug
+ From<i32>
+ Sum
+ Send
+ Sync
+ ToBigInt
+ 'static,
{
pub fn new(
graph_and_costs: &Vec<Vec<Option<T>>>,
supply: &Vec<T>,
arc_mixing: bool,
greater_eq_supply: bool,
) -> Self {
let node_num = supply.len();
assert!(
graph_and_costs.len() == node_num,
"Graph size and supply size mismatch"
);
for row in graph_and_costs.iter() {
assert!(row.len() == node_num, "Graph matrix not square");
}
let node_id: Vec<usize> = (0..node_num).collect();
let supply = (*supply).clone();
let mut source = vec![];
let mut target = vec![];
let mut cost = vec![];
for i in 0..node_num {
for j in 0..node_num {
if let Some(c) = &graph_and_costs[i][j] {
assert!(i != j, "Tried to add arc from node to itself");
source.push(i);
target.push(j);
cost.push((*c).clone());
}
}
}
let arc_num = cost.len();
if arc_mixing {
let mut arcs: Vec<_> = source
.iter()
.zip(target.iter())
.zip(cost.iter())
.map(|((src, tgt), cst)| (*src, *tgt, cst.clone()))
.collect();
let mut rng = rng();
arcs.shuffle(&mut rng);
source.clear();
target.clear();
cost.clear();
for (src, tgt, cst) in arcs {
source.push(src);
target.push(tgt);
cost.push(cst);
}
}
let block_size_factor = 1.0;
let min_block_size = 10;
let block_size =
((block_size_factor * (arc_num as f64).sqrt()) as usize).max(min_block_size);
let supply_type = if greater_eq_supply {
SupplyType::GEQ
} else {
SupplyType::LEQ
};
let ns = NetworkSimplex {
node_num,
all_node_num: node_num,
arc_num,
all_arc_num: arc_num,
search_arc_num: 0,
sum_supply: T::zero(),
node_id,
source,
target,
cost,
supply,
flow: vec![],
pi: vec![],
parent: vec![],
predecessor: vec![],
thread: vec![],
reverse_thread: vec![],
successor_num: vec![],
last_successor: vec![],
predecessor_direction: vec![],
state: vec![],
dirty_revs: vec![],
root: 0,
in_arc: 0,
join: 0,
u_in: 0,
v_in: 0,
u_out: 0,
v_out: 0,
delta: T::zero(),
max: T::one(),
block_size,
next_arc: 0,
problem_type: None,
supply_type,
};
ns
}
pub fn visualize_network(&self) {
let mut nodes_output = String::new();
for i in 0..self.node_id.len() {
let node = self.node_id[i];
let supply = &self.supply[i];
nodes_output.push_str(&format!("{}({})", node, supply));
if i < self.node_id.len() - 1 {
nodes_output.push_str(", ");
}
}
let mut arcs_output = String::new();
for i in 0..self.all_arc_num {
let source = self.source[i];
let target = self.target[i];
let cost = &self.cost[i];
arcs_output.push_str(&format!("{}--({})-->{}", source, cost, target));
if i < self.all_arc_num - 1 {
arcs_output.push_str(", ");
}
}
}
pub fn visualize_tree_graphviz(&self) -> String {
let mut graphviz_code = String::new();
graphviz_code.push_str("digraph Tree {\n");
graphviz_code.push_str(&format!(
" {} [label=\"{} (Root)\", shape=box];\n",
self.root, self.root
));
for i in 0..self.all_node_num {
if self.parent[i] != None {
let parent = self.parent[i].unwrap();
let direction = &self.predecessor_direction[i];
let flow = &self.flow[self.predecessor[i].unwrap()];
if direction.value() == &T::from(1) {
graphviz_code
.push_str(&format!(" {} -> {} [label=\"{}\"];\n", i, parent, *flow));
} else {
graphviz_code
.push_str(&format!(" {} -> {} [label=\"{}\"];\n", parent, i, *flow));
}
}
}
graphviz_code.push_str("}\n");
graphviz_code
}
pub fn run(&mut self, guarantee_network_feasibility: bool) -> ProblemType {
if !self.initialize_feasible_solution() {
self.problem_type = Some(ProblemType::Infeasible);
log::info!("Could not initialize feasible solution");
return ProblemType::Infeasible;
}
let mut iter = 1;
let num_threads = rayon::current_num_threads();
let pool = rayon::ThreadPoolBuilder::new()
.num_threads(num_threads)
.build()
.unwrap();
while self.find_entering_arc_par(&pool) {
iter += 1;
self.find_join_node();
let change = self.find_leaving_arc();
if self.delta >= self.max {
self.problem_type = Some(ProblemType::Unbounded);
log::info!("The current Network is unbounded");
return ProblemType::Unbounded;
}
self.change_flow(change);
if change {
self.update_tree_structure();
self.update_potential(); }
}
log::info!("Network Simplex finished in {} iterations", iter);
if !guarantee_network_feasibility {
if !T::is_exact(&self.sum_supply) {
for e in self.search_arc_num..self.all_arc_num {
if self.flow[e] > T::one().mul_with_float(&EPSILON) {
self.problem_type = Some(ProblemType::Infeasible);
log::info!(
"The current Network is infeasible, flow remains on artificial arcs"
);
return ProblemType::Infeasible;
}
}
} else {
for e in self.search_arc_num..self.all_arc_num {
if self.flow[e] != T::zero() {
self.problem_type = Some(ProblemType::Infeasible);
log::info!(
"The current Network is infeasible, flow remains on artificial arcs"
);
return ProblemType::Infeasible;
}
}
}
}
self.problem_type = Some(ProblemType::Optimal);
log::info!("Optimal solution found");
return ProblemType::Optimal;
}
fn find_entering_arc(&mut self) -> bool {
let mut cost: T;
let mut min_cost = T::zero();
let mut count = self.block_size;
for e in self.next_arc..self.search_arc_num {
cost = self.cost[e].clone();
cost += &self.pi[self.source[e]];
cost -= &self.pi[self.target[e]];
cost *= self.state[e].value();
log::trace!(
"{}-->{}, cost: {} = {} * ({} + {} - {})",
self.source[e],
self.target[e],
cost,
self.state[e].value(),
self.cost[e],
self.pi[self.source[e]],
self.pi[self.target[e]]
);
if cost < min_cost {
min_cost = cost;
self.in_arc = e;
}
count -= 1;
if count == 0 {
if !T::is_exact(&min_cost) {
let source_value = self.pi[self.source[self.in_arc]].clone().abs();
let target_value = self.pi[self.target[self.in_arc]].clone().abs();
let cost_value = self.cost[self.in_arc].clone().abs();
let mut a = if source_value > target_value {
source_value
} else {
target_value
};
a = if a > cost_value { a } else { cost_value };
if min_cost < -a.mul_with_float(&EPSILON) {
self.next_arc = e;
return true;
}
} else {
if min_cost < T::zero() {
self.next_arc = e;
return true;
}
}
count = self.block_size;
}
}
for e in 0..self.next_arc {
cost = self.cost[e].clone();
cost += &self.pi[self.source[e]];
cost -= &self.pi[self.target[e]];
cost *= self.state[e].value();
log::trace!(
"{}-->{}, cost: {} = {} * ({} + {} - {})",
self.source[e],
self.target[e],
cost,
self.state[e].value(),
self.cost[e],
self.pi[self.source[e]],
self.pi[self.target[e]]
);
if cost < min_cost {
min_cost = cost;
self.in_arc = e;
}
count -= 1;
if count == 0 {
if min_cost < T::zero() {
self.next_arc = e;
return true;
}
count = self.block_size;
}
if count == 0 {
if !T::is_exact(&min_cost) {
let source_value = self.pi[self.source[self.in_arc]].clone().abs();
let target_value = self.pi[self.target[self.in_arc]].clone().abs();
let cost_value = self.cost[self.in_arc].clone().abs();
let mut a = if source_value > target_value {
source_value
} else {
target_value
};
a = if a > cost_value { a } else { cost_value };
if min_cost < -a.mul_with_float(&EPSILON) {
self.next_arc = e;
return true;
}
} else {
if min_cost < T::zero() {
self.next_arc = e;
return true;
}
}
count = self.block_size;
}
}
if !T::is_exact(&min_cost) {
let source_value = self.pi[self.source[self.in_arc]].clone().abs();
let target_value = self.pi[self.target[self.in_arc]].clone().abs();
let cost_value = self.cost[self.in_arc].clone().abs();
let mut a = if source_value > target_value {
source_value
} else {
target_value
};
a = if a > cost_value { a } else { cost_value };
if min_cost >= -a.mul_with_float(&EPSILON) {
return false;
}
} else {
if min_cost >= T::zero() {
return false;
}
}
true
}
fn find_entering_arc_par(&mut self, pool: &ThreadPool) -> bool {
self.find_entering_arc_par_recursive(pool, 0)
}
fn find_entering_arc_par_recursive(&mut self, pool: &ThreadPool, arcs_visited: usize) -> bool {
let num_threads = pool.current_num_threads();
let block_size_per_thread = self.block_size / num_threads;
let arcs_per_thread = self.search_arc_num / num_threads;
if block_size_per_thread == 0 || arcs_per_thread < 500 {
return self.find_entering_arc();
}
let min_cost = Arc::new(parking_lot::Mutex::new(T::zero()));
let min_arc = Arc::new(AtomicUsize::new(0));
let cost = &self.cost;
let pi = &self.pi;
let source = &self.source;
let target = &self.target;
let state = &self.state;
let next_arc = self.next_arc;
let search_arc_num = self.search_arc_num;
pool.install(|| {
pool.scope(|scope| {
for thread_idx in 0..num_threads {
let min_cost = Arc::clone(&min_cost);
let min_arc = Arc::clone(&min_arc);
scope.spawn(move |_| {
let start = next_arc + thread_idx * arcs_per_thread;
let end = start + arcs_per_thread;
let mut thread_min_cost = T::zero();
let mut thread_min_arc = start;
let mut first_iteration = true;
let mut current = start;
while current < end {
let e = if current >= search_arc_num {
current - search_arc_num
} else {
current
};
let mut cost = cost[e].clone();
cost += &pi[source[e]];
cost -= &pi[target[e]];
cost *= state[e].value();
if first_iteration || cost < thread_min_cost {
thread_min_cost = cost;
thread_min_arc = e;
first_iteration = false;
}
current += 1;
}
if thread_min_cost < T::zero() {
let mut global_min = min_cost.lock();
if first_iteration || thread_min_cost < *global_min {
*global_min = thread_min_cost;
min_arc.store(thread_min_arc, Ordering::Relaxed);
}
}
});
}
});
});
let final_min_cost = min_cost.lock().clone();
let final_min_arc = min_arc.load(Ordering::Relaxed);
let arcs_in_block = num_threads * arcs_per_thread;
let new_arcs_visited = arcs_visited + arcs_in_block;
self.next_arc = (next_arc + arcs_in_block) % search_arc_num;
self.in_arc = final_min_arc;
let valid_arc_found = if !T::is_exact(&final_min_cost) {
let source_value = self.pi[self.source[self.in_arc]].clone().abs();
let target_value = self.pi[self.target[self.in_arc]].clone().abs();
let cost_value = self.cost[self.in_arc].clone().abs();
let mut a = if source_value > target_value {
source_value
} else {
target_value
};
a = if a > cost_value { a } else { cost_value };
final_min_cost < -a.mul_with_float(&EPSILON)
} else {
final_min_cost < T::zero()
};
if valid_arc_found {
true
} else if new_arcs_visited >= search_arc_num {
false
} else {
self.find_entering_arc_par_recursive(pool, new_arcs_visited)
}
}
fn find_leaving_arc(&mut self) -> bool {
let first;
let second;
if self.state[self.in_arc].value() == &T::from(1) {
first = self.source[self.in_arc];
second = self.target[self.in_arc];
} else {
first = self.target[self.in_arc];
second = self.source[self.in_arc];
}
self.delta = self.max.clone();
let mut result = 0;
let mut d;
let mut e;
let mut u = Some(first);
while let Some(u_node) = u {
if u_node == self.join {
break;
}
e = self.predecessor[u_node].unwrap();
d = &self.flow[e];
if self.predecessor_direction[u_node].value() == &T::from(-1) {
d = &self.max;
}
if *d < self.delta {
self.delta = d.clone();
self.u_out = u_node;
result = 1;
}
u = self.parent[u_node];
}
let mut u = Some(second);
while let Some(u_node) = u {
if u_node == self.join {
break;
}
e = self.predecessor[u_node].unwrap();
d = &self.flow[e];
if self.predecessor_direction[u_node].value() == &T::from(1) {
d = &self.max;
}
if *d < self.delta {
self.delta = d.clone();
self.u_out = u_node;
result = 2;
}
u = self.parent[u_node];
}
if result == 1 {
self.u_in = first;
self.v_in = second;
} else {
self.u_in = second;
self.v_in = first;
}
return result != 0;
}
fn update_potential(&mut self) {
let mut sigma = -self.cost[self.in_arc].clone();
sigma *= &(self.predecessor_direction[self.u_in].value());
sigma += &self.pi[self.v_in];
sigma -= &self.pi[self.u_in];
let end = self.thread[self.last_successor[self.u_in]];
let mut u = self.u_in;
while u != end {
self.pi[u] += σ
u = self.thread[u];
}
}
fn initialize_feasible_solution(&mut self) -> bool {
if self.node_num == 0 {
log::info!("No nodes in the graph");
return false;
}
self.sum_supply = T::zero();
for i in 0..self.node_num {
self.sum_supply += &self.supply[i];
if self.supply[i].is_positive() {
self.max += &self.supply[i]
}
}
if !((self.supply_type == SupplyType::GEQ && self.sum_supply <= T::zero())
|| (self.supply_type == SupplyType::LEQ && self.sum_supply >= T::zero()))
{
log::info!("Sum of supply is invalid, try changing supply type");
return false;
}
let mut max_cost = self.find_max_cost();
max_cost += &T::one();
max_cost *= &T::from(self.node_num as i32);
let art_cost: T = max_cost;
self.all_node_num = self.node_num + 1;
let max_arc_num = self.arc_num + 2 * self.node_num;
self.all_arc_num = self.arc_num + self.node_num;
self.source.resize(max_arc_num, 0);
self.target.resize(max_arc_num, 0);
self.flow.resize(max_arc_num, T::zero());
self.state.resize(max_arc_num, ArcState::lower());
self.cost.resize(max_arc_num, T::zero());
self.supply.resize(self.all_node_num, T::zero());
self.pi.resize(self.all_node_num, T::zero());
self.parent.resize(self.all_node_num, Some(0));
self.predecessor.resize(self.all_node_num, Some(0));
self.predecessor_direction
.resize(self.all_node_num, ArcDirection::up());
self.thread.resize(self.all_node_num, 0);
self.reverse_thread.resize(self.all_node_num, 0);
self.successor_num.resize(self.all_node_num, 0);
self.last_successor.resize(self.all_node_num, 0);
for i in 0..self.node_num {
self.flow[i] = T::zero();
self.state[i] = ArcState::lower();
}
self.root = self.node_num;
self.node_id.push(self.root);
self.parent[self.root] = None;
self.predecessor[self.root] = None;
self.thread[self.root] = 0;
self.reverse_thread[0] = self.root;
self.successor_num[self.root] = self.node_num + 1; self.last_successor[self.root] = self.node_num - 1;
self.supply[self.root] = -self.sum_supply.clone();
self.pi[self.root] = T::zero();
if self.sum_supply == T::zero() {
self.search_arc_num = self.arc_num;
let mut e = self.arc_num;
for u in 0..self.node_num {
self.parent[u] = Some(self.root);
self.predecessor[u] = Some(e);
self.thread[u] = u + 1;
self.reverse_thread[u + 1] = u;
self.successor_num[u] = 1;
self.last_successor[u] = u;
self.state[e] = ArcState::tree();
if !self.supply[u].is_negative() {
self.predecessor_direction[u] = ArcDirection::up();
self.pi[u] = T::zero();
self.source[e] = u;
self.target[e] = self.root;
self.flow[e] = self.supply[u].clone();
self.cost[e] = T::zero();
} else {
self.predecessor_direction[u] = ArcDirection::down();
self.pi[u] = art_cost.clone();
self.source[e] = self.root;
self.target[e] = u;
self.flow[e] = -self.supply[u].clone();
self.cost[e] = art_cost.clone();
}
e += 1;
}
} else if self.sum_supply > T::zero() {
self.search_arc_num = self.arc_num + self.node_num;
let mut f = self.arc_num + self.node_num;
for u in 0..self.node_num {
self.parent[u] = Some(self.root);
self.thread[u] = u + 1;
self.reverse_thread[u + 1] = u;
self.successor_num[u] = 1;
self.last_successor[u] = u;
if !self.supply[u].is_negative() {
self.predecessor_direction[u] = ArcDirection::up();
self.pi[u] = T::zero();
self.predecessor[u] = Some(self.arc_num + u);
self.source[self.arc_num + u] = u;
self.target[self.arc_num + u] = self.root;
self.state[self.arc_num + u] = ArcState::tree();
self.flow[self.arc_num + u] = self.supply[u].clone();
self.cost[self.arc_num + u] = T::zero();
} else {
self.predecessor_direction[u] = ArcDirection::down();
self.pi[u] = art_cost.clone();
self.predecessor[u] = Some(f);
self.source[f] = self.root;
self.target[f] = u;
self.state[f] = ArcState::tree();
self.flow[f] = -self.supply[u].clone();
self.cost[f] = art_cost.clone();
self.source[self.arc_num + u] = u;
self.target[self.arc_num + u] = self.root;
self.state[self.arc_num + u] = ArcState::lower();
self.flow[self.arc_num + u] = T::zero();
self.cost[self.arc_num + u] = T::zero();
f += 1;
}
}
self.all_arc_num = f;
} else {
self.search_arc_num = self.arc_num + self.node_num;
let mut f = self.arc_num + self.node_num;
for u in 0..self.node_num {
self.parent[u] = Some(self.root);
self.thread[u] = u + 1;
self.reverse_thread[u + 1] = u;
self.successor_num[u] = 1;
self.last_successor[u] = u;
if !self.supply[u].is_positive() {
self.predecessor_direction[u] = ArcDirection::down();
self.pi[u] = T::zero();
self.predecessor[u] = Some(self.arc_num + u);
self.source[self.arc_num + u] = self.root;
self.target[self.arc_num + u] = u;
self.state[self.arc_num + u] = ArcState::tree();
self.flow[self.arc_num + u] = -self.supply[u].clone();
self.cost[self.arc_num + u] = T::zero();
} else {
self.predecessor_direction[u] = ArcDirection::up();
self.pi[u] = -art_cost.clone();
self.predecessor[u] = Some(f);
self.source[f] = u;
self.target[f] = self.root;
self.state[f] = ArcState::tree();
self.flow[f] = self.supply[u].clone();
self.cost[f] = art_cost.clone();
self.source[self.arc_num + u] = self.root;
self.target[self.arc_num + u] = u;
self.state[self.arc_num + u] = ArcState::lower();
self.flow[self.arc_num + u] = T::zero();
self.cost[self.arc_num + u] = T::zero();
f += 1;
}
}
self.all_arc_num = f;
}
return true;
}
fn find_join_node(&mut self) {
let mut u = self.source[self.in_arc];
let mut v = self.target[self.in_arc];
while u != v {
if self.successor_num[u] < self.successor_num[v] {
u = self.parent[u].unwrap();
} else {
v = self.parent[v].unwrap();
}
}
self.join = u;
}
fn change_flow(&mut self, change: bool) {
if self.delta > T::zero() {
let mut value = self.state[self.in_arc].value().clone();
value *= &self.delta;
self.flow[self.in_arc] += &value;
let mut u = self.source[self.in_arc];
while u != self.join {
let mut reduce_by = self.predecessor_direction[u].value().clone();
reduce_by *= &value;
self.flow[self.predecessor[u].unwrap()] -= &reduce_by;
u = self.parent[u].unwrap();
}
u = self.target[self.in_arc];
while u != self.join {
let mut increase_by = self.predecessor_direction[u].value().clone();
increase_by *= &value;
self.flow[self.predecessor[u].unwrap()] += &increase_by;
u = self.parent[u].unwrap();
}
}
if change {
self.state[self.in_arc] = ArcState::tree();
if self.flow[self.predecessor[self.u_out].unwrap()] == T::zero() {
self.state[self.predecessor[self.u_out].unwrap()] = ArcState::lower();
} else {
self.state[self.predecessor[self.u_out].unwrap()] = ArcState::upper();
}
} else {
if self.state[self.in_arc] == ArcState::lower() {
self.state[self.in_arc] = ArcState::upper();
} else {
self.state[self.in_arc] = ArcState::lower();
}
}
}
fn update_tree_structure(&mut self) {
let old_reverse_thread = self.reverse_thread[self.u_out];
let old_successor_num = self.successor_num[self.u_out];
let old_last_successor = self.last_successor[self.u_out];
self.v_out = self.parent[self.u_out].unwrap();
if self.u_in == self.u_out {
self.parent[self.u_in] = Some(self.v_in);
self.predecessor[self.u_in] = Some(self.in_arc);
self.predecessor_direction[self.u_in] = if self.u_in == self.source[self.in_arc] {
ArcDirection::up()
} else {
ArcDirection::down()
};
if self.thread[self.v_in] != self.u_out {
let mut after = self.thread[old_last_successor];
self.thread[old_reverse_thread] = after;
self.reverse_thread[after] = old_reverse_thread;
after = self.thread[self.v_in];
self.thread[self.v_in] = self.u_out;
self.reverse_thread[self.u_out] = self.v_in;
self.thread[old_last_successor] = after;
self.reverse_thread[after] = old_last_successor;
}
} else {
let thread_continue = if old_reverse_thread == self.v_in {
self.thread[old_last_successor]
} else {
self.thread[self.v_in]
};
let mut stem = self.u_in;
let mut stem_parent = self.v_in;
let mut next_stem;
let mut last = self.last_successor[self.u_in];
let mut before;
let mut after = self.thread[last];
self.thread[self.v_in] = self.u_in;
self.dirty_revs.clear();
self.dirty_revs.push(self.v_in);
while stem != self.u_out {
next_stem = self.parent[stem].unwrap();
self.thread[last] = next_stem;
self.dirty_revs.push(last);
before = self.reverse_thread[stem];
self.thread[before] = after;
self.reverse_thread[after] = before;
self.parent[stem] = Some(stem_parent);
stem_parent = stem;
stem = next_stem;
last = if self.last_successor[stem] == self.last_successor[stem_parent] {
self.reverse_thread[stem_parent]
} else {
self.last_successor[stem]
};
after = self.thread[last];
}
self.parent[self.u_out] = Some(stem_parent);
self.thread[last] = thread_continue;
self.reverse_thread[thread_continue] = last;
self.last_successor[self.u_out] = last;
if old_reverse_thread != self.v_in {
self.thread[old_reverse_thread] = after;
self.reverse_thread[after] = old_reverse_thread;
}
for i in 0..self.dirty_revs.len() {
let u = self.dirty_revs[i];
self.reverse_thread[self.thread[u]] = u;
}
let mut temp_successor_num = 0;
let temp_last_successor = self.last_successor[self.u_out];
let mut u = self.u_out;
let mut p = self.parent[u];
while u != self.u_in {
self.predecessor[u] = self.predecessor[p.unwrap()];
self.predecessor_direction[u] =
if self.predecessor_direction[p.unwrap()] == ArcDirection::up() {
ArcDirection::down()
} else {
ArcDirection::up()
};
temp_successor_num += self.successor_num[u] - self.successor_num[p.unwrap()];
self.successor_num[u] = temp_successor_num;
self.last_successor[p.unwrap()] = temp_last_successor;
u = p.unwrap();
p = self.parent[u];
}
self.predecessor[self.u_in] = Some(self.in_arc);
self.predecessor_direction[self.u_in] = if self.u_in == self.source[self.in_arc] {
ArcDirection::up()
} else {
ArcDirection::down()
};
self.successor_num[self.u_in] = old_successor_num;
}
let up_limit_out = if self.last_successor[self.join] == self.v_in {
Some(self.join)
} else {
None
};
let last_successor_out = self.last_successor[self.u_out];
let mut u = Some(self.v_in);
while u != None && self.last_successor[u.unwrap()] == self.v_in {
self.last_successor[u.unwrap()] = last_successor_out;
u = self.parent[u.unwrap()];
}
if self.join != old_reverse_thread && self.v_in != old_reverse_thread {
u = Some(self.v_out);
while u != None
&& u != up_limit_out
&& self.last_successor[u.unwrap()] == old_last_successor
{
self.last_successor[u.unwrap()] = old_reverse_thread;
u = self.parent[u.unwrap()];
}
} else if last_successor_out != old_last_successor {
u = Some(self.v_out);
while u != None
&& u != up_limit_out
&& self.last_successor[u.unwrap()] == old_last_successor
{
self.last_successor[u.unwrap()] = last_successor_out;
u = self.parent[u.unwrap()];
}
}
let mut u = self.v_in;
while u != self.join {
self.successor_num[u] += old_successor_num;
u = self.parent[u].unwrap();
}
u = self.v_out;
while u != self.join {
self.successor_num[u] -= old_successor_num;
u = self.parent[u].unwrap();
}
}
pub fn get_result(&self) -> Option<T> {
if let Some(problem_type) = &self.problem_type {
if problem_type == &ProblemType::Optimal {
let flow_cost = self.flow.iter().zip(self.cost.iter());
let mut result = T::zero();
for (flow, cost) in flow_cost {
let mut arc_result = flow.clone();
arc_result *= cost;
result += &arc_result;
}
return Some(result);
}
}
return None;
}
pub fn get_bigint_result(&self) -> Option<Integer> {
if let Some(problem_type) = &self.problem_type {
if problem_type == &ProblemType::Optimal {
let flow_cost = self.flow.iter().zip(self.cost.iter());
let mut result = Integer::zero();
for (flow, cost) in flow_cost {
let mut arc_result = flow.to_big_int();
arc_result *= cost.to_big_int();
result += arc_result;
}
return Some(result);
}
}
None
}
pub fn get_flow(&self) -> Vec<T> {
self.flow.clone()
}
pub fn get_cost(&self) -> Vec<T> {
self.cost.clone()
}
fn find_max_cost(&self) -> T {
select_max(&self.cost).expect("Cost vector cannot be empty")
}
}
pub fn select_max<T>(values: &[T]) -> Option<T>
where
T: PartialOrd + Clone,
{
values
.iter()
.filter(|x| x.partial_cmp(x).is_some()) .max_by(|a, b| a.partial_cmp(b).unwrap())
.cloned()
}
#[cfg(test)]
mod tests {
use crate::network_simplex::NetworkSimplex;
use ebi_arithmetic::malachite::Integer;
#[test]
fn network_simplex_int() {
let supply: Vec<i64> = vec![20, 0, 0, -5, -14];
let graph_and_costs: Vec<Vec<Option<i64>>> = vec![
vec![None, Some(4), Some(4), None, None],
vec![None, None, Some(2), Some(2), Some(6)],
vec![None, None, None, Some(1), Some(3)],
vec![None, None, None, None, Some(2)],
vec![None, None, Some(3), None, None],
];
let mut ns = NetworkSimplex::new(&graph_and_costs, &supply, true, false);
_ = ns.run(false);
assert_eq!(ns.get_result().unwrap(), 123);
}
#[test]
fn network_simplex_bigint() {
let supply: Vec<Integer> = vec![20.into(), 0.into(), 0.into(), (-5).into(), (-14).into()];
let graph_and_costs: Vec<Vec<Option<Integer>>> = vec![
vec![None, Some(4), Some(4), None, None],
vec![None, None, Some(2), Some(2), Some(6)],
vec![None, None, None, Some(1), Some(3)],
vec![None, None, None, None, Some(2)],
vec![None, None, Some(3), None, None],
]
.into_iter()
.map(|row| {
row.into_iter()
.map(|x| x.map(|cost| Integer::from(cost)))
.collect()
})
.collect();
let mut ns = NetworkSimplex::new(&graph_and_costs, &supply, true, false);
_ = ns.run(false);
assert_eq!(ns.get_result().unwrap(), Integer::from(123));
}
#[test]
fn network_simplex_float() {
let supply: Vec<f64> = vec![20, 0, 0, -5, -14]
.into_iter()
.map(|s| s.into())
.collect();
let graph_and_costs: Vec<Vec<Option<f64>>> = vec![
vec![None, Some(4), Some(4), None, None],
vec![None, None, Some(2), Some(2), Some(6)],
vec![None, None, None, Some(1), Some(3)],
vec![None, None, None, None, Some(2)],
vec![None, None, Some(3), None, None],
]
.into_iter()
.map(|row| row.into_iter().map(|x| x.map(|cost| cost.into())).collect())
.collect();
let mut ns = NetworkSimplex::new(&graph_and_costs, &supply, true, false);
_ = ns.run(false);
let result = ns.get_result().unwrap();
assert_eq!(result, 123.0);
}
}