use std::collections::{HashMap, HashSet};
use crate::hyperpath::ALPHA;
use crate::hyperpath_queue::PriorityQueue;
use crate::transit_network::Link;
fn intern(name: &str, n_id: &mut HashMap<String, usize>, n_name: &mut Vec<String>) -> usize {
if let Some(&id) = n_id.get(name) {
return id;
}
let id = n_name.len();
n_id.insert(name.to_string(), id);
n_name.push(name.to_string());
id
}
pub struct Graph {
n_name: Vec<String>,
n_id: HashMap<String, usize>,
n: usize,
m: usize,
from: Vec<usize>,
to: Vec<usize>,
cost: Vec<f64>,
head: Vec<f64>,
adj_by_to: Vec<Vec<usize>>,
}
impl Graph {
pub fn new(all_links: &[Link], all_stops: &HashSet<String>) -> Graph {
let mut n_id: HashMap<String, usize> = HashMap::with_capacity(all_stops.len());
let mut n_name: Vec<String> = Vec::with_capacity(all_stops.len());
for stop in all_stops {
intern(stop.as_str(), &mut n_id, &mut n_name);
}
let m = all_links.len();
let mut from = vec![0usize; m];
let mut to = vec![0usize; m];
let mut cost = vec![0.0f64; m];
let mut head = vec![0.0f64; m];
for (k, link) in all_links.iter().enumerate() {
from[k] = intern(link.from_node.as_str(), &mut n_id, &mut n_name);
to[k] = intern(link.to_node.as_str(), &mut n_id, &mut n_name);
cost[k] = link.travel_cost;
head[k] = link.headway;
}
let n = n_name.len();
let mut adj_by_to: Vec<Vec<usize>> = vec![Vec::new(); n];
for k in 0..m {
adj_by_to[to[k]].push(k);
}
Graph {
n_name,
n_id,
n,
m,
from,
to,
cost,
head,
adj_by_to,
}
}
pub fn num_nodes(&self) -> usize {
self.n
}
pub fn num_links(&self) -> usize {
self.m
}
pub fn node_index(&self, name: &str) -> Option<usize> {
self.n_id.get(name).copied()
}
pub fn node_name(&self, id: usize) -> &str {
&self.n_name[id]
}
pub fn new_workspace(&self) -> Workspace<'_> {
Workspace {
g: self,
u: vec![0.0; self.n],
f: vec![0.0; self.n],
pq: PriorityQueue::with_capacity(self.m),
overline_a: Vec::with_capacity(self.m / 2),
a_set_idx: vec![Vec::new(); self.n],
a_set: Vec::with_capacity(self.m / 2),
link_vol: vec![0.0; self.m],
node_vol: vec![0.0; self.n],
cols: vec![Vec::new(); self.n],
}
}
}
pub struct DestResult<'w> {
pub dest_id: usize,
pub labels: &'w [f64],
pub freqs: &'w [f64],
pub a_set: &'w [usize],
pub link_vol: &'w [f64],
pub node_vol: &'w [f64],
}
pub struct Workspace<'g> {
g: &'g Graph,
u: Vec<f64>,
f: Vec<f64>,
pq: PriorityQueue,
overline_a: Vec<Option<usize>>,
a_set_idx: Vec<Vec<usize>>,
a_set: Vec<usize>,
link_vol: Vec<f64>,
node_vol: Vec<f64>,
cols: Vec<Vec<(usize, f64)>>,
}
impl Workspace<'_> {
fn find_strategy(&mut self, dest_id: usize) {
let g = self.g;
for id in 0..g.n {
self.f[id] = 0.0;
self.u[id] = if id == dest_id { 0.0 } else { f64::INFINITY };
self.a_set_idx[id].clear();
}
self.overline_a.clear();
self.pq.clear();
for k in 0..g.m {
self.pq.push(k, self.u[g.to[k]] + g.cost[k]);
}
self.pq.init();
while self.pq.len() > 0 {
let entry_id = match self.pq.pop() {
Some(id) => id,
None => break,
};
let priority = self.pq.priority(entry_id);
if priority.is_infinite() && priority > 0.0 {
break;
}
let k = self.pq.link(entry_id);
let i = g.from[k];
let j = g.to[k];
let sum_uc = self.u[j] + g.cost[k];
if self.f[i].is_infinite() {
continue;
}
if self.u[i] <= sum_uc {
continue;
}
if g.head[k] <= 0.0 {
self.u[i] = sum_uc;
self.f[i] = f64::INFINITY;
for &idx in &self.a_set_idx[i] {
self.overline_a[idx] = None;
}
self.a_set_idx[i].clear();
self.overline_a.push(Some(k));
self.a_set_idx[i].push(self.overline_a.len() - 1);
} else {
let freq = 1.0 / g.head[k];
let new_u = if self.f[i] == 0.0 {
(ALPHA + freq * sum_uc) / freq
} else {
(self.f[i] * self.u[i] + freq * sum_uc) / (self.f[i] + freq)
};
self.u[i] = new_u;
self.f[i] += freq;
self.overline_a.push(Some(k));
self.a_set_idx[i].push(self.overline_a.len() - 1);
}
for &kk in &g.adj_by_to[i] {
self.pq.update(kk, self.u[i] + g.cost[kk]);
}
}
self.a_set.clear();
for &opt in &self.overline_a {
if let Some(k) = opt {
self.a_set.push(k);
}
}
}
fn load(&mut self) {
let g = self.g;
for k in 0..g.m {
self.link_vol[k] = 0.0;
}
for idx in (0..self.a_set.len()).rev() {
let k = self.a_set[idx];
let i = g.from[k];
let f_i = self.f[i];
let va = if f_i.is_infinite() {
self.node_vol[i]
} else {
let freq = 1.0 / g.head[k];
(freq / f_i) * self.node_vol[i]
};
self.link_vol[k] = va;
self.node_vol[g.to[k]] += va;
}
}
pub fn assign(&mut self, dest_id: usize, demand: &[f64]) -> DestResult<'_> {
self.find_strategy(dest_id);
let mut total = 0.0;
for (id, (nv, &d)) in self.node_vol.iter_mut().zip(demand.iter()).enumerate() {
if id != dest_id && d != 0.0 {
*nv = d;
total += d;
} else {
*nv = 0.0;
}
}
self.node_vol[dest_id] = -total;
self.load();
DestResult {
dest_id,
labels: &self.u,
freqs: &self.f,
a_set: &self.a_set,
link_vol: &self.link_vol,
node_vol: &self.node_vol,
}
}
pub fn solve_each<F: FnMut(&DestResult)>(
&mut self,
od: &HashMap<String, HashMap<String, f64>>,
mut callback: F,
) {
for c in self.cols.iter_mut() {
c.clear();
}
for (origin, row) in od {
let oid = match self.g.node_index(origin) {
Some(id) => id,
None => continue,
};
for (dest, &d) in row {
if d == 0.0 {
continue;
}
if let Some(did) = self.g.node_index(dest) {
self.cols[did].push((oid, d));
}
}
}
let n = self.g.n;
for did in 0..n {
if self.cols[did].is_empty() {
continue;
}
self.find_strategy(did);
let mut total = 0.0;
for id in 0..n {
self.node_vol[id] = 0.0;
}
for idx in 0..self.cols[did].len() {
let (oid, d) = self.cols[did][idx];
if oid != did {
self.node_vol[oid] = d;
total += d;
}
}
self.node_vol[did] = -total;
self.load();
let res = DestResult {
dest_id: did,
labels: &self.u,
freqs: &self.f,
a_set: &self.a_set,
link_vol: &self.link_vol,
node_vol: &self.node_vol,
};
callback(&res);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::spiess_floarian::compute_sf;
use crate::testutil::{gen_grid_network, grid_stops};
#[test]
fn test_solver_parity() {
let (links, nodes, dest, od) = gen_grid_network(4, 4, 6.0, 3.0);
let reference = compute_sf(&links, &nodes, &dest, &od);
let g = Graph::new(&links, &nodes);
let mut w = g.new_workspace();
let dest_id = g.node_index(&dest).unwrap();
let mut demand = vec![0.0; g.num_nodes()];
for (origin, row) in &od {
if let Some(&v) = row.get(&dest) {
demand[g.node_index(origin).unwrap()] = v;
}
}
let got = w.assign(dest_id, &demand);
const EPS: f64 = 1e-9;
for (name, &want) in &reference.strategy.labels {
let id = g.node_index(name).unwrap();
assert!((got.labels[id] - want).abs() < EPS, "label {}", name);
let wf = reference.strategy.freqs[name];
if wf.is_infinite() {
assert!(got.freqs[id].is_infinite());
} else {
assert!((got.freqs[id] - wf).abs() < EPS, "freq {}", name);
}
}
for (k, link) in links.iter().enumerate() {
let want = reference.volumes.links[&link.from_node][&link.to_node];
assert!(
(got.link_vol[k] - want).abs() < EPS,
"linkvol {}->{}",
link.from_node,
link.to_node
);
}
}
#[test]
fn test_solve_each_parity() {
let (links, nodes, _, _) = gen_grid_network(4, 4, 6.0, 3.0);
let stops = grid_stops(4, 4);
let mut od: HashMap<String, HashMap<String, f64>> = HashMap::new();
for o in &stops {
let mut row = HashMap::new();
for d in &stops {
if d != o {
row.insert(d.clone(), 1.0);
}
}
od.insert(o.clone(), row);
}
let mut want_total = 0.0;
for dest in &stops {
let mut col: HashMap<String, HashMap<String, f64>> = HashMap::new();
for o in &stops {
if o != dest {
col.insert(o.clone(), HashMap::from([(dest.clone(), 1.0)]));
}
}
let res = compute_sf(&links, &nodes, dest, &col);
for m in res.volumes.links.values() {
for v in m.values() {
want_total += v;
}
}
}
let g = Graph::new(&links, &nodes);
let mut w = g.new_workspace();
let mut got_total = 0.0;
w.solve_each(&od, |res| {
for &v in res.link_vol {
got_total += v;
}
});
assert!((got_total - want_total).abs() < 1e-6, "{} {}", got_total, want_total);
}
#[test]
fn test_concurrent_shared_graph() {
let (links, nodes, _, _) = gen_grid_network(6, 6, 6.0, 3.0);
let stops = grid_stops(6, 6);
let mut od: HashMap<String, HashMap<String, f64>> = HashMap::new();
for o in &stops {
let mut row = HashMap::new();
for d in &stops {
if d != o {
row.insert(d.clone(), 1.0);
}
}
od.insert(o.clone(), row);
}
let graph = Graph::new(&links, &nodes);
let mut ref_total = 0.0;
{
let mut w = graph.new_workspace();
w.solve_each(&od, |res| {
for &v in res.link_vol {
ref_total += v;
}
});
}
std::thread::scope(|s| {
let handles: Vec<_> = (0..8)
.map(|_| {
let g = &graph;
let od = &od;
s.spawn(move || {
let mut w = g.new_workspace();
let mut total = 0.0;
w.solve_each(od, |res| {
for &v in res.link_vol {
total += v;
}
});
total
})
})
.collect();
for h in handles {
let total = h.join().unwrap();
assert!((total - ref_total).abs() < 1e-6);
}
});
}
}