use crate::autodiff::Scalar;
use crate::iter_maybe_parallel;
use crate::matrix::FdMatrix;
#[cfg(feature = "parallel")]
use rayon::iter::ParallelIterator;
use super::{cross_distance_matrix, self_distance_matrix};
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct SoftDtwBarycenterResult {
pub barycenter: Vec<f64>,
pub n_iter: usize,
pub converged: bool,
}
#[inline]
pub(super) fn softmin3(a: f64, b: f64, c: f64, gamma: f64) -> f64 {
let min_val = a.min(b).min(c);
if !min_val.is_finite() {
return min_val;
}
let neg_inv_gamma = -1.0 / gamma;
let ea = ((a - min_val) * neg_inv_gamma).exp();
let eb = ((b - min_val) * neg_inv_gamma).exp();
let ec = ((c - min_val) * neg_inv_gamma).exp();
min_val - gamma * (ea + eb + ec).ln()
}
#[inline]
fn softmin3_generic<S: Scalar>(a: S, b: S, c: S, gamma: f64) -> S {
let min_val = if a <= b {
if a <= c {
a
} else {
c
}
} else if b <= c {
b
} else {
c
};
if min_val >= S::infinity() {
return min_val;
}
let neg_inv_gamma = S::from_f64(-1.0 / gamma);
let term = |v: S| -> S {
if v >= S::infinity() {
S::zero()
} else {
S::exp((v - min_val) * neg_inv_gamma)
}
};
let ea = term(a);
let eb = term(b);
let ec = term(c);
min_val - S::from_f64(gamma) * S::ln(ea + eb + ec)
}
fn soft_dtw_distance_inner<S: Scalar>(x: &[S], y: &[S], gamma: f64) -> S {
let n = x.len();
let m = y.len();
if n == 0 || m == 0 {
return S::zero();
}
let mut prev = vec![S::infinity(); m + 1];
let mut curr = vec![S::infinity(); m + 1];
prev[0] = S::zero();
for i in 1..=n {
for v in curr.iter_mut() {
*v = S::infinity();
}
for j in 1..=m {
let d = x[i - 1] - y[j - 1];
let cost = d * d;
curr[j] = cost + softmin3_generic(prev[j], curr[j - 1], prev[j - 1], gamma);
}
std::mem::swap(&mut prev, &mut curr);
}
prev[m]
}
pub fn soft_dtw_distance(x: &[f64], y: &[f64], gamma: f64) -> f64 {
soft_dtw_distance_inner(x, y, gamma)
}
pub fn soft_dtw_distance_generic<S: Scalar>(x: &[S], y: &[S], gamma: f64) -> S {
soft_dtw_distance_inner(x, y, gamma)
}
pub fn soft_dtw_divergence(x: &[f64], y: &[f64], gamma: f64) -> f64 {
let xy = soft_dtw_distance(x, y, gamma);
let xx = soft_dtw_distance(x, x, gamma);
let yy = soft_dtw_distance(y, y, gamma);
xy - 0.5 * (xx + yy)
}
pub fn soft_dtw_self_1d(data: &FdMatrix, gamma: f64) -> FdMatrix {
let n = data.nrows();
if n == 0 || data.ncols() == 0 {
return FdMatrix::zeros(0, 0);
}
let rows: Vec<Vec<f64>> = (0..n).map(|i| data.row(i)).collect();
let mut dist = self_distance_matrix(n, |i, j| soft_dtw_distance(&rows[i], &rows[j], gamma));
for i in 0..n {
dist[(i, i)] = soft_dtw_distance(&rows[i], &rows[i], gamma);
}
dist
}
pub fn soft_dtw_cross_1d(data1: &FdMatrix, data2: &FdMatrix, gamma: f64) -> FdMatrix {
let n1 = data1.nrows();
let n2 = data2.nrows();
if n1 == 0 || n2 == 0 || data1.ncols() == 0 || data2.ncols() == 0 {
return FdMatrix::zeros(0, 0);
}
let rows1: Vec<Vec<f64>> = (0..n1).map(|i| data1.row(i)).collect();
let rows2: Vec<Vec<f64>> = (0..n2).map(|i| data2.row(i)).collect();
cross_distance_matrix(n1, n2, |i, j| {
soft_dtw_distance(&rows1[i], &rows2[j], gamma)
})
}
pub fn soft_dtw_div_self_1d(data: &FdMatrix, gamma: f64) -> FdMatrix {
let n = data.nrows();
if n == 0 || data.ncols() == 0 {
return FdMatrix::zeros(0, 0);
}
let rows: Vec<Vec<f64>> = (0..n).map(|i| data.row(i)).collect();
let self_dists: Vec<f64> = iter_maybe_parallel!(0..n)
.map(|i| soft_dtw_distance(&rows[i], &rows[i], gamma))
.collect();
self_distance_matrix(n, |i, j| {
let xy = soft_dtw_distance(&rows[i], &rows[j], gamma);
xy - 0.5 * (self_dists[i] + self_dists[j])
})
}
pub fn soft_dtw_div_cross_1d(data1: &FdMatrix, data2: &FdMatrix, gamma: f64) -> FdMatrix {
let n1 = data1.nrows();
let n2 = data2.nrows();
if n1 == 0 || n2 == 0 || data1.ncols() == 0 || data2.ncols() == 0 {
return FdMatrix::zeros(0, 0);
}
let rows1: Vec<Vec<f64>> = (0..n1).map(|i| data1.row(i)).collect();
let rows2: Vec<Vec<f64>> = (0..n2).map(|i| data2.row(i)).collect();
let self1: Vec<f64> = iter_maybe_parallel!(0..n1)
.map(|i| soft_dtw_distance(&rows1[i], &rows1[i], gamma))
.collect();
let self2: Vec<f64> = iter_maybe_parallel!(0..n2)
.map(|j| soft_dtw_distance(&rows2[j], &rows2[j], gamma))
.collect();
cross_distance_matrix(n1, n2, |i, j| {
let xy = soft_dtw_distance(&rows1[i], &rows2[j], gamma);
xy - 0.5 * (self1[i] + self2[j])
})
}
fn soft_dtw_forward(x: &[f64], y: &[f64], gamma: f64) -> Vec<Vec<f64>> {
let n = x.len();
let m = y.len();
let mut r = vec![vec![f64::INFINITY; m + 1]; n + 1];
r[0][0] = 0.0;
for i in 1..=n {
for j in 1..=m {
let d = x[i - 1] - y[j - 1];
let cost = d * d;
r[i][j] = cost + softmin3(r[i - 1][j], r[i][j - 1], r[i - 1][j - 1], gamma);
}
}
r
}
fn soft_dtw_backward(x: &[f64], y: &[f64], r: &[Vec<f64>], gamma: f64) -> Vec<Vec<f64>> {
let n = x.len();
let m = y.len();
let mut e = vec![vec![0.0; m + 2]; n + 2];
e[n][m] = 1.0;
for i in (1..=n).rev() {
for j in (1..=m).rev() {
if i == n && j == m {
continue;
}
let a = if i < n {
e[i + 1][j]
* (-(r[i][j] - r[i + 1][j] + r[i + 1][j] - softmin3_val(r, i + 1, j, gamma))
/ gamma)
.exp()
} else {
0.0
};
let b = if j < m {
e[i][j + 1]
* (-(r[i][j] - r[i][j + 1] + r[i][j + 1] - softmin3_val(r, i, j + 1, gamma))
/ gamma)
.exp()
} else {
0.0
};
let c = if i < n && j < m {
e[i + 1][j + 1]
* (-(r[i][j] - r[i + 1][j + 1] + r[i + 1][j + 1]
- softmin3_val(r, i + 1, j + 1, gamma))
/ gamma)
.exp()
} else {
0.0
};
e[i][j] = a + b + c;
}
}
e
}
#[inline]
fn softmin3_val(r: &[Vec<f64>], i: usize, j: usize, gamma: f64) -> f64 {
softmin3(
if i > 0 { r[i - 1][j] } else { f64::INFINITY },
if j > 0 { r[i][j - 1] } else { f64::INFINITY },
if i > 0 && j > 0 {
r[i - 1][j - 1]
} else {
f64::INFINITY
},
gamma,
)
}
fn soft_dtw_accumulate_gradient_and_weight(
bary: &[f64],
xi: &[f64],
gamma: f64,
grad: &mut [f64],
weight: &mut [f64],
) {
let m = bary.len();
let r = soft_dtw_forward(bary, xi, gamma);
let e = soft_dtw_backward(bary, xi, &r, gamma);
for k in 1..=m {
let mut g = 0.0;
let mut w = 0.0;
for j in 1..=xi.len() {
g += e[k][j] * 2.0 * (bary[k - 1] - xi[j - 1]);
w += e[k][j];
}
grad[k - 1] += g;
weight[k - 1] += w;
}
}
fn update_barycenter(bary: &mut [f64], grad: &[f64], weight: &[f64], tol: f64) -> bool {
let mut max_change = 0.0_f64;
let mut max_val = 0.0_f64;
for ((b, &g), &w) in bary.iter_mut().zip(grad.iter()).zip(weight.iter()) {
let update = if w > 1e-12 { g / (2.0 * w) } else { 0.0 };
*b -= update;
max_change = max_change.max(update.abs());
max_val = max_val.max(b.abs());
}
max_val > 0.0 && max_change / max_val < tol
}
fn init_barycenter_mean(rows: &[Vec<f64>]) -> Vec<f64> {
let n = rows.len();
let m = rows[0].len();
let mut bary = vec![0.0; m];
for row in rows {
for (j, val) in row.iter().enumerate() {
bary[j] += val;
}
}
for v in &mut bary {
*v /= n as f64;
}
bary
}
pub fn soft_dtw_barycenter(
data: &FdMatrix,
gamma: f64,
max_iter: usize,
tol: f64,
) -> SoftDtwBarycenterResult {
let (n, m) = data.shape();
if n == 0 || m == 0 {
return SoftDtwBarycenterResult {
barycenter: Vec::new(),
n_iter: 0,
converged: true,
};
}
let rows: Vec<Vec<f64>> = (0..n).map(|i| data.row(i)).collect();
let mut bary = init_barycenter_mean(&rows);
let mut converged = false;
let mut n_iter = 0;
for iter in 0..max_iter {
n_iter = iter + 1;
let mut grad = vec![0.0; m];
let mut weight = vec![0.0; m];
for row in &rows {
soft_dtw_accumulate_gradient_and_weight(&bary, row, gamma, &mut grad, &mut weight);
}
if update_barycenter(&mut bary, &grad, &weight, tol) {
converged = true;
break;
}
}
SoftDtwBarycenterResult {
barycenter: bary,
n_iter,
converged,
}
}
#[cfg(test)]
mod differentiable_tests {
use super::*;
use crate::autodiff::Dual;
const X: [f64; 5] = [0.1, 0.4, 0.9, 1.2, 0.7];
const Y: [f64; 5] = [0.2, 0.3, 1.0, 1.1, 0.6];
fn dual_grad(x: &[f64], y: &[f64], gamma: f64) -> Vec<f64> {
let y_dual: Vec<Dual> = y.iter().map(|&v| Dual::constant(v)).collect();
(0..x.len())
.map(|k| {
let x_dual: Vec<Dual> = x
.iter()
.enumerate()
.map(|(i, &v)| {
if i == k {
Dual::seed(v)
} else {
Dual::constant(v)
}
})
.collect();
soft_dtw_distance_generic(&x_dual, &y_dual, gamma)
.extract()
.1
})
.collect()
}
fn corrected_oracle_gradient(bary: &[f64], xi: &[f64], gamma: f64) -> Vec<f64> {
let n = bary.len();
let m = xi.len();
let mut r = vec![vec![f64::INFINITY; m + 1]; n + 1];
r[0][0] = 0.0;
for i in 1..=n {
for j in 1..=m {
let d = bary[i - 1] - xi[j - 1];
let cost = d * d;
r[i][j] = cost + softmin3(r[i - 1][j], r[i][j - 1], r[i - 1][j - 1], gamma);
}
}
let mut e = vec![vec![0.0; m + 2]; n + 2];
e[n][m] = 1.0;
for i in (1..=n).rev() {
for j in (1..=m).rev() {
if i == n && j == m {
continue; }
let a = if i < n {
e[i + 1][j] * (-(r[i][j] - softmin3_val(&r, i + 1, j, gamma)) / gamma).exp()
} else {
0.0
};
let b = if j < m {
e[i][j + 1] * (-(r[i][j] - softmin3_val(&r, i, j + 1, gamma)) / gamma).exp()
} else {
0.0
};
let c = if i < n && j < m {
e[i + 1][j + 1]
* (-(r[i][j] - softmin3_val(&r, i + 1, j + 1, gamma)) / gamma).exp()
} else {
0.0
};
e[i][j] = a + b + c;
}
}
let mut grad = vec![0.0; n];
for k in 1..=n {
let mut g = 0.0;
for j in 1..=m {
g += e[k][j] * 2.0 * (bary[k - 1] - xi[j - 1]);
}
grad[k - 1] = g;
}
grad
}
#[test]
fn dual_gradient_vs_oracle() {
let gamma = 1.0;
let dual = dual_grad(&X, &Y, gamma);
let oracle = corrected_oracle_gradient(&X, &Y, gamma);
for k in 0..X.len() {
assert!(
(dual[k] - oracle[k]).abs() <= 1e-9,
"k={k}: dual {} vs oracle {}",
dual[k],
oracle[k]
);
}
}
#[test]
fn f64_parity() {
for &gamma in &[0.1, 1.0, 10.0] {
let generic = soft_dtw_distance_generic::<f64>(&X, &Y, gamma);
let original = soft_dtw_distance(&X, &Y, gamma);
assert!(
(generic - original).abs() <= 1e-12,
"gamma={gamma}: generic {generic} vs original {original}"
);
assert_eq!(generic, original, "gamma={gamma}: not bit-identical");
}
let xl: Vec<f64> = (0..20).map(|i| (i as f64 * 0.31).sin()).collect();
let yl: Vec<f64> = (0..20).map(|i| (i as f64 * 0.27).cos()).collect();
assert_eq!(
soft_dtw_distance_generic::<f64>(&xl, &yl, 1.0),
soft_dtw_distance(&xl, &yl, 1.0)
);
}
#[test]
fn soft_dtw_backward_nonzero_and_matches_oracle() {
let gamma = 1.0;
let r = soft_dtw_forward(&X, &Y, gamma);
let e = soft_dtw_backward(&X, &Y, &r, gamma);
let e_nonzero = e.iter().flatten().any(|&v| v.abs() > 1e-12);
assert!(
e_nonzero,
"E matrix must be non-zero after CORR-01 fix (all-zero indicates endpoint seed was zeroed)"
);
let n = X.len();
let mut shipped_grad = vec![0.0; n];
let mut scratch_weight = vec![0.0; n];
soft_dtw_accumulate_gradient_and_weight(
&X,
&Y,
gamma,
&mut shipped_grad,
&mut scratch_weight,
);
let oracle = corrected_oracle_gradient(&X, &Y, gamma);
let dual = dual_grad(&X, &Y, gamma);
for k in 0..n {
let rel_oracle = (shipped_grad[k] - oracle[k]).abs() / oracle[k].abs().max(1e-12);
assert!(
rel_oracle <= 1e-6,
"k={k}: shipped grad {:.8} vs oracle {:.8} (rel err {:.2e} > 1e-6)",
shipped_grad[k],
oracle[k],
rel_oracle
);
let rel_dual = (shipped_grad[k] - dual[k]).abs() / dual[k].abs().max(1e-12);
assert!(
rel_dual <= 1e-6,
"k={k}: shipped grad {:.8} vs Dual {:.8} (rel err {:.2e} > 1e-6)",
shipped_grad[k],
dual[k],
rel_dual
);
}
}
#[test]
fn dual_gradient_vs_fd() {
let h = 1e-8;
for &gamma in &[0.5, 1.0, 2.0] {
let dual = dual_grad(&X, &Y, gamma);
for k in 0..X.len() {
let mut xp = X.to_vec();
let mut xm = X.to_vec();
xp[k] += h;
xm[k] -= h;
let fd = (soft_dtw_distance_generic::<f64>(&xp, &Y, gamma)
- soft_dtw_distance_generic::<f64>(&xm, &Y, gamma))
/ (2.0 * h);
assert!(
(dual[k] - fd).abs() <= 1e-6,
"gamma={gamma} k={k}: dual {} vs fd {}",
dual[k],
fd
);
}
}
}
}