Skip to main content

fdars_core/frechet/
space.rs

1//! Metric-space abstraction and the 1D-Wasserstein (density-response) backend.
2//!
3//! Provides the [`MetricSpace`] trait (a distance + weighted-Fréchet-mean solver)
4//! and its first concrete implementation, [`WassersteinDensitySpace`], whose
5//! objects are probability densities on a shared strictly-increasing grid and
6//! whose metric is the 1D 2-Wasserstein distance ([`wasserstein2_distance`]).
7//!
8//! The density backend reuses DENS-01's quantile/Wasserstein machinery
9//! ([`crate::density_fda::wasserstein_barycenter`] and its density→quantile→
10//! density back-map) rather than re-deriving it.
11
12use 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
17/// A metric space: a distance function plus a weighted-Fréchet-mean solver over
18/// its objects. Regression / statistics routines are generic over this trait.
19///
20/// Implementors must be `Send + Sync` so the statistics routines can parallelize.
21/// `weighted_frechet_mean` expects **non-negative** weights (they are normalized
22/// to sum to 1 by callers); the signed-weight regression path uses the private
23/// [`signed_quantile_average`] helper instead, never `weighted_frechet_mean`.
24pub trait MetricSpace: Send + Sync {
25    /// The object type living in this metric space (e.g. a density on a grid).
26    type Object;
27
28    /// Distance between two objects.
29    ///
30    /// # Errors
31    /// Returns [`FdarError`] on dimension mismatch or degenerate input.
32    fn distance(&self, a: &Self::Object, b: &Self::Object) -> Result<f64, FdarError>;
33
34    /// Weighted Fréchet mean (barycenter) of `objects` under non-negative `weights`.
35    ///
36    /// # Errors
37    /// Returns [`FdarError`] on empty input, weight/length mismatch, or a
38    /// degenerate barycenter.
39    fn weighted_frechet_mean(
40        &self,
41        objects: &[Self::Object],
42        weights: &[f64],
43    ) -> Result<Self::Object, FdarError>;
44}
45
46/// The 1D-Wasserstein (density-response) metric space.
47///
48/// Objects are probability densities sampled on the shared strictly-increasing
49/// grid `argvals`; the metric is the 1D 2-Wasserstein distance and the weighted
50/// Fréchet mean is the Wasserstein barycenter (quantile average).
51#[derive(Debug, Clone, PartialEq)]
52#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
53pub struct WassersteinDensitySpace {
54    /// Shared strictly-increasing evaluation grid for all density objects.
55    pub argvals: Vec<f64>,
56}
57
58impl WassersteinDensitySpace {
59    /// Construct a density space over a strictly-increasing `argvals` grid.
60    ///
61    /// # Errors
62    /// Returns [`FdarError::InvalidDimension`] if `argvals` has fewer than 2
63    /// points, or [`FdarError::InvalidParameter`] if it is not strictly
64    /// increasing.
65    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        // Non-negative-weight sample barycenter — reuse DENS-01's solver.
118        // (Signed-weight regression uses `signed_quantile_average` instead.)
119        wasserstein_barycenter(&mat, &self.argvals, Some(weights))
120    }
121}
122
123/// The 1D 2-Wasserstein distance between two densities on a shared grid.
124///
125/// Computed as the L² distance between quantile functions,
126/// `W₂(F,G) = (∫₀¹ (Q_F(t) − Q_G(t))² dt)^{1/2}`, reusing the density→CDF→quantile
127/// machinery of [`crate::density_fda`]. This is the metric behind
128/// [`MetricSpace::distance`] for [`WassersteinDensitySpace`].
129///
130/// # Errors
131/// Returns [`FdarError::InvalidDimension`] if `a`, `b`, and `argvals` lengths
132/// differ or `argvals` has fewer than 2 points.
133#[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/// Quantile function `Q(t)` of a density `row` on `argvals`, evaluated at `t_grid`.
163///
164/// Normalizes the density to integrate to 1, forms its CDF via
165/// [`cumulative_trapz`], and inverts by interpolating the CDF at each probability
166/// `t` — replicating the density→quantile step of
167/// [`crate::density_fda::wasserstein_barycenter`].
168#[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
184/// Signed weighted quantile average → density, with a sort-based monotone
185/// (isotonic) projection. **Reserved for the signed-weight regression path**
186/// (global/local Fréchet regression); the non-negative-weight sample Fréchet mean
187/// ([`WassersteinDensitySpace::weighted_frechet_mean`]) uses
188/// [`crate::density_fda::wasserstein_barycenter`] instead and never calls this.
189///
190/// Computes `Q̄(t) = Σᵢ wᵢ · Qᵢ(t)` with possibly-**negative** weights (so it does
191/// NOT call `wasserstein_barycenter`, which rejects negative weights), then sorts
192/// `Q̄` to restore monotonicity and inverts it back to a density on `argvals`
193/// using the same back-map as the Wasserstein barycenter.
194///
195/// # Divergence from R `frechet`
196///
197/// R's `GloWassReg`/`LocWassReg` enforce a monotone quantile via an `osqp`
198/// quadratic-program projection. To avoid a new crate dependency this uses a
199/// sort-based isotonic projection — equivalent on smooth quantile averages,
200/// slightly more conservative on non-smooth ones.
201///
202/// # Errors
203/// Returns [`FdarError`] on dimension mismatch or a degenerate (zero-range)
204/// quantile average.
205pub(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    // Signed weighted average of quantile functions.
235    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    // Sort-based monotone (isotonic) projection — the no-osqp alternative.
246    q_bar.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
247
248    // Clamp the averaged quantile to the target support so the inverted density
249    // stays on `argvals` (signed extrapolation weights can push Q̄ slightly past
250    // the grid edges).
251    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    // Invert Q̄ → density directly in x-units (Q̄ is a weighted average of quantile
265    // functions, already on the argvals x-scale — no rescale-to-full-support, which
266    // would spuriously stretch a narrow barycenter across the whole grid).
267    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        // For a location family, W₂ between N(0,1) and N(δ,1) equals δ.
326        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        // The mean is the Wasserstein barycenter; its density→quantile→density
353        // reconstruction (reused from DENS-01) has an inherent ~0.1 W₂ round-trip
354        // floor, so recovery is within that documented tolerance, not machine eps.
355        let w2 = wasserstein2_distance(&mean, &d, &argvals).unwrap();
356        assert!(w2 < 0.15, "w2 = {w2}");
357        // Exact agreement with DENS-01's Wasserstein mean (same underlying call).
358        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        // With uniform non-negative weights, the signed quantile average of two
389        // unit Gaussians at ±1 is the true Wasserstein barycenter N(0,1) (the
390        // location-family barycenter is the mean-location Gaussian).
391        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        // Negative weights must NOT error (the whole point of this helper).
410        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]; // sums to 1, has negatives
421        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}