use crate::autodiff::Scalar;
fn generic_linear_interp<S: Scalar>(argvals: &[f64], curve: &[S], t: f64) -> S {
let m = argvals.len();
if m == 0 {
return S::zero();
}
if t <= argvals[0] {
return curve[0];
}
if t >= argvals[m - 1] {
return curve[m - 1];
}
let j = argvals
.partition_point(|&a| a <= t)
.saturating_sub(1)
.min(m - 2);
let dt = argvals[j + 1] - argvals[j];
if dt <= 0.0 {
return curve[j];
}
let alpha = S::from_f64((t - argvals[j]) / dt);
curve[j] * (S::one() - alpha) + curve[j + 1] * alpha
}
fn generic_srsf_central_diff<S: Scalar>(curve: &[S], argvals: &[f64]) -> Vec<S> {
let m = curve.len();
let h = if m > 1 {
(argvals[m - 1] - argvals[0]) / (m - 1) as f64
} else {
1.0
};
let inv_h = S::from_f64(1.0 / h);
let inv_2h = S::from_f64(1.0 / (2.0 * h));
let mut q = vec![S::zero(); m];
for (j, qj) in q.iter_mut().enumerate() {
let deriv = if m == 1 {
S::zero()
} else if j == 0 {
(curve[1] - curve[0]) * inv_h
} else if j == m - 1 {
(curve[m - 1] - curve[m - 2]) * inv_h
} else {
(curve[j + 1] - curve[j - 1]) * inv_2h
};
*qj = S::signum(deriv) * S::sqrt(S::abs(deriv));
}
q
}
fn generic_l2_srsf_distance<S: Scalar>(q1: &[f64], q2: &[S], weights: &[f64]) -> S {
let mut dist_sq = S::zero();
for j in 0..q1.len() {
let diff = S::from_f64(q1[j]) - q2[j];
dist_sq += diff * diff * S::from_f64(weights[j]);
}
S::sqrt(dist_sq)
}
pub fn amplitude_distance_at_warp_generic<S: Scalar>(
q1_ref: &[f64],
curve2: &[S],
warping: &[f64],
argvals: &[f64],
weights: &[f64],
) -> S {
let m = argvals.len();
let f2_aligned: Vec<S> = (0..m)
.map(|j| generic_linear_interp(argvals, curve2, warping[j]))
.collect();
let q2_aligned = generic_srsf_central_diff(&f2_aligned, argvals);
generic_l2_srsf_distance(q1_ref, &q2_aligned, weights)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::autodiff::Dual;
use crate::helpers::simpsons_weights;
use std::f64::consts::PI;
fn setup() -> (Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>, Vec<f64>) {
let n = 20;
let argvals: Vec<f64> = (0..n)
.map(|i| 0.1 + 0.8 * i as f64 / (n - 1) as f64)
.collect();
let f1: Vec<f64> = argvals.iter().map(|&t| (2.0 * PI * t).sin()).collect();
let q1_ref = srsf_central_diff_f64(&f1, &argvals);
let curve2: Vec<f64> = argvals
.iter()
.map(|&t| (2.0 * PI * t).cos() + 0.1)
.collect();
let warping = argvals.clone();
let weights = simpsons_weights(&argvals);
(argvals, q1_ref, curve2, warping, weights)
}
fn srsf_central_diff_f64(curve: &[f64], argvals: &[f64]) -> Vec<f64> {
generic_srsf_central_diff(curve, argvals)
}
#[test]
fn amplitude_f64_parity() {
let (argvals, q1_ref, curve2, warping, weights) = setup();
let generic = amplitude_distance_at_warp_generic::<f64>(
&q1_ref, &curve2, &warping, &argvals, &weights,
);
let m = argvals.len();
let f2_aligned: Vec<f64> = (0..m)
.map(|j| generic_linear_interp(&argvals, &curve2, warping[j]))
.collect();
let q2 = srsf_central_diff_f64(&f2_aligned, &argvals);
let mut dist_sq = 0.0;
for j in 0..m {
let d = q1_ref[j] - q2[j];
dist_sq += d * d * weights[j];
}
let reference = dist_sq.sqrt();
assert!(
(generic - reference).abs() <= 1e-10,
"generic {generic} vs reference {reference}"
);
}
#[test]
fn amplitude_gradient_vs_fd() {
let (argvals, q1_ref, curve2, warping, weights) = setup();
let m = curve2.len();
let h = 1e-8;
for j in 0..m {
let c2_dual: Vec<Dual> = curve2
.iter()
.enumerate()
.map(|(i, &v)| {
if i == j {
Dual::seed(v)
} else {
Dual::constant(v)
}
})
.collect();
let dual =
amplitude_distance_at_warp_generic(&q1_ref, &c2_dual, &warping, &argvals, &weights)
.extract()
.1;
let mut cp = curve2.clone();
let mut cm = curve2.clone();
cp[j] += h;
cm[j] -= h;
let fp = amplitude_distance_at_warp_generic::<f64>(
&q1_ref, &cp, &warping, &argvals, &weights,
);
let fm = amplitude_distance_at_warp_generic::<f64>(
&q1_ref, &cm, &warping, &argvals, &weights,
);
let fd = (fp - fm) / (2.0 * h);
assert!((dual - fd).abs() <= 1e-5, "j={j}: dual {dual} vs fd {fd}");
}
}
}