fdars_core/frechet/spaces/
point_process.rs1use crate::error::FdarError;
17use crate::frechet::MetricSpace;
18use crate::helpers::NUMERICAL_EPS;
19
20#[derive(Debug, Clone, PartialEq)]
22#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
23pub struct PointProcessSpace {
24 pub m: usize,
26}
27
28impl PointProcessSpace {
29 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}