use scirs2_core::linalg::solve_ndarray;
use scirs2_core::ndarray::{Array1, Array2};
pub(crate) fn smooth_whittaker(y: &[f64], lambda: f64, order: usize) -> Vec<f64> {
let n = y.len();
if n == 0 || n <= order {
return y.to_vec();
}
let w: Vec<f64> = y
.iter()
.map(|&v| if v.is_nan() { 0.0 } else { 1.0 })
.collect();
if w.iter().all(|&wi| wi == 0.0) {
return y.to_vec();
}
let y_clean: Array1<f64> =
Array1::from_iter(y.iter().map(|&v| if v.is_nan() { 0.0 } else { v }));
let d = build_difference_matrix(n, order);
let mut w_mat = Array2::<f64>::zeros((n, n));
for i in 0..n {
w_mat[[i, i]] = w[i];
}
let dtd = d.t().dot(&d);
let a: Array2<f64> = w_mat + lambda * dtd;
let b: Array1<f64> = Array1::from_iter((0..n).map(|i| w[i] * y_clean[i]));
match solve_ndarray(&a, &b) {
Ok(z) => z.to_vec(),
Err(_) => y.to_vec(), }
}
fn build_difference_matrix(n: usize, order: usize) -> Array2<f64> {
if order == 0 || n <= order {
return Array2::eye(n);
}
let mut d = build_first_difference(n);
for _ in 1..order {
let rows = d.nrows();
let d1 = build_first_difference(rows);
d = d1.dot(&d);
}
d
}
fn build_first_difference(n: usize) -> Array2<f64> {
let mut d = Array2::<f64>::zeros((n - 1, n));
for i in 0..n - 1 {
d[[i, i]] = -1.0;
d[[i, i + 1]] = 1.0;
}
d
}
#[cfg(test)]
mod unit_tests {
use super::*;
#[test]
fn constant_series_reproduced_exactly() {
let y = vec![5.0_f64; 20];
let z = smooth_whittaker(&y, 100.0, 2);
for (&zi, &yi) in z.iter().zip(y.iter()) {
assert!(
(zi - yi).abs() < 1e-8,
"constant series not reproduced: got {zi}, expected {yi}"
);
}
}
#[test]
fn empty_input_returns_empty() {
let z = smooth_whittaker(&[], 100.0, 2);
assert!(z.is_empty());
}
#[test]
fn short_series_at_boundary_falls_back() {
let y = vec![1.0, 2.0]; let z = smooth_whittaker(&y, 100.0, 2);
assert_eq!(z, y);
}
#[test]
fn all_nan_returns_input() {
let y = vec![f64::NAN; 10];
let z = smooth_whittaker(&y, 100.0, 2);
assert_eq!(z.len(), y.len());
assert!(z.iter().all(|v| v.is_nan()));
}
#[test]
fn build_first_difference_shape() {
let d = build_first_difference(5);
assert_eq!(d.nrows(), 4);
assert_eq!(d.ncols(), 5);
}
#[test]
fn build_difference_matrix_order2_shape() {
let d = build_difference_matrix(10, 2);
assert_eq!(d.nrows(), 8); assert_eq!(d.ncols(), 10);
}
}