use crate::density_fda::{dedup_adjacent, quantile_density_from_q, wasserstein_barycenter};
use crate::error::FdarError;
use crate::helpers::{cumulative_trapz, linear_interp, trapz};
use crate::matrix::FdMatrix;
pub trait MetricSpace: Send + Sync {
type Object;
fn distance(&self, a: &Self::Object, b: &Self::Object) -> Result<f64, FdarError>;
fn weighted_frechet_mean(
&self,
objects: &[Self::Object],
weights: &[f64],
) -> Result<Self::Object, FdarError>;
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct WassersteinDensitySpace {
pub argvals: Vec<f64>,
}
impl WassersteinDensitySpace {
pub fn new(argvals: Vec<f64>) -> Result<Self, FdarError> {
if argvals.len() < 2 {
return Err(FdarError::InvalidDimension {
parameter: "argvals",
expected: "at least 2 grid points".to_string(),
actual: format!("{} points", argvals.len()),
});
}
if argvals.windows(2).any(|w| w[1] <= w[0]) {
return Err(FdarError::InvalidParameter {
parameter: "argvals",
message: "argvals must be strictly increasing".to_string(),
});
}
Ok(Self { argvals })
}
}
impl MetricSpace for WassersteinDensitySpace {
type Object = Vec<f64>;
fn distance(&self, a: &Self::Object, b: &Self::Object) -> Result<f64, FdarError> {
wasserstein2_distance(a, b, &self.argvals)
}
fn weighted_frechet_mean(
&self,
objects: &[Self::Object],
weights: &[f64],
) -> Result<Self::Object, FdarError> {
let m = self.argvals.len();
if objects.is_empty() {
return Err(FdarError::InvalidDimension {
parameter: "objects",
expected: "at least 1 object".to_string(),
actual: "0 objects".to_string(),
});
}
let n = objects.len();
let mut mat = FdMatrix::zeros(n, m);
for (i, obj) in objects.iter().enumerate() {
if obj.len() != m {
return Err(FdarError::InvalidDimension {
parameter: "objects",
expected: format!("each object has {m} points"),
actual: format!("object {i} has {} points", obj.len()),
});
}
for j in 0..m {
mat[(i, j)] = obj[j];
}
}
wasserstein_barycenter(&mat, &self.argvals, Some(weights))
}
}
#[must_use = "returns the 2-Wasserstein distance; result should be examined"]
pub fn wasserstein2_distance(a: &[f64], b: &[f64], argvals: &[f64]) -> Result<f64, FdarError> {
let m = argvals.len();
if m < 2 {
return Err(FdarError::InvalidDimension {
parameter: "argvals",
expected: "at least 2 grid points".to_string(),
actual: format!("{m} points"),
});
}
if a.len() != m || b.len() != m {
return Err(FdarError::InvalidDimension {
parameter: "a/b",
expected: format!("both length {m} (matching argvals)"),
actual: format!("a={}, b={}", a.len(), b.len()),
});
}
let n_q = m.max(101);
let t_grid: Vec<f64> = (0..n_q).map(|i| i as f64 / (n_q - 1) as f64).collect();
let qa = density_to_quantile(a, argvals, &t_grid);
let qb = density_to_quantile(b, argvals, &t_grid);
let sq_diff: Vec<f64> = qa
.iter()
.zip(qb.iter())
.map(|(&x, &y)| (x - y) * (x - y))
.collect();
Ok(trapz(&sq_diff, &t_grid).sqrt())
}
#[inline]
fn density_to_quantile(row: &[f64], argvals: &[f64], t_grid: &[f64]) -> Vec<f64> {
let integral = trapz(row, argvals);
let inv = if integral.abs() < 1e-300 {
1.0
} else {
1.0 / integral
};
let norm: Vec<f64> = row.iter().map(|&v| v * inv).collect();
let cdf = cumulative_trapz(&norm, argvals);
t_grid
.iter()
.map(|&t| linear_interp(&cdf, argvals, t))
.collect()
}
pub(crate) fn signed_quantile_average(
density_matrix: &FdMatrix,
argvals: &[f64],
weights: &[f64],
n_q: usize,
) -> Result<Vec<f64>, FdarError> {
let (n, m) = density_matrix.shape();
if m != argvals.len() {
return Err(FdarError::InvalidDimension {
parameter: "density_matrix",
expected: format!("{} columns (matching argvals)", argvals.len()),
actual: format!("{m} columns"),
});
}
if weights.len() != n {
return Err(FdarError::InvalidDimension {
parameter: "weights",
expected: format!("{n} weights (matching rows)"),
actual: format!("{} weights", weights.len()),
});
}
if n_q < 2 {
return Err(FdarError::InvalidParameter {
parameter: "n_q",
message: "n_q must be at least 2".to_string(),
});
}
let t_grid: Vec<f64> = (0..n_q).map(|i| i as f64 / (n_q - 1) as f64).collect();
let mut q_bar = vec![0.0_f64; n_q];
for i in 0..n {
let row: Vec<f64> = (0..m).map(|j| density_matrix[(i, j)]).collect();
let qi = density_to_quantile(&row, argvals, &t_grid);
let wi = weights[i];
for j in 0..n_q {
q_bar[j] += wi * qi[j];
}
}
q_bar.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let lb = argvals[0];
let ub = argvals[m - 1];
for v in q_bar.iter_mut() {
*v = v.clamp(lb, ub);
}
let q_range = q_bar[n_q - 1] - q_bar[0];
if q_range < 1e-15 {
return Err(FdarError::ComputationFailed {
operation: "signed_quantile_average",
detail: "quantile average has zero range; degenerate weighted input".to_string(),
});
}
let dens_raw = quantile_density_from_q(&q_bar, &t_grid);
let (q_dedup, dens_dedup) = dedup_adjacent(&q_bar, &dens_raw);
let dens: Vec<f64> = argvals
.iter()
.map(|&x| linear_interp(&q_dedup, &dens_dedup, x))
.collect();
let integral = trapz(&dens, argvals);
if integral < 1e-15 {
return Err(FdarError::ComputationFailed {
operation: "signed_quantile_average",
detail: "reconstructed density integrates to zero".to_string(),
});
}
Ok(dens.iter().map(|&d| d / integral).collect())
}
#[cfg(test)]
mod tests {
use super::*;
fn uniform_grid(m: usize, lb: f64, ub: f64) -> Vec<f64> {
(0..m)
.map(|j| lb + (ub - lb) * j as f64 / (m - 1) as f64)
.collect()
}
fn gaussian(argvals: &[f64], mu: f64) -> Vec<f64> {
let raw: Vec<f64> = argvals
.iter()
.map(|&x| (-(x - mu).powi(2) / 2.0).exp())
.collect();
let integral = trapz(&raw, argvals);
raw.iter().map(|&d| d / integral).collect()
}
#[test]
fn space_new_validates_grid() {
assert!(WassersteinDensitySpace::new(uniform_grid(50, -5.0, 5.0)).is_ok());
assert!(matches!(
WassersteinDensitySpace::new(vec![0.0, 1.0, 0.5]).unwrap_err(),
FdarError::InvalidParameter { parameter, .. } if parameter == "argvals"
));
assert!(matches!(
WassersteinDensitySpace::new(vec![0.0]).unwrap_err(),
FdarError::InvalidDimension { .. }
));
}
#[test]
fn w2_identical_is_zero() {
let argvals = uniform_grid(101, -5.0, 5.0);
let d = gaussian(&argvals, 0.0);
let w2 = wasserstein2_distance(&d, &d, &argvals).unwrap();
assert!(w2 < 1e-8, "w2 = {w2}");
}
#[test]
fn w2_matches_location_shift() {
let argvals = uniform_grid(201, -8.0, 8.0);
let d0 = gaussian(&argvals, 0.0);
let d1 = gaussian(&argvals, 0.5);
let w2 = wasserstein2_distance(&d0, &d1, &argvals).unwrap();
assert!((w2 - 0.5).abs() < 0.05, "w2 = {w2}");
}
#[test]
fn distance_delegates_to_w2() {
let argvals = uniform_grid(101, -5.0, 5.0);
let space = WassersteinDensitySpace::new(argvals.clone()).unwrap();
let d0 = gaussian(&argvals, 0.0);
let d1 = gaussian(&argvals, 0.3);
let via_trait = space.distance(&d0, &d1).unwrap();
let direct = wasserstein2_distance(&d0, &d1, &argvals).unwrap();
assert!((via_trait - direct).abs() < 1e-12);
}
#[test]
fn weighted_frechet_mean_of_identical_recovers_object() {
let argvals = uniform_grid(101, -5.0, 5.0);
let space = WassersteinDensitySpace::new(argvals.clone()).unwrap();
let d = gaussian(&argvals, 0.0);
let objects = vec![d.clone(), d.clone(), d.clone()];
let weights = vec![1.0 / 3.0; 3];
let mean = space.weighted_frechet_mean(&objects, &weights).unwrap();
let w2 = wasserstein2_distance(&mean, &d, &argvals).unwrap();
assert!(w2 < 0.15, "w2 = {w2}");
let bary = wasserstein_barycenter(
&{
let mut m = FdMatrix::zeros(3, argvals.len());
for i in 0..3 {
for j in 0..argvals.len() {
m[(i, j)] = d[j];
}
}
m
},
&argvals,
Some(&weights),
)
.unwrap();
assert_eq!(mean, bary);
}
#[test]
fn w2_rejects_length_mismatch() {
let argvals = uniform_grid(50, -5.0, 5.0);
let a = vec![0.0; 50];
let b = vec![0.0; 49];
assert!(matches!(
wasserstein2_distance(&a, &b, &argvals).unwrap_err(),
FdarError::InvalidDimension { .. }
));
}
#[test]
fn signed_quantile_average_uniform_weights_recovers_true_barycenter() {
let argvals = uniform_grid(101, -5.0, 5.0);
let d0 = gaussian(&argvals, -1.0);
let d1 = gaussian(&argvals, 1.0);
let mut mat = FdMatrix::zeros(2, argvals.len());
for j in 0..argvals.len() {
mat[(0, j)] = d0[j];
mat[(1, j)] = d1[j];
}
let w = vec![0.5, 0.5];
let n_q = argvals.len().max(101);
let signed = signed_quantile_average(&mat, &argvals, &w, n_q).unwrap();
let truth = gaussian(&argvals, 0.0);
let diff = wasserstein2_distance(&signed, &truth, &argvals).unwrap();
assert!(diff < 0.15, "diff = {diff}");
}
#[test]
fn signed_quantile_average_accepts_negative_weights() {
let argvals = uniform_grid(101, -6.0, 6.0);
let d0 = gaussian(&argvals, -1.0);
let d1 = gaussian(&argvals, 0.0);
let d2 = gaussian(&argvals, 1.0);
let mut mat = FdMatrix::zeros(3, argvals.len());
for j in 0..argvals.len() {
mat[(0, j)] = d0[j];
mat[(1, j)] = d1[j];
mat[(2, j)] = d2[j];
}
let w = vec![-0.2, 1.4, -0.2]; let n_q = argvals.len().max(101);
let res = signed_quantile_average(&mat, &argvals, &w, n_q).unwrap();
assert_eq!(res.len(), argvals.len());
assert!(res.iter().all(|v| v.is_finite() && *v >= -1e-9));
}
}