use graphops::{bfs_distances, GraphRef};
use ndarray::{Array1, Array2};
#[derive(Debug, Clone, Copy)]
pub struct CurvatureConfig {
pub alpha: f64,
pub reg: f32,
pub max_iter: usize,
}
impl Default for CurvatureConfig {
fn default() -> Self {
Self {
alpha: 0.0,
reg: 0.01,
max_iter: 500,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct EdgeCurvature {
pub u: usize,
pub v: usize,
pub kappa: f32,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum CurvatureError {
NotSquare,
NotSymmetric,
InvalidWeight,
InvalidAlpha,
}
impl std::fmt::Display for CurvatureError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::NotSquare => write!(f, "adjacency matrix must be square"),
Self::NotSymmetric => write!(f, "adjacency matrix must be symmetric (undirected)"),
Self::InvalidWeight => write!(f, "edge weights must be finite and non-negative"),
Self::InvalidAlpha => write!(f, "alpha must be in [0, 1]"),
}
}
}
impl std::error::Error for CurvatureError {}
struct NeighborLists(Vec<Vec<usize>>);
impl GraphRef for NeighborLists {
fn node_count(&self) -> usize {
self.0.len()
}
fn neighbors_ref(&self, node: usize) -> &[usize] {
&self.0[node]
}
}
pub fn ollivier_ricci_curvatures(
adj: &Array2<f64>,
config: &CurvatureConfig,
) -> Result<Vec<EdgeCurvature>, CurvatureError> {
let n = adj.nrows();
if adj.ncols() != n {
return Err(CurvatureError::NotSquare);
}
if !(0.0..=1.0).contains(&config.alpha) {
return Err(CurvatureError::InvalidAlpha);
}
for i in 0..n {
for j in 0..n {
let w = adj[[i, j]];
if !w.is_finite() || w < 0.0 {
return Err(CurvatureError::InvalidWeight);
}
if (w - adj[[j, i]]).abs() > 1e-9 * w.abs().max(1.0) {
return Err(CurvatureError::NotSymmetric);
}
}
}
let transition = lapl::transition_matrix(adj);
let neighbors = NeighborLists(
(0..n)
.map(|i| (0..n).filter(|&j| adj[[i, j]] > 0.0).collect())
.collect(),
);
let hops: Vec<Vec<Option<usize>>> = (0..n).map(|i| bfs_distances(&neighbors, i)).collect();
let mut out = Vec::new();
for u in 0..n {
for v in (u + 1)..n {
if adj[[u, v]] <= 0.0 {
continue;
}
let mu_u = lazy_measure(&transition, u, config.alpha);
let mu_v = lazy_measure(&transition, v, config.alpha);
let support: Vec<usize> = (0..n).filter(|&i| mu_u[i] > 0.0 || mu_v[i] > 0.0).collect();
let m = support.len();
let a = Array1::from_iter(support.iter().map(|&i| mu_u[i] as f32));
let b = Array1::from_iter(support.iter().map(|&i| mu_v[i] as f32));
let mut cost = Array2::zeros((m, m));
for (si, &i) in support.iter().enumerate() {
for (sj, &j) in support.iter().enumerate() {
cost[[si, sj]] = hops[i][j].unwrap_or(usize::MAX) as f32;
}
}
let (_, w1) = wass::sinkhorn_log(&a, &b, &cost, config.reg, config.max_iter);
out.push(EdgeCurvature {
u,
v,
kappa: 1.0 - w1,
});
}
}
Ok(out)
}
fn lazy_measure(transition: &Array2<f64>, x: usize, alpha: f64) -> Array1<f64> {
let mut mu = transition.row(x).to_owned() * (1.0 - alpha);
mu[x] += alpha;
mu
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::array;
fn kappa_of(kappas: &[EdgeCurvature], u: usize, v: usize) -> f32 {
kappas
.iter()
.find(|e| (e.u, e.v) == (u.min(v), u.max(v)))
.expect("edge present")
.kappa
}
#[test]
fn triangle_edges_curve_positively() {
let adj = array![[0., 1., 1.], [1., 0., 1.], [1., 1., 0.]];
let kappas = ollivier_ricci_curvatures(&adj, &CurvatureConfig::default()).unwrap();
assert_eq!(kappas.len(), 3);
for e in &kappas {
assert!((e.kappa - 0.5).abs() < 0.05, "kappa {}", e.kappa);
}
}
#[test]
fn four_cycle_edges_are_flat() {
let adj = array![
[0., 1., 0., 1.],
[1., 0., 1., 0.],
[0., 1., 0., 1.],
[1., 0., 1., 0.]
];
let kappas = ollivier_ricci_curvatures(&adj, &CurvatureConfig::default()).unwrap();
assert_eq!(kappas.len(), 4);
for e in &kappas {
assert!(e.kappa.abs() < 0.05, "kappa {}", e.kappa);
}
}
#[test]
fn double_star_bridge_curves_negatively() {
let mut adj = Array2::zeros((6, 6));
for &(i, j) in &[(1usize, 2usize), (1, 0), (1, 4), (2, 3), (2, 5)] {
adj[[i, j]] = 1.0;
adj[[j, i]] = 1.0;
}
let kappas = ollivier_ricci_curvatures(&adj, &CurvatureConfig::default()).unwrap();
let bridge = kappa_of(&kappas, 1, 2);
assert!((bridge - (-2.0 / 3.0)).abs() < 0.05, "bridge {bridge}");
for e in &kappas {
if (e.u, e.v) != (1, 2) {
assert!(e.kappa > bridge, "bridge should be the most negative");
}
}
}
#[test]
fn full_laziness_zeroes_curvature() {
let adj = array![[0., 1., 1.], [1., 0., 1.], [1., 1., 0.]];
let cfg = CurvatureConfig {
alpha: 1.0,
..CurvatureConfig::default()
};
for e in &ollivier_ricci_curvatures(&adj, &cfg).unwrap() {
assert!(e.kappa.abs() < 0.05, "kappa {}", e.kappa);
}
}
#[test]
fn rejects_bad_inputs() {
let rect = Array2::<f64>::zeros((2, 3));
assert_eq!(
ollivier_ricci_curvatures(&rect, &CurvatureConfig::default()),
Err(CurvatureError::NotSquare)
);
let asym = array![[0., 1.], [0., 0.]];
assert_eq!(
ollivier_ricci_curvatures(&asym, &CurvatureConfig::default()),
Err(CurvatureError::NotSymmetric)
);
let neg = array![[0., -1.], [-1., 0.]];
assert_eq!(
ollivier_ricci_curvatures(&neg, &CurvatureConfig::default()),
Err(CurvatureError::InvalidWeight)
);
let adj = array![[0., 1.], [1., 0.]];
let bad_alpha = CurvatureConfig {
alpha: 1.5,
..CurvatureConfig::default()
};
assert_eq!(
ollivier_ricci_curvatures(&adj, &bad_alpha),
Err(CurvatureError::InvalidAlpha)
);
}
}