Skip to main content

fdars_core/frechet/spaces/
point_process.rs

1//! Point-process (intensity/count) `MetricSpace` backend (FRE-02-05).
2//!
3//! Objects are intensity or count vectors on a shared grid, stored as `Vec<f64>`
4//! of length `m`. The metric is the L2 (Euclidean) distance; the weighted Fréchet
5//! mean is the weighted average. Non-negativity of intensities is not enforced —
6//! supplying valid intensities/counts is the caller's responsibility; only
7//! dimensions are validated.
8//!
9//! # Divergence from R `frechet` 0.3.0
10//!
11//! R's point-process response geometry may use an intensity-transform or
12//! Fisher–Rao metric; this backend uses a plain L2 metric on the intensity
13//! vector. The capability (distance + weighted Fréchet mean over point-process
14//! responses) matches.
15
16use crate::error::FdarError;
17use crate::frechet::MetricSpace;
18use crate::helpers::NUMERICAL_EPS;
19
20/// Point-process response space over length-`m` intensity/count vectors (FRE-02-05).
21#[derive(Debug, Clone, PartialEq)]
22#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
23pub struct PointProcessSpace {
24    /// Grid length `m` (objects are intensity/count vectors of length `m`).
25    pub m: usize,
26}
27
28impl PointProcessSpace {
29    /// Construct a point-process space over length-`m` intensity vectors.
30    ///
31    /// # Errors
32    /// [`FdarError::InvalidParameter`] if `m < 1`.
33    pub fn new(m: usize) -> Result<Self, FdarError> {
34        if m < 1 {
35            return Err(FdarError::InvalidParameter {
36                parameter: "m",
37                message: "intensity grid length must be >= 1".to_string(),
38            });
39        }
40        Ok(Self { m })
41    }
42
43    fn check_len(&self, obj: &[f64], name: &'static str) -> Result<(), FdarError> {
44        if obj.len() != self.m {
45            return Err(FdarError::InvalidDimension {
46                parameter: name,
47                expected: format!("{} elements", self.m),
48                actual: format!("{} elements", obj.len()),
49            });
50        }
51        Ok(())
52    }
53}
54
55impl MetricSpace for PointProcessSpace {
56    type Object = Vec<f64>;
57
58    fn distance(&self, a: &Self::Object, b: &Self::Object) -> Result<f64, FdarError> {
59        self.check_len(a, "a")?;
60        self.check_len(b, "b")?;
61        Ok(a.iter()
62            .zip(b.iter())
63            .map(|(x, y)| (x - y) * (x - y))
64            .sum::<f64>()
65            .sqrt())
66    }
67
68    fn weighted_frechet_mean(
69        &self,
70        objects: &[Self::Object],
71        weights: &[f64],
72    ) -> Result<Self::Object, FdarError> {
73        if objects.is_empty() {
74            return Err(FdarError::InvalidDimension {
75                parameter: "objects",
76                expected: "at least 1 object".to_string(),
77                actual: "0 objects".to_string(),
78            });
79        }
80        if weights.len() != objects.len() {
81            return Err(FdarError::InvalidDimension {
82                parameter: "weights",
83                expected: format!("{} weights (matching objects)", objects.len()),
84                actual: format!("{} weights", weights.len()),
85            });
86        }
87        for (i, o) in objects.iter().enumerate() {
88            if o.len() != self.m {
89                return Err(FdarError::InvalidDimension {
90                    parameter: "objects",
91                    expected: format!("each object has {} elements", self.m),
92                    actual: format!("object {i} has {} elements", o.len()),
93                });
94            }
95        }
96        let sw: f64 = weights.iter().sum();
97        if sw.abs() < NUMERICAL_EPS {
98            return Err(FdarError::ComputationFailed {
99                operation: "PointProcessSpace::weighted_frechet_mean",
100                detail: "sum of weights is ~0; cannot normalize the barycenter".to_string(),
101            });
102        }
103        let mut m = vec![0.0f64; self.m];
104        for (o, &w) in objects.iter().zip(weights.iter()) {
105            for (k, mk) in m.iter_mut().enumerate() {
106                *mk += w * o[k];
107            }
108        }
109        for x in &mut m {
110            *x /= sw;
111        }
112        Ok(m)
113    }
114}
115
116#[cfg(test)]
117mod tests {
118    use super::*;
119
120    #[test]
121    fn point_process_distance_of_identical_is_zero() {
122        let s = PointProcessSpace::new(3).unwrap();
123        let a = vec![0.5, 1.5, 2.0];
124        assert!(s.distance(&a, &a).unwrap() < 1e-12);
125    }
126
127    #[test]
128    fn point_process_distance_orthonormal_is_sqrt2() {
129        let s = PointProcessSpace::new(3).unwrap();
130        let a = vec![1.0, 0.0, 0.0];
131        let b = vec![0.0, 1.0, 0.0];
132        assert!((s.distance(&a, &b).unwrap() - 2f64.sqrt()).abs() < 1e-12);
133    }
134
135    #[test]
136    fn point_process_mean_of_identical_recovers() {
137        let s = PointProcessSpace::new(3).unwrap();
138        let a = vec![0.5, 1.5, 2.0];
139        let m = s
140            .weighted_frechet_mean(&[a.clone(), a.clone()], &[0.3, 0.7])
141            .unwrap();
142        for (x, y) in m.iter().zip(a.iter()) {
143            assert!((x - y).abs() < 1e-10);
144        }
145    }
146
147    #[test]
148    fn point_process_rejects_dimension_mismatch() {
149        let s = PointProcessSpace::new(3).unwrap();
150        let a = vec![1.0, 0.0, 0.0];
151        let bad = vec![1.0, 0.0];
152        assert!(matches!(
153            s.distance(&a, &bad),
154            Err(FdarError::InvalidDimension { .. })
155        ));
156    }
157}