Skip to main content

holos_tda/atlas/
point.rs

1use crate::{Error, PointCloudGraph, PointCloudParams, Result, RipsParams, SparseDistanceMatrix};
2
3use super::model::{
4    AtlasEvaluation, EdgeKey, EndpointGradient, LineageId, PersistenceAtlas, UpdateMode,
5};
6
7/// One coordinate derivative of a point-distance endpoint.
8#[derive(Debug, Clone, PartialEq)]
9pub struct CoordinateDerivative {
10    /// Point index.
11    pub point: usize,
12    /// Coordinate index.
13    pub coordinate: usize,
14    /// Analytic derivative of the Euclidean distance before `f64` rounding.
15    pub value: f64,
16}
17
18/// Point-coordinate derivatives for one finite barcode endpoint.
19#[derive(Debug, Clone, PartialEq)]
20pub struct PointEndpointGradient {
21    /// Edge whose distance controls the endpoint.
22    pub edge: EdgeKey,
23    /// Nonzero derivatives for both endpoints of the edge.
24    pub terms: Vec<CoordinateDerivative>,
25}
26
27/// Point-coordinate sensitivity of one persistent class space.
28#[derive(Debug, Clone, PartialEq)]
29pub struct PointClassSensitivity {
30    /// Atlas lineage.
31    pub lineage: LineageId,
32    /// Birth derivative. `None` when a distance tie prevents one gradient.
33    pub birth: Option<PointEndpointGradient>,
34    /// Death derivative. `None` for an essential death or a distance tie.
35    pub death: Option<PointEndpointGradient>,
36}
37
38/// Exact sparse point-cloud atlas with a conservative coordinate radius.
39#[derive(Debug, Clone)]
40pub struct PointPersistenceAtlas {
41    points: Vec<Vec<f64>>,
42    threshold: f64,
43    atlas: PersistenceAtlas,
44    coordinate_radius: f64,
45}
46
47impl PointPersistenceAtlas {
48    /// Build a finite-threshold point-cloud atlas.
49    pub fn build(points: &[Vec<f64>], params: &RipsParams) -> Result<Self> {
50        let threshold = params.threshold.ok_or_else(|| {
51            Error::InvalidInput("a point persistence atlas requires an explicit threshold".into())
52        })?;
53        if !threshold.is_finite() || threshold < 0.0 {
54            return Err(Error::InvalidInput(
55                "a point persistence atlas requires a finite non-negative threshold".into(),
56            ));
57        }
58        let graph = PointCloudGraph::build(
59            points,
60            PointCloudParams::new(threshold).with_threads(params.threads),
61        )?;
62        let atlas = PersistenceAtlas::build(graph.matrix(), params)?;
63        let coordinate_radius = coordinate_radius(points, threshold)?;
64        Ok(Self {
65            points: points.to_vec(),
66            threshold,
67            atlas,
68            coordinate_radius,
69        })
70    }
71
72    /// Conservative per-point Euclidean radius. Within it, pairwise distance
73    /// relations and threshold memberships stay fixed. A zero radius forces a
74    /// rebuild for every changed cloud. Pairwise distance overflow sets it to
75    /// zero.
76    pub fn coordinate_radius(&self) -> f64 {
77        self.coordinate_radius
78    }
79
80    /// Underlying edge-weight atlas.
81    pub fn edge_atlas(&self) -> &PersistenceAtlas {
82        &self.atlas
83    }
84
85    /// Analytic point-coordinate gradients at the compiled point cloud.
86    pub fn sensitivities(&self) -> Vec<PointClassSensitivity> {
87        let evaluation = self
88            .atlas
89            .evaluate(&self.original_graph())
90            .expect("point atlas contains its own valid graph");
91        evaluation
92            .sensitivities
93            .into_iter()
94            .map(|sensitivity| PointClassSensitivity {
95                lineage: sensitivity.lineage,
96                birth: point_gradient(&self.points, sensitivity.birth),
97                death: point_gradient(&self.points, sensitivity.death),
98            })
99            .collect()
100    }
101
102    /// Evaluate new coordinates when the displacement radius holds.
103    pub fn evaluate(&self, points: &[Vec<f64>]) -> Result<AtlasEvaluation> {
104        let displacement = point_displacement(&self.points, points)?;
105        let unchanged = self
106            .points
107            .iter()
108            .zip(points)
109            .all(|(a, b)| a.iter().zip(b).all(|(a, b)| a.to_bits() == b.to_bits()));
110        if !unchanged && (self.coordinate_radius == 0.0 || displacement >= self.coordinate_radius) {
111            return Err(Error::InvalidInput(format!(
112                "point displacement {displacement} reaches atlas radius {}",
113                self.coordinate_radius
114            )));
115        }
116        let dense = crate::DistanceMatrix::from_points(points)?;
117        let triplets: Vec<_> = self
118            .atlas
119            .topology
120            .iter()
121            .map(|edge| (edge.u, edge.v, dense.get(edge.u, edge.v)))
122            .collect();
123        let updated = SparseDistanceMatrix::from_triplets(points.len(), &triplets)?;
124        self.atlas.evaluate(&updated)
125    }
126
127    /// Reuse the atlas when the displacement is inside the current radius.
128    /// Any other valid point set rebuilds it.
129    pub fn update(&self, points: &[Vec<f64>]) -> Result<PointAtlasUpdate> {
130        if let Ok(evaluation) = self.evaluate(points) {
131            return Ok(PointAtlasUpdate {
132                atlas: self.clone(),
133                evaluation,
134                mode: UpdateMode::Reused,
135            });
136        }
137        let mut params = self.atlas.params.clone();
138        params.threshold = Some(self.threshold);
139        let atlas = Self::build(points, &params)?;
140        let evaluation = atlas.evaluate(points)?;
141        Ok(PointAtlasUpdate {
142            atlas,
143            evaluation,
144            mode: UpdateMode::Recomputed,
145        })
146    }
147
148    fn original_graph(&self) -> SparseDistanceMatrix {
149        let dense = crate::DistanceMatrix::from_points(&self.points)
150            .expect("point atlas contains validated finite points");
151        let triplets: Vec<_> = self
152            .atlas
153            .topology
154            .iter()
155            .map(|edge| (edge.u, edge.v, dense.get(edge.u, edge.v)))
156            .collect();
157        SparseDistanceMatrix::from_triplets(self.points.len(), &triplets)
158            .expect("point atlas topology is canonical")
159    }
160}
161
162/// Result of applying a new point cloud to a point atlas.
163#[derive(Debug, Clone)]
164pub struct PointAtlasUpdate {
165    /// Atlas valid at the new coordinates.
166    pub atlas: PointPersistenceAtlas,
167    /// Exact result at the new coordinates.
168    pub evaluation: AtlasEvaluation,
169    /// How the atlas was updated.
170    pub mode: UpdateMode,
171}
172
173fn coordinate_radius(points: &[Vec<f64>], threshold: f64) -> Result<f64> {
174    let dense = crate::DistanceMatrix::from_points(points)?;
175    let mut distances = Vec::new();
176    let mut threshold_gap = f64::INFINITY;
177    for v in 1..points.len() {
178        for u in 0..v {
179            let distance = dense.get(u, v);
180            if !distance.is_finite() {
181                return Ok(0.0);
182            }
183            distances.push(distance);
184            threshold_gap = threshold_gap.min((distance - threshold).abs());
185        }
186    }
187    distances.sort_by(f64::total_cmp);
188    let mut order_gap = f64::INFINITY;
189    for pair in distances.windows(2) {
190        let gap = pair[1] - pair[0];
191        if gap == 0.0 {
192            return Ok(0.0);
193        }
194        order_gap = order_gap.min(gap);
195    }
196    let radius = (order_gap / 4.0).min(threshold_gap / 2.0);
197    Ok(if radius.is_nan() { 0.0 } else { radius })
198}
199
200fn point_displacement(original: &[Vec<f64>], updated: &[Vec<f64>]) -> Result<f64> {
201    if original.len() != updated.len() {
202        return Err(Error::InvalidInput(format!(
203            "point count changed from {} to {}",
204            original.len(),
205            updated.len()
206        )));
207    }
208    let mut maximum = 0.0f64;
209    for (index, (a, b)) in original.iter().zip(updated).enumerate() {
210        if a.len() != b.len() {
211            return Err(Error::InvalidInput(format!(
212                "point {index} changed dimension from {} to {}",
213                a.len(),
214                b.len()
215            )));
216        }
217        for &b in b {
218            if !b.is_finite() {
219                return Err(Error::InvalidInput(format!(
220                    "point {index} has a non-finite coordinate"
221                )));
222            }
223        }
224        maximum = maximum.max(scaled_difference_norm(a, b));
225    }
226    Ok(maximum)
227}
228
229fn point_gradient(
230    points: &[Vec<f64>],
231    gradient: EndpointGradient,
232) -> Option<PointEndpointGradient> {
233    let EndpointGradient::Edge(edge) = gradient else {
234        return None;
235    };
236    let distance = scaled_difference_norm(&points[edge.u], &points[edge.v]);
237    if distance == 0.0 || !distance.is_finite() {
238        return None;
239    }
240    let mut terms = Vec::new();
241    for (coordinate, (&a, &b)) in points[edge.u].iter().zip(&points[edge.v]).enumerate() {
242        let derivative = (a - b) / distance;
243        if derivative != 0.0 {
244            terms.push(CoordinateDerivative {
245                point: edge.u,
246                coordinate,
247                value: derivative,
248            });
249            terms.push(CoordinateDerivative {
250                point: edge.v,
251                coordinate,
252                value: -derivative,
253            });
254        }
255    }
256    Some(PointEndpointGradient { edge, terms })
257}
258
259pub(crate) fn scaled_difference_norm(a: &[f64], b: &[f64]) -> f64 {
260    let mut scale = 0.0f64;
261    let mut sum = 1.0f64;
262    for (&a, &b) in a.iter().zip(b) {
263        let difference = (a - b).abs();
264        if difference == 0.0 {
265            continue;
266        }
267        if !difference.is_finite() {
268            return f64::INFINITY;
269        }
270        if scale < difference {
271            let ratio = scale / difference;
272            sum = 1.0 + sum * ratio * ratio;
273            scale = difference;
274        } else {
275            let ratio = difference / scale;
276            sum += ratio * ratio;
277        }
278    }
279    if scale == 0.0 {
280        0.0
281    } else {
282        scale * sum.sqrt()
283    }
284}