use super::{FrechetGlobalRegResult, FrechetLocalRegResult};
use crate::error::FdarError;
use crate::frechet::space::{signed_quantile_average, MetricSpace};
use crate::helpers::{gaussian_kernel, NUMERICAL_EPS};
use crate::linalg::{cholesky_factor, cholesky_forward_back, cholesky_solve};
use crate::matrix::FdMatrix;
pub(crate) fn compute_global_weights(
predictors: &FdMatrix,
xout: &FdMatrix,
) -> Result<(Vec<Vec<f64>>, Vec<f64>), FdarError> {
let (n, p) = predictors.shape();
if n == 0 {
return Err(FdarError::InvalidDimension {
parameter: "predictors",
expected: "at least 1 observation".to_string(),
actual: "0 rows".to_string(),
});
}
if xout.ncols() != p {
return Err(FdarError::InvalidDimension {
parameter: "xout",
expected: format!("{p} columns (matching predictors)"),
actual: format!("{} columns", xout.ncols()),
});
}
let n_out = xout.nrows();
let mut x_bar = vec![0.0; p];
for j in 0..p {
let mut s = 0.0;
for i in 0..n {
s += predictors[(i, j)];
}
x_bar[j] = s / n as f64;
}
let denom = if n > 1 { (n - 1) as f64 } else { 1.0 };
let mut sigma = vec![0.0; p * p];
for i in 0..n {
for a in 0..p {
let da = predictors[(i, a)] - x_bar[a];
for b in 0..p {
let db = predictors[(i, b)] - x_bar[b];
sigma[a * p + b] += da * db;
}
}
}
for v in sigma.iter_mut() {
*v /= denom;
}
for j in 0..p {
sigma[j * p + j] += 1e-6;
}
let chol = cholesky_factor(&sigma, p)?;
let mut weights_per_row = Vec::with_capacity(n_out);
for r in 0..n_out {
let diff_x: Vec<f64> = (0..p).map(|j| xout[(r, j)] - x_bar[j]).collect();
let v = cholesky_forward_back(&chol, &diff_x, p); let mut weights = vec![0.0; n];
for i in 0..n {
let mut dot = 0.0;
for j in 0..p {
dot += (predictors[(i, j)] - x_bar[j]) * v[j];
}
weights[i] = (1.0 + dot) / n as f64;
}
weights_per_row.push(weights);
}
Ok((weights_per_row, x_bar))
}
pub(crate) fn compute_local_weights(
predictors: &FdMatrix,
x0: &[f64],
bandwidth: f64,
n: usize,
p: usize,
) -> Result<Vec<f64>, FdarError> {
let mut kern = vec![0.0; n];
for i in 0..n {
let mut k = 1.0;
for j in 0..p {
k *= gaussian_kernel(predictors[(i, j)] - x0[j], bandwidth);
}
kern[i] = k;
}
let mut mu1 = vec![0.0; p];
let mut mu2 = vec![0.0; p * p];
for i in 0..n {
let ki = kern[i];
for a in 0..p {
let da = predictors[(i, a)] - x0[a];
mu1[a] += ki * da;
for b in 0..p {
let db = predictors[(i, b)] - x0[b];
mu2[a * p + b] += ki * da * db;
}
}
}
for v in mu1.iter_mut() {
*v /= n as f64;
}
for v in mu2.iter_mut() {
*v /= n as f64;
}
for j in 0..p {
mu2[j * p + j] += 1e-6;
}
let a_vec = cholesky_solve(&mu2, &mu1, p)?;
let mut weights = vec![0.0; n];
for i in 0..n {
let mut corr = 0.0;
for j in 0..p {
corr += (predictors[(i, j)] - x0[j]) * a_vec[j];
}
weights[i] = kern[i] * (1.0 - corr);
}
let sum_w: f64 = weights.iter().sum();
if sum_w.abs() < NUMERICAL_EPS {
return Err(FdarError::ComputationFailed {
operation: "frechet_local_reg",
detail: "local weights sum to zero (bandwidth too small or no nearby points)"
.to_string(),
});
}
for w in weights.iter_mut() {
*w /= sum_w;
}
Ok(weights)
}
fn validate_reg_input(
predictors: &FdMatrix,
responses: &FdMatrix,
argvals: &[f64],
xout: &FdMatrix,
) -> Result<(usize, usize, usize), FdarError> {
let (n, p) = predictors.shape();
let (nr, m) = responses.shape();
if n == 0 {
return Err(FdarError::InvalidDimension {
parameter: "predictors",
expected: "at least 1 observation".to_string(),
actual: "0 rows".to_string(),
});
}
if nr != n {
return Err(FdarError::InvalidDimension {
parameter: "responses",
expected: format!("{n} rows (matching predictors)"),
actual: format!("{nr} rows"),
});
}
if argvals.len() != m {
return Err(FdarError::InvalidDimension {
parameter: "argvals",
expected: format!("{m} elements (matching responses columns)"),
actual: format!("{} elements", 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(),
});
}
if xout.ncols() != p {
return Err(FdarError::InvalidDimension {
parameter: "xout",
expected: format!("{p} columns (matching predictors)"),
actual: format!("{} columns", xout.ncols()),
});
}
Ok((n, p, m))
}
#[must_use = "expensive regression — store or use the returned prediction"]
pub fn frechet_global_reg(
predictors: &FdMatrix,
responses: &FdMatrix,
argvals: &[f64],
xout: &FdMatrix,
) -> Result<FrechetGlobalRegResult, FdarError> {
let (_n, _p, m) = validate_reg_input(predictors, responses, argvals, xout)?;
let n_out = xout.nrows();
let n_q = m.max(101);
let (weights_per_row, x_bar) = compute_global_weights(predictors, xout)?;
let mut predicted = FdMatrix::zeros(n_out, m);
for (r, weights) in weights_per_row.iter().enumerate() {
let dens = signed_quantile_average(responses, argvals, weights, n_q)?;
for j in 0..m {
predicted[(r, j)] = dens[j];
}
}
Ok(FrechetGlobalRegResult {
predicted,
xout: xout.clone(),
x_bar,
})
}
#[must_use = "expensive regression — store or use the returned predictions"]
pub fn frechet_global_reg_space<S: MetricSpace>(
space: &S,
predictors: &FdMatrix,
responses: &[S::Object],
xout: &FdMatrix,
) -> Result<Vec<S::Object>, FdarError> {
let n = predictors.nrows();
if responses.len() != n {
return Err(FdarError::InvalidDimension {
parameter: "responses",
expected: format!("{n} objects (matching predictor rows)"),
actual: format!("{} objects", responses.len()),
});
}
let (weights_per_row, _x_bar) = compute_global_weights(predictors, xout)?;
let mut out = Vec::with_capacity(weights_per_row.len());
for weights in &weights_per_row {
out.push(space.weighted_frechet_mean(responses, weights)?);
}
Ok(out)
}
#[must_use = "expensive regression — store or use the returned prediction"]
pub fn frechet_local_reg(
predictors: &FdMatrix,
responses: &FdMatrix,
argvals: &[f64],
xout: &FdMatrix,
bandwidth: f64,
) -> Result<FrechetLocalRegResult, FdarError> {
let (n, p, m) = validate_reg_input(predictors, responses, argvals, xout)?;
if bandwidth <= 0.0 || !bandwidth.is_finite() {
return Err(FdarError::InvalidParameter {
parameter: "bandwidth",
message: format!("bandwidth must be positive and finite, got {bandwidth}"),
});
}
let n_out = xout.nrows();
let n_q = m.max(101);
let mut predicted = FdMatrix::zeros(n_out, m);
for r in 0..n_out {
let x0: Vec<f64> = (0..p).map(|j| xout[(r, j)]).collect();
let weights = compute_local_weights(predictors, &x0, bandwidth, n, p)?;
let dens = signed_quantile_average(responses, argvals, &weights, n_q)?;
for j in 0..m {
predicted[(r, j)] = dens[j];
}
}
Ok(FrechetLocalRegResult {
predicted,
xout: xout.clone(),
bandwidth,
})
}
#[must_use = "expensive regression — store or use the returned predictions"]
pub fn frechet_local_reg_space<S: MetricSpace>(
space: &S,
predictors: &FdMatrix,
responses: &[S::Object],
xout: &FdMatrix,
bandwidth: f64,
) -> Result<Vec<S::Object>, FdarError> {
let (n, p) = predictors.shape();
if n == 0 {
return Err(FdarError::InvalidDimension {
parameter: "predictors",
expected: "at least 1 observation".to_string(),
actual: "0 rows".to_string(),
});
}
if responses.len() != n {
return Err(FdarError::InvalidDimension {
parameter: "responses",
expected: format!("{n} objects (matching predictor rows)"),
actual: format!("{} objects", responses.len()),
});
}
if xout.ncols() != p {
return Err(FdarError::InvalidDimension {
parameter: "xout",
expected: format!("{p} columns (matching predictors)"),
actual: format!("{} columns", xout.ncols()),
});
}
if bandwidth <= 0.0 || !bandwidth.is_finite() {
return Err(FdarError::InvalidParameter {
parameter: "bandwidth",
message: format!("bandwidth must be positive and finite, got {bandwidth}"),
});
}
let mut out = Vec::with_capacity(xout.nrows());
for r in 0..xout.nrows() {
let x0: Vec<f64> = (0..p).map(|j| xout[(r, j)]).collect();
let weights = compute_local_weights(predictors, &x0, bandwidth, n, p)?;
out.push(space.weighted_frechet_mean(responses, &weights)?);
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::frechet::space::wasserstein2_distance;
use crate::helpers::trapz;
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 truncated_gaussian(argvals: &[f64], mu: f64, sigma: f64) -> Vec<f64> {
let raw: Vec<f64> = argvals
.iter()
.map(|&x| (-(x - mu).powi(2) / (2.0 * sigma * sigma)).exp())
.collect();
let integral = trapz(&raw, argvals);
raw.iter().map(|&d| d / integral).collect()
}
fn synthetic(n: usize, m: usize, sigma: f64) -> (FdMatrix, FdMatrix, Vec<f64>) {
let argvals = uniform_grid(m, -6.0, 6.0);
let mut predictors = FdMatrix::zeros(n, 1);
let mut responses = FdMatrix::zeros(n, m);
for i in 0..n {
let xi = -1.5 + 3.0 * i as f64 / (n - 1) as f64;
predictors[(i, 0)] = xi;
let dens = truncated_gaussian(&argvals, xi, sigma);
for j in 0..m {
responses[(i, j)] = dens[j];
}
}
(predictors, responses, argvals)
}
fn is_monotone_nondecreasing_cdf(density: &[f64], argvals: &[f64]) -> bool {
density.iter().all(|&v| v >= -1e-9) && (trapz(density, argvals) - 1.0).abs() < 1e-6
}
#[test]
fn global_tracks_known_relationship() {
let (predictors, responses, argvals) = synthetic(31, 101, 1.0);
let mut xout = FdMatrix::zeros(1, 1);
xout[(0, 0)] = 0.5;
let res = frechet_global_reg(&predictors, &responses, &argvals, &xout).unwrap();
assert_eq!(res.predicted.shape(), (1, 101));
let pred: Vec<f64> = (0..101).map(|j| res.predicted[(0, j)]).collect();
let truth = truncated_gaussian(&argvals, 0.5, 1.0);
let w2 = wasserstein2_distance(&pred, &truth, &argvals).unwrap();
assert!(w2 < 0.25, "w2 = {w2}");
}
#[test]
fn global_accepts_negative_weights() {
let (predictors, responses, argvals) = synthetic(31, 101, 1.0);
let mut xout = FdMatrix::zeros(1, 1);
xout[(0, 0)] = 4.0; let res = frechet_global_reg(&predictors, &responses, &argvals, &xout).unwrap();
let pred: Vec<f64> = (0..101).map(|j| res.predicted[(0, j)]).collect();
assert!(is_monotone_nondecreasing_cdf(&pred, &argvals));
}
#[test]
fn global_rejects_bad_input() {
let (predictors, responses, argvals) = synthetic(10, 40, 1.0);
let xout = {
let mut x = FdMatrix::zeros(1, 1);
x[(0, 0)] = 0.0;
x
};
let bad_resp = FdMatrix::zeros(9, 40);
assert!(matches!(
frechet_global_reg(&predictors, &bad_resp, &argvals, &xout).unwrap_err(),
FdarError::InvalidDimension { .. }
));
let mut bad_arg = argvals.clone();
bad_arg[1] = bad_arg[0];
assert!(matches!(
frechet_global_reg(&predictors, &responses, &bad_arg, &xout).unwrap_err(),
FdarError::InvalidParameter { parameter, .. } if parameter == "argvals"
));
}
#[test]
fn local_tracks_known_relationship() {
let (predictors, responses, argvals) = synthetic(31, 101, 1.0);
let mut xout = FdMatrix::zeros(1, 1);
xout[(0, 0)] = 0.0;
let res = frechet_local_reg(&predictors, &responses, &argvals, &xout, 0.6).unwrap();
assert_eq!(res.predicted.shape(), (1, 101));
let pred: Vec<f64> = (0..101).map(|j| res.predicted[(0, j)]).collect();
let truth = truncated_gaussian(&argvals, 0.0, 1.0);
let w2 = wasserstein2_distance(&pred, &truth, &argvals).unwrap();
assert!(w2 < 0.25, "w2 = {w2}");
}
#[test]
fn local_accepts_negative_weights() {
let (predictors, responses, argvals) = synthetic(31, 101, 1.0);
let mut xout = FdMatrix::zeros(1, 1);
xout[(0, 0)] = 1.2; let res = frechet_local_reg(&predictors, &responses, &argvals, &xout, 0.5).unwrap();
let pred: Vec<f64> = (0..101).map(|j| res.predicted[(0, j)]).collect();
assert!(is_monotone_nondecreasing_cdf(&pred, &argvals));
}
#[test]
fn local_rejects_bad_bandwidth() {
let (predictors, responses, argvals) = synthetic(20, 40, 1.0);
let mut xout = FdMatrix::zeros(1, 1);
xout[(0, 0)] = 0.0;
for bad in [0.0, -1.0, f64::NAN, f64::INFINITY] {
assert!(matches!(
frechet_local_reg(&predictors, &responses, &argvals, &xout, bad).unwrap_err(),
FdarError::InvalidParameter { parameter, .. } if parameter == "bandwidth"
));
}
}
use crate::frechet::{SpdMatrixSpace, SpdMetric};
fn spd_constant_setup(n: usize, a: &[f64]) -> (FdMatrix, Vec<Vec<f64>>) {
let mut predictors = FdMatrix::zeros(n, 1);
let mut responses = Vec::with_capacity(n);
for i in 0..n {
predictors[(i, 0)] = -1.5 + 3.0 * i as f64 / (n - 1) as f64;
responses.push(a.to_vec());
}
(predictors, responses)
}
#[test]
fn spd_global_reg_constant_response_predicts_constant() {
let space = SpdMatrixSpace::new(2, SpdMetric::Frobenius).unwrap();
let a = vec![2.0, 0.5, 0.5, 3.0];
let (predictors, responses) = spd_constant_setup(21, &a);
let mut xout = FdMatrix::zeros(1, 1);
xout[(0, 0)] = 0.7;
let preds = frechet_global_reg_space(&space, &predictors, &responses, &xout).unwrap();
assert_eq!(preds.len(), 1);
let d = space.distance(&preds[0], &a).unwrap();
assert!(d < 1e-8, "constant-response prediction off by {d}");
}
#[test]
fn spd_global_reg_rejects_response_count_mismatch() {
let space = SpdMatrixSpace::new(2, SpdMetric::Frobenius).unwrap();
let a = vec![2.0, 0.5, 0.5, 3.0];
let (predictors, mut responses) = spd_constant_setup(10, &a);
responses.pop(); let mut xout = FdMatrix::zeros(1, 1);
xout[(0, 0)] = 0.0;
assert!(matches!(
frechet_global_reg_space(&space, &predictors, &responses, &xout),
Err(FdarError::InvalidDimension {
parameter: "responses",
..
})
));
}
#[test]
fn spd_local_reg_returns_object_per_xout() {
let space = SpdMatrixSpace::new(2, SpdMetric::Frobenius).unwrap();
let a = vec![2.0, 0.5, 0.5, 3.0];
let (predictors, responses) = spd_constant_setup(21, &a);
let mut xout = FdMatrix::zeros(2, 1);
xout[(0, 0)] = -0.3;
xout[(1, 0)] = 0.4;
let preds = frechet_local_reg_space(&space, &predictors, &responses, &xout, 0.6).unwrap();
assert_eq!(preds.len(), 2);
for p in &preds {
assert_eq!(p.len(), 4);
assert!(p.iter().all(|x| x.is_finite()));
}
}
#[test]
fn spd_local_reg_rejects_nonpositive_bandwidth() {
let space = SpdMatrixSpace::new(2, SpdMetric::Frobenius).unwrap();
let a = vec![2.0, 0.5, 0.5, 3.0];
let (predictors, responses) = spd_constant_setup(10, &a);
let mut xout = FdMatrix::zeros(1, 1);
xout[(0, 0)] = 0.0;
assert!(matches!(
frechet_local_reg_space(&space, &predictors, &responses, &xout, 0.0),
Err(FdarError::InvalidParameter {
parameter: "bandwidth",
..
})
));
}
}