use crate::distances::Distances;
use crate::field::{Fp, MODULUS_LIMIT, Z2, is_prime};
use crate::reduce::{Engine, RawH1Class};
use crate::{Diagram, Error, Result, RipsParams};
pub(crate) fn compute<D: Distances + Sync>(dist: &D, params: &RipsParams) -> Result<Diagram> {
compute_in(dist, params, None)
}
pub(crate) fn compute_in<D: Distances + Sync>(
dist: &D,
params: &RipsParams,
pool: Option<rayon::ThreadPool>,
) -> Result<Diagram> {
let p = validate_params(params)?;
if p == 2 {
compute_impl(dist, Z2, params, pool)
} else {
compute_impl(dist, Fp::new(p), params, pool)
}
}
fn compute_impl<C, D>(
dist: &D,
ops: C,
params: &RipsParams,
pool: Option<rayon::ThreadPool>,
) -> Result<Diagram>
where
C: crate::field::Coeffs + Sync,
D: Distances + Sync,
{
let mut diagram = Diagram::default();
if dist.len() == 0 {
return Ok(diagram);
}
let engine = match pool {
Some(_) => Engine::new_in(dist, params, ops, pool)?,
None => Engine::new(dist, params, ops)?,
};
engine.run(&mut diagram);
diagram.canonicalize();
Ok(diagram)
}
pub(crate) fn compute_with_h1_classes<D: Distances + Sync>(
dist: &D,
params: &RipsParams,
) -> Result<(Diagram, Vec<RawH1Class>)> {
let p = validate_params(params)?;
if p == 2 {
compute_with_h1_impl(dist, Z2, params)
} else {
compute_with_h1_impl(dist, Fp::new(p), params)
}
}
fn validate_params(params: &RipsParams) -> Result<u64> {
if let Some(t) = params.threshold {
if t.is_nan() || t < 0.0 {
return Err(Error::InvalidInput(format!(
"threshold must be non-negative, got {t}"
)));
}
}
let p = params.modulus as u64;
if !is_prime(p) || p >= MODULUS_LIMIT {
return Err(Error::InvalidInput(format!(
"modulus must be a prime below {MODULUS_LIMIT}, got {p}"
)));
}
Ok(p)
}
fn compute_with_h1_impl<C, D>(
dist: &D,
ops: C,
params: &RipsParams,
) -> Result<(Diagram, Vec<RawH1Class>)>
where
C: crate::field::Coeffs + Sync,
D: Distances + Sync,
{
let mut diagram = Diagram::default();
let mut classes = Vec::new();
if dist.len() == 0 {
return Ok((diagram, classes));
}
let engine = Engine::new(dist, params, ops)?;
engine.run_with_h1_classes(&mut diagram, &mut classes);
diagram.canonicalize();
Ok((diagram, classes))
}
#[cfg(test)]
mod tests {
use crate::{DistanceMatrix, RipsParams, rips_persistence};
#[test]
fn triangle_unit_distances() {
let dist = DistanceMatrix::from_condensed(vec![1.0, 1.0, 1.0]).unwrap();
let d = rips_persistence(&dist, &RipsParams::new(1)).unwrap();
let h0: Vec<_> = d.in_dim(0).collect();
assert_eq!(h0.len(), 3);
assert_eq!(h0.iter().filter(|b| b.is_essential()).count(), 1);
assert_eq!(h0.iter().filter(|b| b.death == 1.0).count(), 2);
assert_eq!(d.in_dim(1).count(), 0);
}
#[test]
fn square_has_one_h1_bar() {
let s = 2.0f64.sqrt();
let dist = DistanceMatrix::from_condensed(vec![1.0, s, 1.0, 1.0, s, 1.0]).unwrap();
let d = rips_persistence(&dist, &RipsParams::new(1).with_threshold(2.0)).unwrap();
assert_eq!(d.in_dim(0).filter(|b| b.is_essential()).count(), 1);
assert_eq!(d.in_dim(0).count(), 4);
let h1: Vec<_> = d.in_dim(1).collect();
assert_eq!(h1.len(), 1);
assert_eq!(h1[0].birth, 1.0);
assert_eq!(h1[0].death, s);
}
#[test]
fn two_components() {
let inf = f64::INFINITY;
let dist = DistanceMatrix::from_condensed(vec![1.0, inf, inf, inf, inf, 1.0]).unwrap();
let d = rips_persistence(&dist, &RipsParams::new(1).with_threshold(10.0)).unwrap();
assert_eq!(d.in_dim(0).filter(|b| b.is_essential()).count(), 2);
}
#[test]
fn square_diagram_is_field_independent() {
let s = 2.0f64.sqrt();
let dist = DistanceMatrix::from_condensed(vec![1.0, s, 1.0, 1.0, s, 1.0]).unwrap();
let base = rips_persistence(&dist, &RipsParams::new(1)).unwrap();
for p in [3u32, 5, 7, 13] {
let mut params = RipsParams::new(1);
params.modulus = p;
let d = rips_persistence(&dist, ¶ms).unwrap();
assert_eq!(d.bars, base.bars, "diagram differs at p = {p}");
}
}
#[test]
fn invalid_modulus_is_rejected() {
let dist = DistanceMatrix::from_condensed(vec![1.0]).unwrap();
for bad in [0u32, 1, 4, 6, 9, 32768, 32770] {
let mut params = RipsParams::new(1);
params.modulus = bad;
assert!(
rips_persistence(&dist, ¶ms).is_err(),
"modulus {bad} must be rejected"
);
}
}
}