1use crate::{Error, PointCloudGraph, PointCloudParams, Result, RipsParams, SparseDistanceMatrix};
2
3use super::model::{
4 AtlasEvaluation, EdgeKey, EndpointGradient, LineageId, PersistenceAtlas, UpdateMode,
5};
6
7#[derive(Debug, Clone, PartialEq)]
9pub struct CoordinateDerivative {
10 pub point: usize,
12 pub coordinate: usize,
14 pub value: f64,
16}
17
18#[derive(Debug, Clone, PartialEq)]
20pub struct PointEndpointGradient {
21 pub edge: EdgeKey,
23 pub terms: Vec<CoordinateDerivative>,
25}
26
27#[derive(Debug, Clone, PartialEq)]
29pub struct PointClassSensitivity {
30 pub lineage: LineageId,
32 pub birth: Option<PointEndpointGradient>,
34 pub death: Option<PointEndpointGradient>,
36}
37
38#[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 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 pub fn coordinate_radius(&self) -> f64 {
77 self.coordinate_radius
78 }
79
80 pub fn edge_atlas(&self) -> &PersistenceAtlas {
82 &self.atlas
83 }
84
85 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 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 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, ¶ms)?;
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#[derive(Debug, Clone)]
164pub struct PointAtlasUpdate {
165 pub atlas: PointPersistenceAtlas,
167 pub evaluation: AtlasEvaluation,
169 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}