use std::fmt::Debug;
use std::ops::{Add, AddAssign, Div, Mul, MulAssign, Neg, Sub, SubAssign};
pub trait Scalar:
Copy
+ Clone
+ Debug
+ PartialOrd
+ Add<Output = Self>
+ Sub<Output = Self>
+ Mul<Output = Self>
+ Div<Output = Self>
+ Neg<Output = Self>
+ AddAssign
+ SubAssign
+ MulAssign
{
fn zero() -> Self;
fn one() -> Self;
fn from_f64(v: f64) -> Self;
fn infinity() -> Self;
fn sqrt(self) -> Self;
fn exp(self) -> Self;
fn ln(self) -> Self;
fn sin(self) -> Self;
fn cos(self) -> Self;
fn powf(self, p: f64) -> Self;
fn abs(self) -> Self;
fn signum(self) -> Self;
}
impl Scalar for f64 {
#[inline]
fn zero() -> Self {
0.0
}
#[inline]
fn one() -> Self {
1.0
}
#[inline]
fn from_f64(v: f64) -> Self {
v
}
#[inline]
fn infinity() -> Self {
f64::INFINITY
}
#[inline]
fn sqrt(self) -> Self {
f64::sqrt(self)
}
#[inline]
fn exp(self) -> Self {
f64::exp(self)
}
#[inline]
fn ln(self) -> Self {
f64::ln(self)
}
#[inline]
fn sin(self) -> Self {
f64::sin(self)
}
#[inline]
fn cos(self) -> Self {
f64::cos(self)
}
#[inline]
fn powf(self, p: f64) -> Self {
f64::powf(self, p)
}
#[inline]
fn abs(self) -> Self {
f64::abs(self)
}
#[inline]
fn signum(self) -> Self {
f64::signum(self)
}
}
#[derive(Debug, Clone, Copy)]
pub struct Dual {
pub value: f64,
pub tangent: f64,
}
impl Dual {
#[inline]
#[must_use]
pub fn seed(value: f64) -> Self {
Dual {
value,
tangent: 1.0,
}
}
#[inline]
#[must_use]
pub fn constant(value: f64) -> Self {
Dual {
value,
tangent: 0.0,
}
}
#[inline]
#[must_use]
pub fn extract(self) -> (f64, f64) {
(self.value, self.tangent)
}
}
impl Add for Dual {
type Output = Self;
#[inline]
fn add(self, rhs: Self) -> Self {
Dual {
value: self.value + rhs.value,
tangent: self.tangent + rhs.tangent,
}
}
}
impl Sub for Dual {
type Output = Self;
#[inline]
fn sub(self, rhs: Self) -> Self {
Dual {
value: self.value - rhs.value,
tangent: self.tangent - rhs.tangent,
}
}
}
impl Mul for Dual {
type Output = Self;
#[inline]
fn mul(self, rhs: Self) -> Self {
Dual {
value: self.value * rhs.value,
tangent: self.tangent * rhs.value + self.value * rhs.tangent,
}
}
}
impl Div for Dual {
type Output = Self;
#[inline]
fn div(self, rhs: Self) -> Self {
let v2 = rhs.value * rhs.value;
Dual {
value: self.value / rhs.value,
tangent: (self.tangent * rhs.value - self.value * rhs.tangent) / v2,
}
}
}
impl Neg for Dual {
type Output = Self;
#[inline]
fn neg(self) -> Self {
Dual {
value: -self.value,
tangent: -self.tangent,
}
}
}
impl AddAssign for Dual {
#[inline]
fn add_assign(&mut self, rhs: Self) {
*self = *self + rhs;
}
}
impl SubAssign for Dual {
#[inline]
fn sub_assign(&mut self, rhs: Self) {
*self = *self - rhs;
}
}
impl MulAssign for Dual {
#[inline]
fn mul_assign(&mut self, rhs: Self) {
*self = *self * rhs;
}
}
impl PartialEq for Dual {
#[inline]
fn eq(&self, other: &Self) -> bool {
self.value == other.value
}
}
impl PartialOrd for Dual {
#[inline]
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
self.value.partial_cmp(&other.value)
}
}
impl Scalar for Dual {
#[inline]
fn zero() -> Self {
Dual {
value: 0.0,
tangent: 0.0,
}
}
#[inline]
fn one() -> Self {
Dual {
value: 1.0,
tangent: 0.0,
}
}
#[inline]
fn from_f64(v: f64) -> Self {
Dual {
value: v,
tangent: 0.0,
}
}
#[inline]
fn infinity() -> Self {
Dual {
value: f64::INFINITY,
tangent: 0.0,
}
}
#[inline]
fn sqrt(self) -> Self {
let s = self.value.sqrt();
Dual {
value: s,
tangent: self.tangent / (2.0 * s),
}
}
#[inline]
fn exp(self) -> Self {
let e = self.value.exp();
Dual {
value: e,
tangent: self.tangent * e,
}
}
#[inline]
fn ln(self) -> Self {
Dual {
value: self.value.ln(),
tangent: self.tangent / self.value,
}
}
#[inline]
fn sin(self) -> Self {
Dual {
value: self.value.sin(),
tangent: self.tangent * self.value.cos(),
}
}
#[inline]
fn cos(self) -> Self {
Dual {
value: self.value.cos(),
tangent: -self.tangent * self.value.sin(),
}
}
#[inline]
fn powf(self, p: f64) -> Self {
Dual {
value: self.value.powf(p),
tangent: self.tangent * p * self.value.powf(p - 1.0),
}
}
#[inline]
fn abs(self) -> Self {
let sub = if self.value == 0.0 {
0.0
} else {
self.value.signum()
};
Dual {
value: self.value.abs(),
tangent: self.tangent * sub,
}
}
#[inline]
fn signum(self) -> Self {
Dual {
value: self.value.signum(),
tangent: 0.0,
}
}
}
#[must_use]
pub fn diff<F: Fn(Dual) -> Dual>(f: F, x: f64) -> (f64, f64) {
f(Dual::seed(x)).extract()
}
#[must_use]
pub fn grad<F: Fn(&[Dual]) -> Dual>(f: F, x: &[f64]) -> (f64, Vec<f64>) {
let m = x.len();
if m == 0 {
return (f(&[]).value, Vec::new());
}
let mut gradient = vec![0.0; m];
let mut value = 0.0;
for k in 0..m {
let duals: Vec<Dual> = (0..m)
.map(|j| {
if j == k {
Dual::seed(x[j])
} else {
Dual::constant(x[j])
}
})
.collect();
let (v, t) = f(&duals).extract();
if k == 0 {
value = v;
}
gradient[k] = t;
}
(value, gradient)
}
#[must_use]
pub fn jacobian<F: Fn(&[Dual]) -> Vec<Dual>>(f: F, x: &[f64]) -> (Vec<f64>, Vec<Vec<f64>>) {
let m = x.len();
if m == 0 {
let outputs = f(&[]);
let values: Vec<f64> = outputs.iter().map(|d| d.value).collect();
let rows = values.len();
return (values, vec![Vec::new(); rows]);
}
let mut values: Vec<f64> = Vec::new();
let mut jac: Vec<Vec<f64>> = Vec::new();
for k in 0..m {
let duals: Vec<Dual> = (0..m)
.map(|j| {
if j == k {
Dual::seed(x[j])
} else {
Dual::constant(x[j])
}
})
.collect();
let outputs = f(&duals);
if k == 0 {
values = outputs.iter().map(|d| d.value).collect();
jac = vec![vec![0.0; m]; outputs.len()];
}
for (i, out) in outputs.iter().enumerate() {
jac[i][k] = out.tangent;
}
}
(values, jac)
}
#[must_use]
pub fn directional_derivative<F: Fn(&[Dual]) -> Dual>(
f: F,
x: &[f64],
direction: &[f64],
) -> (f64, f64) {
assert_eq!(
direction.len(),
x.len(),
"direction length must match input length"
);
let duals: Vec<Dual> = x
.iter()
.zip(direction.iter())
.map(|(&value, &tangent)| Dual { value, tangent })
.collect();
f(&duals).extract()
}
#[cfg(test)]
mod tests {
use super::*;
use std::f64::consts::PI;
const TOL: f64 = 1e-10;
#[test]
fn dual_mul_known_answer() {
let d = Dual::seed(3.0);
let r = d * d;
assert!((r.value - 9.0).abs() < TOL, "primal {} != 9.0", r.value);
assert!(
(r.tangent - 6.0).abs() < TOL,
"tangent {} != 6.0",
r.tangent
);
}
#[test]
fn dual_sqrt_known_answer() {
let r = Scalar::sqrt(Dual::seed(4.0));
assert!((r.value - 2.0).abs() < TOL);
assert!((r.tangent - 0.25).abs() < TOL);
}
#[test]
fn dual_exp_known_answer() {
let e = std::f64::consts::E;
let r = Scalar::exp(Dual::seed(1.0));
assert!((r.value - e).abs() < TOL);
assert!((r.tangent - e).abs() < TOL);
}
#[test]
fn dual_ln_known_answer() {
let r = Scalar::ln(Dual::seed(2.0));
assert!((r.value - 2.0_f64.ln()).abs() < TOL);
assert!((r.tangent - 0.5).abs() < TOL);
}
#[test]
fn dual_sin_known_answer() {
let expected = 2.0_f64.sqrt() / 2.0;
let r = Scalar::sin(Dual::seed(PI / 4.0));
assert!((r.value - expected).abs() < TOL);
assert!((r.tangent - expected).abs() < TOL);
}
#[test]
fn dual_cos_known_answer() {
let expected = 2.0_f64.sqrt() / 2.0;
let r = Scalar::cos(Dual::seed(PI / 4.0));
assert!((r.value - expected).abs() < TOL);
assert!((r.tangent + expected).abs() < TOL);
}
#[test]
fn dual_powf_known_answer() {
let r = Scalar::powf(Dual::seed(4.0), 1.5);
assert!((r.value - 8.0).abs() < TOL);
assert!((r.tangent - 3.0).abs() < TOL);
}
#[test]
fn dual_abs_known_answer() {
let r = Scalar::abs(Dual::seed(2.0));
assert!((r.value - 2.0).abs() < TOL);
assert!((r.tangent - 1.0).abs() < TOL);
let rn = Scalar::abs(Dual::seed(-3.0));
assert!((rn.value - 3.0).abs() < TOL);
assert!((rn.tangent + 1.0).abs() < TOL);
}
#[test]
fn dual_sub_div_neg_known_answer() {
let d = Dual::seed(5.0);
let r = (d - Dual::constant(1.0)) / Dual::constant(2.0);
assert!((r.value - 2.0).abs() < TOL);
assert!((r.tangent - 0.5).abs() < TOL);
let n = -Dual::seed(7.0);
assert!((n.value + 7.0).abs() < TOL);
assert!((n.tangent + 1.0).abs() < TOL);
}
#[test]
fn dual_assign_ops() {
let mut acc = Dual::constant(0.0);
let x = Dual::seed(2.0);
acc += x; acc *= x; acc -= Dual::constant(1.0); assert!((acc.value - 3.0).abs() < TOL);
assert!((acc.tangent - 4.0).abs() < TOL);
}
#[test]
fn dual_composed_chain_known_answer() {
let x0 = 1.0_f64;
let (value, deriv) = diff(
|x| {
let e = Scalar::exp(x);
let s = Scalar::sin(x);
let x2 = x * x;
Scalar::sqrt(e * s + x2)
},
x0,
);
let f = (x0.exp() * x0.sin() + x0 * x0).sqrt();
let expected_deriv = (1.0 / (2.0 * f)) * (x0.exp() * (x0.sin() + x0.cos()) + 2.0 * x0);
assert!((value - f).abs() < TOL, "value {value} != {f}");
assert!(
(deriv - expected_deriv).abs() < TOL,
"deriv {deriv} != {expected_deriv}"
);
}
#[test]
fn dual_partial_cmp_value_only() {
use std::cmp::Ordering;
let a = Dual {
value: 1.0,
tangent: 0.5,
};
let b = Dual {
value: 1.0,
tangent: 0.3,
};
assert_eq!(a.partial_cmp(&b), Some(Ordering::Equal));
let small = Dual {
value: 0.5,
tangent: 99.0,
};
let big = Dual {
value: 2.0,
tangent: -99.0,
};
assert_eq!(small.partial_cmp(&big), Some(Ordering::Less));
assert!(small < big);
}
#[test]
fn dual_eq_is_value_only_and_consistent_with_ord() {
use std::cmp::Ordering;
let a = Dual {
value: 1.0,
tangent: 0.5,
};
let b = Dual {
value: 1.0,
tangent: -7.0,
};
assert_eq!(a, b);
assert_eq!(a.partial_cmp(&b), Some(Ordering::Equal));
let c = Dual {
value: 2.0,
tangent: 0.5,
};
assert_ne!(a, c);
}
#[test]
fn dual_abs_at_zero_tangent_is_zero() {
let r = Scalar::abs(Dual::seed(0.0));
assert_eq!(r.value, 0.0);
assert_eq!(r.tangent, 0.0);
let rn = Scalar::abs(Dual::seed(-0.0));
assert_eq!(rn.value, 0.0);
assert_eq!(rn.tangent, 0.0);
}
#[test]
fn dual_signum_tangent_is_zero_value_is_f64_signum() {
let rp = Scalar::signum(Dual::seed(3.0));
assert_eq!(rp.value, 1.0);
assert_eq!(rp.tangent, 0.0);
let rn = Scalar::signum(Dual::seed(-3.0));
assert_eq!(rn.value, -1.0);
assert_eq!(rn.tangent, 0.0);
let rz = Scalar::signum(Dual::seed(0.0));
assert_eq!(rz.value, 1.0);
assert_eq!(rz.tangent, 0.0);
}
#[test]
fn dual_sqrt_at_zero_tangent_is_nonfinite() {
let r = Scalar::sqrt(Dual::seed(0.0));
assert_eq!(r.value, 0.0);
assert!(
!r.tangent.is_finite(),
"tangent {} should be non-finite",
r.tangent
);
}
#[test]
fn dual_ln_at_zero_tangent_is_nonfinite() {
let r = Scalar::ln(Dual::seed(0.0));
assert!(r.value.is_infinite() && r.value < 0.0);
assert!(
!r.tangent.is_finite(),
"tangent {} should be non-finite",
r.tangent
);
}
#[test]
fn dual_powf_singular_and_linear_edges() {
let r = Scalar::powf(Dual::seed(0.0), 0.5);
assert_eq!(r.value, 0.0);
assert!(
r.tangent.is_infinite(),
"tangent {} should be infinite",
r.tangent
);
let lin = Scalar::powf(Dual::seed(0.0), 1.0);
assert_eq!(lin.value, 0.0);
assert_eq!(lin.tangent, 1.0);
let nan = Scalar::powf(Dual::seed(-2.0), 0.5);
assert!(nan.value.is_nan());
assert!(nan.tangent.is_nan());
}
fn central_fd(f: impl Fn(f64) -> f64, x: f64) -> f64 {
let h = 1e-8_f64;
(f(x + h) - f(x - h)) / (2.0 * h)
}
#[test]
fn finite_diff_cross_check_composed() {
let x0 = 1.0_f64;
let (_, ad) = diff(
|x| {
let e = Scalar::exp(x);
let s = Scalar::sin(x);
Scalar::sqrt(e * s + x * x)
},
x0,
);
let fd = central_fd(|x| (x.exp() * x.sin() + x * x).sqrt(), x0);
assert!((ad - fd).abs() < 1e-6, "AD {ad} FD {fd}");
}
#[test]
fn finite_diff_cross_check_log_trig() {
let x0 = 2.0_f64;
let (_, ad) = diff(|x| Scalar::ln(x) * Scalar::cos(x), x0);
let fd = central_fd(|x| x.ln() * x.cos(), x0);
assert!((ad - fd).abs() < 1e-6, "AD {ad} FD {fd}");
}
#[test]
fn finite_diff_cross_check_powf_exp() {
let x0 = 1.5_f64;
let (_, ad) = diff(|x| Scalar::powf(x, 1.5) / Scalar::exp(x), x0);
let fd = central_fd(|x| x.powf(1.5) / x.exp(), x0);
assert!((ad - fd).abs() < 1e-6, "AD {ad} FD {fd}");
}
#[test]
fn f64_parity_transcendentals() {
let x = 2.5_f64;
assert_eq!(<f64 as Scalar>::sqrt(x), x.sqrt());
assert_eq!(<f64 as Scalar>::exp(x), x.exp());
assert_eq!(<f64 as Scalar>::ln(x), x.ln());
assert_eq!(<f64 as Scalar>::sin(x), x.sin());
assert_eq!(<f64 as Scalar>::cos(x), x.cos());
assert_eq!(<f64 as Scalar>::powf(x, 1.5), x.powf(1.5));
assert_eq!(<f64 as Scalar>::abs(-x), (-x).abs());
assert_eq!(<f64 as Scalar>::signum(-x), (-x).signum());
}
#[test]
fn f64_parity_constants() {
assert_eq!(<f64 as Scalar>::zero(), 0.0);
assert_eq!(<f64 as Scalar>::one(), 1.0);
assert_eq!(<f64 as Scalar>::from_f64(3.25), 3.25);
assert_eq!(<f64 as Scalar>::infinity(), f64::INFINITY);
}
#[test]
fn grad_sum_of_squares_closed_form() {
let (value, gradient) = grad(
|x| {
let mut acc = Dual::constant(0.0);
for &xi in x {
acc += xi * xi;
}
acc
},
&[1.0, 2.0, 3.0],
);
assert_eq!(gradient.len(), 3);
assert!((value - 14.0).abs() <= 1e-12, "value {value} != 14.0");
for (g, expected) in gradient.iter().zip([2.0, 4.0, 6.0]) {
assert!((g - expected).abs() <= 1e-12, "grad {g} != {expected}");
}
}
#[test]
fn grad_single_input_agrees_with_diff() {
let (value, gradient) = grad(|x| x[0] * x[0], &[3.0]);
assert_eq!(gradient.len(), 1);
assert!((value - 9.0).abs() <= 1e-12);
assert!((gradient[0] - 6.0).abs() <= 1e-12);
let (dv, dd) = diff(|x| x * x, 3.0);
assert!((value - dv).abs() <= 1e-12);
assert!((gradient[0] - dd).abs() <= 1e-12);
}
#[test]
fn grad_empty_input_returns_constant() {
let (value, gradient) = grad(|_x| Dual::constant(7.0), &[]);
assert!((value - 7.0).abs() <= 1e-12);
assert!(gradient.is_empty());
}
#[test]
fn jacobian_known_answer() {
let (values, j) = jacobian(|x| vec![x[0] * x[1], x[0] + x[1]], &[2.0, 3.0]);
assert_eq!(values.len(), 2);
assert!((values[0] - 6.0).abs() <= 1e-12);
assert!((values[1] - 5.0).abs() <= 1e-12);
assert_eq!(j.len(), 2);
assert!((j[0][0] - 3.0).abs() <= 1e-12 && (j[0][1] - 2.0).abs() <= 1e-12);
assert!((j[1][0] - 1.0).abs() <= 1e-12 && (j[1][1] - 1.0).abs() <= 1e-12);
}
#[test]
fn directional_derivative_projects_gradient() {
let (value, dd) =
directional_derivative(|x| x[0] * x[0] + x[1] * x[1], &[1.0, 2.0], &[1.0, 0.0]);
assert!((value - 5.0).abs() <= 1e-12);
assert!((dd - 2.0).abs() <= 1e-12);
}
#[test]
fn grad_composed_objective_matches_finite_diff() {
use crate::matrix::FdMatrix;
use crate::metric::soft_dtw_distance_generic;
use crate::regression::{fdata_to_pc_1d, project_scores_generic};
use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
let m = 24usize;
let n = 40usize;
let ncomp = 3usize;
let gamma = 0.1_f64;
let lambda = 1.0_f64;
let argvals: Vec<f64> = (0..m)
.map(|j| 0.1 + 0.8 * j as f64 / (m - 1) as f64)
.collect();
let mut rng = StdRng::seed_from_u64(20260906);
let mut data = vec![0.0f64; n * m];
for i in 0..n {
let a: f64 = rng.gen_range(-1.0..1.0);
let b: f64 = rng.gen_range(-1.0..1.0);
let c: f64 = rng.gen_range(-1.0..1.0);
for (j, &t) in argvals.iter().enumerate() {
let v = a * (PI * t).sin() + b * (2.0 * PI * t).cos() + c * (3.0 * PI * t).sin();
data[i + j * n] = v;
}
}
let data = FdMatrix::from_column_major(data, n, m).unwrap();
let fpca = fdata_to_pc_1d(&data, ncomp, &argvals).unwrap();
let mean = fpca.mean.clone();
let rotation = fpca.rotation.clone();
let weights = fpca.weights.clone();
let curve: Vec<f64> = argvals
.iter()
.map(|&t| {
0.7 * (PI * t).sin() - 0.4 * (2.0 * PI * t).cos() + 0.3 * (3.0 * PI * t).sin()
})
.collect();
let reference: Vec<f64> = argvals
.iter()
.map(|&t| {
0.2 * (PI * t).sin() + 0.5 * (2.0 * PI * t).cos() - 0.6 * (3.0 * PI * t).sin()
})
.collect();
let reference_duals: Vec<Dual> = reference.iter().map(|&r| Dual::constant(r)).collect();
let objective = |c: &[Dual]| -> Dual {
let sdtw = soft_dtw_distance_generic(c, &reference_duals, gamma);
let scores = project_scores_generic(c, &mean, &rotation, &weights, ncomp);
let mut acc = Dual::constant(0.0);
for s in &scores {
acc += *s * *s;
}
sdtw + Dual::constant(lambda) * acc
};
let (value, gradient) = grad(objective, &curve);
assert_eq!(gradient.len(), m);
assert!(value.is_finite(), "objective value not finite: {value}");
assert!(value > 0.0, "objective value not positive: {value}");
let f64_obj = |c: &[f64]| -> f64 {
let d = soft_dtw_distance_generic::<f64>(c, &reference, gamma);
let sc = project_scores_generic::<f64>(c, &mean, &rotation, &weights, ncomp);
d + lambda * sc.iter().map(|s| s * s).sum::<f64>()
};
assert!(
(value - f64_obj(&curve)).abs() < 1e-12,
"composition parity broke: {value}"
);
let h = 1e-6_f64;
for j in 0..m {
let mut plus = curve.clone();
let mut minus = curve.clone();
plus[j] += h;
minus[j] -= h;
let fd = (f64_obj(&plus) - f64_obj(&minus)) / (2.0 * h);
assert!(
(gradient[j] - fd).abs() < 1e-6,
"component {j}: AD {} vs FD {fd}",
gradient[j]
);
}
}
#[test]
fn dual_constants() {
let z = <Dual as Scalar>::zero();
assert_eq!(z.value, 0.0);
assert_eq!(z.tangent, 0.0);
let o = <Dual as Scalar>::one();
assert_eq!(o.value, 1.0);
assert_eq!(o.tangent, 0.0);
let c = <Dual as Scalar>::from_f64(4.5);
assert_eq!(c.value, 4.5);
assert_eq!(c.tangent, 0.0);
let inf = <Dual as Scalar>::infinity();
assert!(inf.value.is_infinite() && inf.value > 0.0);
assert_eq!(inf.tangent, 0.0);
}
}