pub const COST_EPSILON: f64 = 1e-9;
#[inline]
pub fn interval_dist(v: f64, lo: f64, hi: f64) -> f64 {
(lo - v).max(0.0).max(v - hi)
}
#[inline]
pub fn c_func_merge_lb(a: f64, b: f64, lo: f64, hi: f64, c_const: f64) -> f64 {
let penalty = if (a >= b && hi >= a) || (a <= b && lo <= a) {
0.0
} else {
(a - b).abs().min(interval_dist(a, lo, hi))
};
c_const + penalty
}
#[inline]
pub fn c_func_split_lb(a_lo: f64, a_hi: f64, b: f64, c_lo: f64, c_hi: f64, c_const: f64) -> f64 {
let union_lo = b.min(c_lo);
let union_hi = b.max(c_hi);
let penalty = if a_lo <= union_hi && a_hi >= union_lo {
0.0
} else if a_lo > union_hi {
a_lo - union_hi
} else {
union_lo - a_hi
};
c_const + penalty
}
pub fn step_interval_column(
prev_col: &[f64],
query: &[f64],
curr: (f64, f64),
prev: Option<(f64, f64)>,
c_const: f64,
) -> Vec<f64> {
let m = query.len();
assert!(m > 0, "step_interval_column requires a non-empty query");
let (clo, chi) = curr;
let mut col = vec![f64::INFINITY; m + 1];
match prev {
None => {
col[1] = interval_dist(query[0], clo, chi);
for i in 2..=m {
if col[i - 1].is_finite() {
col[i] =
col[i - 1] + c_func_merge_lb(query[i - 1], query[i - 2], clo, chi, c_const);
}
}
}
Some((plo, phi)) => {
if prev_col[1].is_finite() {
col[1] = prev_col[1] + c_func_split_lb(clo, chi, query[0], plo, phi, c_const);
}
for i in 2..=m {
let mv = if prev_col[i - 1].is_finite() {
prev_col[i - 1] + interval_dist(query[i - 1], clo, chi)
} else {
f64::INFINITY
};
let mg = if col[i - 1].is_finite() {
col[i - 1] + c_func_merge_lb(query[i - 1], query[i - 2], clo, chi, c_const)
} else {
f64::INFINITY
};
let sp = if prev_col[i].is_finite() {
prev_col[i] + c_func_split_lb(clo, chi, query[i - 1], plo, phi, c_const)
} else {
f64::INFINITY
};
col[i] = mv.min(mg).min(sp);
}
}
}
col
}
#[inline]
pub fn column_lower_bound(col: &[f64]) -> f64 {
col[1..]
.iter()
.copied()
.fold(f64::INFINITY, |acc, v| acc.min(v))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::time_series::MsmConfig;
#[test]
fn interval_dist_basic() {
assert_eq!(interval_dist(5.0, 0.0, 10.0), 0.0); assert_eq!(interval_dist(-2.0, 0.0, 10.0), 2.0); assert_eq!(interval_dist(13.0, 0.0, 10.0), 3.0); assert_eq!(interval_dist(5.0, f64::NEG_INFINITY, 10.0), 0.0);
assert_eq!(interval_dist(20.0, 0.0, f64::INFINITY), 0.0);
}
#[test]
fn merge_lb_is_min_over_c() {
let config = MsmConfig::new(1.3);
for &(a, b, lo, hi) in &[
(2.0_f64, 1.0_f64, 0.0_f64, 5.0_f64),
(2.0, 5.0, 0.0, 5.0),
(-3.0, 1.0, 0.0, 5.0),
(9.0, 1.0, 0.0, 5.0),
(3.0, 3.0, 1.0, 2.0),
(0.5, 0.0, 1.0, 4.0),
] {
let lb = c_func_merge_lb(a, b, lo, hi, config.c);
let mut brute = f64::INFINITY;
let steps = 4000;
for s in 0..=steps {
let c = lo + (hi - lo) * (s as f64 / steps as f64);
brute = brute.min(config.c_func(a, b, c));
}
assert!(
lb <= brute + 1e-6 && lb >= brute - 1e-2,
"merge a={a} b={b} [{lo},{hi}]: lb={lb} brute={brute}"
);
}
}
#[test]
fn split_lb_is_min_over_box() {
let config = MsmConfig::new(0.7);
for &(alo, ahi, b, clo, chi) in &[
(0.0_f64, 5.0_f64, 2.0_f64, 0.0_f64, 5.0_f64),
(6.0, 9.0, 2.0, 0.0, 4.0), (-5.0, -2.0, 2.0, 0.0, 4.0), (1.0, 1.5, 3.0, 4.0, 6.0),
(10.0, 12.0, 1.0, 2.0, 3.0),
] {
let lb = c_func_split_lb(alo, ahi, b, clo, chi, config.c);
let mut brute = f64::INFINITY;
let steps = 200;
for sa in 0..=steps {
let a = alo + (ahi - alo) * (sa as f64 / steps as f64);
for sc in 0..=steps {
let c = clo + (chi - clo) * (sc as f64 / steps as f64);
brute = brute.min(config.c_func(a, b, c));
}
}
assert!(
lb <= brute + 1e-6 && lb >= brute - 1e-2,
"split [{alo},{ahi}] b={b} [{clo},{chi}]: lb={lb} brute={brute}"
);
}
}
#[test]
fn degenerate_bins_reproduce_scalar_dp() {
let config = MsmConfig::new(1.0);
let query = vec![1.0, 3.0, 2.0, 5.0];
let target = vec![1.0, 2.5, 4.0];
let m = query.len();
let mut col = vec![f64::INFINITY; m + 1];
let mut prev_interval: Option<(f64, f64)> = None;
for (j, &y) in target.iter().enumerate() {
col = step_interval_column(&col, &query, (y, y), prev_interval, config.c);
prev_interval = Some((y, y));
let _ = j;
}
let exact = config.distance(&query, &target);
assert!(
(col[m] - exact).abs() < 1e-9,
"interval col[m]={} != exact DP {}",
col[m],
exact
);
}
#[test]
fn interval_column_lower_bounds_concrete() {
let config = MsmConfig::new(1.0);
let query = vec![0.4, 2.1, 3.9, 1.2];
let bins = [(0.0, 1.0), (2.0, 3.0), (3.0, 4.0)];
let concretes = [
vec![0.0, 2.0, 3.0],
vec![1.0, 3.0, 4.0],
vec![0.5, 2.5, 3.5],
vec![0.9, 2.1, 3.9],
];
let m = query.len();
let mut col = vec![f64::INFINITY; m + 1];
let mut prev_interval: Option<(f64, f64)> = None;
for &(lo, hi) in &bins {
col = step_interval_column(&col, &query, (lo, hi), prev_interval, config.c);
prev_interval = Some((lo, hi));
}
let lb_final = col[m];
let lb_subtree = column_lower_bound(&col);
for c in &concretes {
let exact = config.distance(&query, c);
assert!(
lb_final <= exact + 1e-9,
"col[m] LB {lb_final} exceeded exact {exact} for {c:?}"
);
assert!(
lb_subtree <= exact + 1e-9,
"subtree LB {lb_subtree} exceeded exact {exact} for {c:?}"
);
}
}
}