1use crate::density_fda::{dedup_adjacent, quantile_density_from_q, wasserstein_barycenter};
13use crate::error::FdarError;
14use crate::helpers::{cumulative_trapz, linear_interp, trapz};
15use crate::matrix::FdMatrix;
16
17pub trait MetricSpace: Send + Sync {
25 type Object;
27
28 fn distance(&self, a: &Self::Object, b: &Self::Object) -> Result<f64, FdarError>;
33
34 fn weighted_frechet_mean(
40 &self,
41 objects: &[Self::Object],
42 weights: &[f64],
43 ) -> Result<Self::Object, FdarError>;
44}
45
46#[derive(Debug, Clone, PartialEq)]
52#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
53pub struct WassersteinDensitySpace {
54 pub argvals: Vec<f64>,
56}
57
58impl WassersteinDensitySpace {
59 pub fn new(argvals: Vec<f64>) -> Result<Self, FdarError> {
66 if argvals.len() < 2 {
67 return Err(FdarError::InvalidDimension {
68 parameter: "argvals",
69 expected: "at least 2 grid points".to_string(),
70 actual: format!("{} points", argvals.len()),
71 });
72 }
73 if argvals.windows(2).any(|w| w[1] <= w[0]) {
74 return Err(FdarError::InvalidParameter {
75 parameter: "argvals",
76 message: "argvals must be strictly increasing".to_string(),
77 });
78 }
79 Ok(Self { argvals })
80 }
81}
82
83impl MetricSpace for WassersteinDensitySpace {
84 type Object = Vec<f64>;
85
86 fn distance(&self, a: &Self::Object, b: &Self::Object) -> Result<f64, FdarError> {
87 wasserstein2_distance(a, b, &self.argvals)
88 }
89
90 fn weighted_frechet_mean(
91 &self,
92 objects: &[Self::Object],
93 weights: &[f64],
94 ) -> Result<Self::Object, FdarError> {
95 let m = self.argvals.len();
96 if objects.is_empty() {
97 return Err(FdarError::InvalidDimension {
98 parameter: "objects",
99 expected: "at least 1 object".to_string(),
100 actual: "0 objects".to_string(),
101 });
102 }
103 let n = objects.len();
104 let mut mat = FdMatrix::zeros(n, m);
105 for (i, obj) in objects.iter().enumerate() {
106 if obj.len() != m {
107 return Err(FdarError::InvalidDimension {
108 parameter: "objects",
109 expected: format!("each object has {m} points"),
110 actual: format!("object {i} has {} points", obj.len()),
111 });
112 }
113 for j in 0..m {
114 mat[(i, j)] = obj[j];
115 }
116 }
117 wasserstein_barycenter(&mat, &self.argvals, Some(weights))
120 }
121}
122
123#[must_use = "returns the 2-Wasserstein distance; result should be examined"]
134pub fn wasserstein2_distance(a: &[f64], b: &[f64], argvals: &[f64]) -> Result<f64, FdarError> {
135 let m = argvals.len();
136 if m < 2 {
137 return Err(FdarError::InvalidDimension {
138 parameter: "argvals",
139 expected: "at least 2 grid points".to_string(),
140 actual: format!("{m} points"),
141 });
142 }
143 if a.len() != m || b.len() != m {
144 return Err(FdarError::InvalidDimension {
145 parameter: "a/b",
146 expected: format!("both length {m} (matching argvals)"),
147 actual: format!("a={}, b={}", a.len(), b.len()),
148 });
149 }
150 let n_q = m.max(101);
151 let t_grid: Vec<f64> = (0..n_q).map(|i| i as f64 / (n_q - 1) as f64).collect();
152 let qa = density_to_quantile(a, argvals, &t_grid);
153 let qb = density_to_quantile(b, argvals, &t_grid);
154 let sq_diff: Vec<f64> = qa
155 .iter()
156 .zip(qb.iter())
157 .map(|(&x, &y)| (x - y) * (x - y))
158 .collect();
159 Ok(trapz(&sq_diff, &t_grid).sqrt())
160}
161
162#[inline]
169fn density_to_quantile(row: &[f64], argvals: &[f64], t_grid: &[f64]) -> Vec<f64> {
170 let integral = trapz(row, argvals);
171 let inv = if integral.abs() < 1e-300 {
172 1.0
173 } else {
174 1.0 / integral
175 };
176 let norm: Vec<f64> = row.iter().map(|&v| v * inv).collect();
177 let cdf = cumulative_trapz(&norm, argvals);
178 t_grid
179 .iter()
180 .map(|&t| linear_interp(&cdf, argvals, t))
181 .collect()
182}
183
184pub(crate) fn signed_quantile_average(
206 density_matrix: &FdMatrix,
207 argvals: &[f64],
208 weights: &[f64],
209 n_q: usize,
210) -> Result<Vec<f64>, FdarError> {
211 let (n, m) = density_matrix.shape();
212 if m != argvals.len() {
213 return Err(FdarError::InvalidDimension {
214 parameter: "density_matrix",
215 expected: format!("{} columns (matching argvals)", argvals.len()),
216 actual: format!("{m} columns"),
217 });
218 }
219 if weights.len() != n {
220 return Err(FdarError::InvalidDimension {
221 parameter: "weights",
222 expected: format!("{n} weights (matching rows)"),
223 actual: format!("{} weights", weights.len()),
224 });
225 }
226 if n_q < 2 {
227 return Err(FdarError::InvalidParameter {
228 parameter: "n_q",
229 message: "n_q must be at least 2".to_string(),
230 });
231 }
232 let t_grid: Vec<f64> = (0..n_q).map(|i| i as f64 / (n_q - 1) as f64).collect();
233
234 let mut q_bar = vec![0.0_f64; n_q];
236 for i in 0..n {
237 let row: Vec<f64> = (0..m).map(|j| density_matrix[(i, j)]).collect();
238 let qi = density_to_quantile(&row, argvals, &t_grid);
239 let wi = weights[i];
240 for j in 0..n_q {
241 q_bar[j] += wi * qi[j];
242 }
243 }
244
245 q_bar.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
247
248 let lb = argvals[0];
252 let ub = argvals[m - 1];
253 for v in q_bar.iter_mut() {
254 *v = v.clamp(lb, ub);
255 }
256 let q_range = q_bar[n_q - 1] - q_bar[0];
257 if q_range < 1e-15 {
258 return Err(FdarError::ComputationFailed {
259 operation: "signed_quantile_average",
260 detail: "quantile average has zero range; degenerate weighted input".to_string(),
261 });
262 }
263
264 let dens_raw = quantile_density_from_q(&q_bar, &t_grid);
268 let (q_dedup, dens_dedup) = dedup_adjacent(&q_bar, &dens_raw);
269 let dens: Vec<f64> = argvals
270 .iter()
271 .map(|&x| linear_interp(&q_dedup, &dens_dedup, x))
272 .collect();
273 let integral = trapz(&dens, argvals);
274 if integral < 1e-15 {
275 return Err(FdarError::ComputationFailed {
276 operation: "signed_quantile_average",
277 detail: "reconstructed density integrates to zero".to_string(),
278 });
279 }
280 Ok(dens.iter().map(|&d| d / integral).collect())
281}
282
283#[cfg(test)]
284mod tests {
285 use super::*;
286
287 fn uniform_grid(m: usize, lb: f64, ub: f64) -> Vec<f64> {
288 (0..m)
289 .map(|j| lb + (ub - lb) * j as f64 / (m - 1) as f64)
290 .collect()
291 }
292
293 fn gaussian(argvals: &[f64], mu: f64) -> Vec<f64> {
294 let raw: Vec<f64> = argvals
295 .iter()
296 .map(|&x| (-(x - mu).powi(2) / 2.0).exp())
297 .collect();
298 let integral = trapz(&raw, argvals);
299 raw.iter().map(|&d| d / integral).collect()
300 }
301
302 #[test]
303 fn space_new_validates_grid() {
304 assert!(WassersteinDensitySpace::new(uniform_grid(50, -5.0, 5.0)).is_ok());
305 assert!(matches!(
306 WassersteinDensitySpace::new(vec![0.0, 1.0, 0.5]).unwrap_err(),
307 FdarError::InvalidParameter { parameter, .. } if parameter == "argvals"
308 ));
309 assert!(matches!(
310 WassersteinDensitySpace::new(vec![0.0]).unwrap_err(),
311 FdarError::InvalidDimension { .. }
312 ));
313 }
314
315 #[test]
316 fn w2_identical_is_zero() {
317 let argvals = uniform_grid(101, -5.0, 5.0);
318 let d = gaussian(&argvals, 0.0);
319 let w2 = wasserstein2_distance(&d, &d, &argvals).unwrap();
320 assert!(w2 < 1e-8, "w2 = {w2}");
321 }
322
323 #[test]
324 fn w2_matches_location_shift() {
325 let argvals = uniform_grid(201, -8.0, 8.0);
327 let d0 = gaussian(&argvals, 0.0);
328 let d1 = gaussian(&argvals, 0.5);
329 let w2 = wasserstein2_distance(&d0, &d1, &argvals).unwrap();
330 assert!((w2 - 0.5).abs() < 0.05, "w2 = {w2}");
331 }
332
333 #[test]
334 fn distance_delegates_to_w2() {
335 let argvals = uniform_grid(101, -5.0, 5.0);
336 let space = WassersteinDensitySpace::new(argvals.clone()).unwrap();
337 let d0 = gaussian(&argvals, 0.0);
338 let d1 = gaussian(&argvals, 0.3);
339 let via_trait = space.distance(&d0, &d1).unwrap();
340 let direct = wasserstein2_distance(&d0, &d1, &argvals).unwrap();
341 assert!((via_trait - direct).abs() < 1e-12);
342 }
343
344 #[test]
345 fn weighted_frechet_mean_of_identical_recovers_object() {
346 let argvals = uniform_grid(101, -5.0, 5.0);
347 let space = WassersteinDensitySpace::new(argvals.clone()).unwrap();
348 let d = gaussian(&argvals, 0.0);
349 let objects = vec![d.clone(), d.clone(), d.clone()];
350 let weights = vec![1.0 / 3.0; 3];
351 let mean = space.weighted_frechet_mean(&objects, &weights).unwrap();
352 let w2 = wasserstein2_distance(&mean, &d, &argvals).unwrap();
356 assert!(w2 < 0.15, "w2 = {w2}");
357 let bary = wasserstein_barycenter(
359 &{
360 let mut m = FdMatrix::zeros(3, argvals.len());
361 for i in 0..3 {
362 for j in 0..argvals.len() {
363 m[(i, j)] = d[j];
364 }
365 }
366 m
367 },
368 &argvals,
369 Some(&weights),
370 )
371 .unwrap();
372 assert_eq!(mean, bary);
373 }
374
375 #[test]
376 fn w2_rejects_length_mismatch() {
377 let argvals = uniform_grid(50, -5.0, 5.0);
378 let a = vec![0.0; 50];
379 let b = vec![0.0; 49];
380 assert!(matches!(
381 wasserstein2_distance(&a, &b, &argvals).unwrap_err(),
382 FdarError::InvalidDimension { .. }
383 ));
384 }
385
386 #[test]
387 fn signed_quantile_average_uniform_weights_recovers_true_barycenter() {
388 let argvals = uniform_grid(101, -5.0, 5.0);
392 let d0 = gaussian(&argvals, -1.0);
393 let d1 = gaussian(&argvals, 1.0);
394 let mut mat = FdMatrix::zeros(2, argvals.len());
395 for j in 0..argvals.len() {
396 mat[(0, j)] = d0[j];
397 mat[(1, j)] = d1[j];
398 }
399 let w = vec![0.5, 0.5];
400 let n_q = argvals.len().max(101);
401 let signed = signed_quantile_average(&mat, &argvals, &w, n_q).unwrap();
402 let truth = gaussian(&argvals, 0.0);
403 let diff = wasserstein2_distance(&signed, &truth, &argvals).unwrap();
404 assert!(diff < 0.15, "diff = {diff}");
405 }
406
407 #[test]
408 fn signed_quantile_average_accepts_negative_weights() {
409 let argvals = uniform_grid(101, -6.0, 6.0);
411 let d0 = gaussian(&argvals, -1.0);
412 let d1 = gaussian(&argvals, 0.0);
413 let d2 = gaussian(&argvals, 1.0);
414 let mut mat = FdMatrix::zeros(3, argvals.len());
415 for j in 0..argvals.len() {
416 mat[(0, j)] = d0[j];
417 mat[(1, j)] = d1[j];
418 mat[(2, j)] = d2[j];
419 }
420 let w = vec![-0.2, 1.4, -0.2]; let n_q = argvals.len().max(101);
422 let res = signed_quantile_average(&mat, &argvals, &w, n_q).unwrap();
423 assert_eq!(res.len(), argvals.len());
424 assert!(res.iter().all(|v| v.is_finite() && *v >= -1e-9));
425 }
426}