use scirs2_core::ndarray::{Array, Array1, Axis, Dimension};
use scirs2_core::numeric::{Float, FromPrimitive, One, Zero};
use std::fmt::Debug;
use crate::error::{NdimageError, NdimageResult};
#[allow(dead_code)]
fn get_spline_poles<T: Float + FromPrimitive>(order: usize) -> Vec<T> {
match order {
0 | 1 => vec![], 2 => {
let sqrt8 = T::from_f64(8.0).expect("Operation failed").sqrt();
let three = T::from_f64(3.0).expect("Operation failed");
vec![sqrt8 - three]
}
3 => {
let sqrt3 = T::from_f64(3.0).expect("Operation failed").sqrt();
let two = T::from_f64(2.0).expect("Operation failed");
vec![sqrt3 - two]
}
4 => {
let val1 = T::from_f64(0.361341225285).expect("Operation failed"); let val2 = T::from_f64(0.013725429297).expect("Operation failed"); vec![val1, val2]
}
5 => {
let val1 = T::from_f64(0.430575347099).expect("Operation failed");
let val2 = T::from_f64(0.043096288203).expect("Operation failed");
vec![val1, val2]
}
_ => vec![], }
}
#[allow(dead_code)]
fn get_initial_causal_coefficient<T: Float + FromPrimitive>(
coeffs: &[T],
pole: T,
tolerance: T,
) -> T {
let mut sum = T::zero();
let mut z_power = T::one();
let _abs_pole = pole.abs();
for &coeff in coeffs {
sum = sum + coeff * z_power;
z_power = z_power * pole;
if z_power.abs() < tolerance {
break;
}
}
sum
}
#[allow(dead_code)]
fn get_initial_anti_causal_coefficient<T: Float + FromPrimitive>(coeffs: &[T], pole: T) -> T {
let n = coeffs.len();
if n < 2 {
return T::zero();
}
let last_idx = n - 1;
(pole / (pole * pole - T::one())) * (pole * coeffs[last_idx] + coeffs[last_idx - 1])
}
#[allow(dead_code)]
fn apply_causal_filter<T: Float + FromPrimitive>(coeffs: &mut [T], pole: T, initialcoeff: T) {
if coeffs.is_empty() {
return;
}
coeffs[0] = initialcoeff;
for i in 1..coeffs.len() {
coeffs[i] = coeffs[i] + pole * coeffs[i - 1];
}
}
#[allow(dead_code)]
fn apply_anti_causal_filter<T: Float + FromPrimitive>(coeffs: &mut [T], pole: T, initialcoeff: T) {
if coeffs.is_empty() {
return;
}
let last_idx = coeffs.len() - 1;
coeffs[last_idx] = initialcoeff;
for i in (0..last_idx).rev() {
coeffs[i] = pole * (coeffs[i + 1] - coeffs[i]);
}
}
#[allow(dead_code)]
pub fn spline_filter<T, D>(input: &Array<T, D>, order: Option<usize>) -> NdimageResult<Array<T, D>>
where
T: Float + FromPrimitive + Debug + std::ops::AddAssign + std::ops::DivAssign + 'static,
D: Dimension + scirs2_core::ndarray::RemoveAxis + 'static,
usize: scirs2_core::ndarray::NdIndex<<D as scirs2_core::ndarray::Dimension>::Smaller>,
{
if input.ndim() == 0 {
return Err(NdimageError::InvalidInput(
"Input array cannot be 0-dimensional".into(),
));
}
let spline_order = order.unwrap_or(3);
if spline_order == 0 || spline_order > 5 {
return Err(NdimageError::InvalidInput(format!(
"Spline order must be between 1 and 5, got {}",
spline_order
)));
}
if spline_order <= 1 {
return Ok(input.to_owned());
}
let mut output = input.to_owned();
for axis in 0..input.ndim() {
spline_filter_axis(&mut output, spline_order, axis)?;
}
Ok(output)
}
#[allow(dead_code)]
pub fn spline_filter1d<T, D>(
input: &Array<T, D>,
order: Option<usize>,
axis: Option<usize>,
) -> NdimageResult<Array<T, D>>
where
T: Float + FromPrimitive + Debug + std::ops::AddAssign + std::ops::DivAssign + 'static,
D: Dimension + scirs2_core::ndarray::RemoveAxis + 'static,
usize: scirs2_core::ndarray::NdIndex<<D as scirs2_core::ndarray::Dimension>::Smaller>,
{
if input.ndim() == 0 {
return Err(NdimageError::InvalidInput(
"Input array cannot be 0-dimensional".into(),
));
}
let spline_order = order.unwrap_or(3);
let axis_val = axis.unwrap_or(0);
if spline_order == 0 || spline_order > 5 {
return Err(NdimageError::InvalidInput(format!(
"Spline order must be between 1 and 5, got {}",
spline_order
)));
}
if axis_val >= input.ndim() {
return Err(NdimageError::InvalidInput(format!(
"Axis {} is out of bounds for array of dimension {}",
axis_val,
input.ndim()
)));
}
if spline_order <= 1 {
return Ok(input.to_owned());
}
let mut output = input.to_owned();
spline_filter_axis(&mut output, spline_order, axis_val)?;
Ok(output)
}
#[allow(dead_code)]
pub fn bspline<T>(
positions: &Array<T, scirs2_core::ndarray::Ix1>,
order: Option<usize>,
derivative: Option<usize>,
) -> NdimageResult<Array<T, scirs2_core::ndarray::Ix1>>
where
T: Float + FromPrimitive + Debug,
{
let spline_order = order.unwrap_or(3);
let deriv = derivative.unwrap_or(0);
if spline_order == 0 || spline_order > 5 {
return Err(NdimageError::InvalidInput(format!(
"Spline order must be between 1 and 5, got {}",
spline_order
)));
}
if deriv > spline_order {
return Err(NdimageError::InvalidInput(format!(
"Derivative order must be less than or equal to spline order (got {} for order {})",
deriv, spline_order
)));
}
let mut result = Array1::<T>::zeros(positions.len());
for (i, &pos) in positions.iter().enumerate() {
result[i] = evaluate_bspline_basis(pos, spline_order, deriv);
}
Ok(result)
}
#[allow(dead_code)]
fn spline_filter_axis<T, D>(data: &mut Array<T, D>, order: usize, axis: usize) -> NdimageResult<()>
where
T: Float + FromPrimitive + Clone,
D: Dimension + scirs2_core::ndarray::RemoveAxis,
usize: scirs2_core::ndarray::NdIndex<<D as scirs2_core::ndarray::Dimension>::Smaller>,
{
let poles = get_spline_poles::<T>(order);
if poles.is_empty() {
return Ok(());
}
let tolerance = T::from_f64(1e-10).expect("Operation failed");
let axis_len = data.shape()[axis];
for mut lane in data.axis_iter_mut(Axis(axis)) {
let mut coeffs: Vec<T> = lane.iter().cloned().collect();
for &pole in &poles {
let initial_causal = get_initial_causal_coefficient(&coeffs, pole, tolerance);
apply_causal_filter(&mut coeffs, pole, initial_causal);
let initial_anti_causal = get_initial_anti_causal_coefficient(&coeffs, pole);
apply_anti_causal_filter(&mut coeffs, pole, initial_anti_causal);
}
for (i, &coeff) in coeffs.iter().enumerate() {
lane[i] = coeff;
}
}
Ok(())
}
#[allow(dead_code)]
fn bspline_knots<T: Float + FromPrimitive>(order: usize) -> Vec<T> {
let half_span = T::from_f64((order + 1) as f64 / 2.0).expect("Operation failed");
(0..=order + 1)
.map(|i| T::from_usize(i).expect("Operation failed") - half_span)
.collect()
}
#[allow(dead_code)]
fn cox_de_boor<T: Float + FromPrimitive>(i: usize, k: usize, x: T, knots: &[T]) -> T {
if k == 0 {
return if knots[i] <= x && x < knots[i + 1] {
T::one()
} else {
T::zero()
};
}
let mut result = T::zero();
let denom_left = knots[i + k] - knots[i];
if denom_left != T::zero() {
result = result + (x - knots[i]) / denom_left * cox_de_boor(i, k - 1, x, knots);
}
let denom_right = knots[i + k + 1] - knots[i + 1];
if denom_right != T::zero() {
result =
result + (knots[i + k + 1] - x) / denom_right * cox_de_boor(i + 1, k - 1, x, knots);
}
result
}
#[allow(dead_code)]
fn cox_de_boor_derivative<T: Float + FromPrimitive>(
i: usize,
k: usize,
x: T,
knots: &[T],
d: usize,
) -> T {
if d == 0 {
return cox_de_boor(i, k, x, knots);
}
if k == 0 {
return T::zero();
}
let k_t = T::from_usize(k).expect("Operation failed");
let mut result = T::zero();
let denom_left = knots[i + k] - knots[i];
if denom_left != T::zero() {
result = result + cox_de_boor_derivative(i, k - 1, x, knots, d - 1) / denom_left;
}
let denom_right = knots[i + k + 1] - knots[i + 1];
if denom_right != T::zero() {
result = result - cox_de_boor_derivative(i + 1, k - 1, x, knots, d - 1) / denom_right;
}
k_t * result
}
#[allow(dead_code)]
fn evaluate_bspline_basis<T: Float + FromPrimitive>(x: T, order: usize, derivative: usize) -> T {
if derivative > order {
return T::zero();
}
let knots: Vec<T> = bspline_knots(order);
cox_de_boor_derivative(0, order, x, &knots, derivative)
}
#[cfg(test)]
mod tests {
use super::*;
use scirs2_core::ndarray::{Array1, Array2};
#[test]
fn test_spline_filter() {
let input: Array2<f64> = Array2::eye(3);
let result = spline_filter(&input, None).expect("Operation failed");
assert_eq!(result.shape(), input.shape());
}
#[test]
fn test_spline_filter1d() {
let input: Array2<f64> = Array2::eye(3);
let result = spline_filter1d(&input, None, None).expect("Operation failed");
assert_eq!(result.shape(), input.shape());
}
#[test]
fn test_bspline() {
let positions = Array1::linspace(0.0, 2.0, 5);
let result = bspline(&positions, None, None).expect("Operation failed");
assert_eq!(result.len(), positions.len());
}
#[test]
fn test_bspline_orders_and_derivatives_vs_analytic() {
let cases: &[(usize, usize, f64, f64)] = &[
(1, 0, 0.3, 0.7),
(1, 0, 1.3, 0.0),
(1, 1, 0.3, -1.0),
(1, 1, 1.3, 0.0),
(2, 0, 0.3, 0.66),
(2, 0, 1.3, 0.02),
(2, 1, 0.3, -0.6),
(2, 1, 1.3, -0.2),
(2, 2, 0.3, -2.0),
(2, 2, 1.3, 1.0),
(3, 0, 0.3, 0.5901666667),
(3, 0, 1.3, 0.0571666667),
(3, 1, 0.3, -0.465),
(3, 1, 1.3, -0.245),
(3, 2, 0.3, -1.1),
(3, 2, 1.3, 0.7),
(3, 3, 0.3, 3.0),
(3, 3, 1.3, -1.0),
(4, 0, 0.3, 0.5447333333),
(4, 0, 1.3, 0.0860666667),
(4, 1, 0.3, -0.348),
(4, 1, 1.3, -0.2813333333),
(4, 2, 0.3, -0.98),
(4, 2, 1.3, 0.62),
(4, 3, 0.3, 1.8),
(4, 3, 1.3, -0.2),
(4, 4, 0.3, 6.0),
(4, 4, 1.3, -4.0),
(5, 0, 0.3, 0.5068225),
(5, 0, 1.3, 0.1099179167),
(5, 1, 0.3, -0.276375),
(5, 1, 1.3, -0.2879791667),
(5, 2, 0.3, -0.775),
(5, 2, 1.3, 0.4758333333),
(5, 3, 0.3, 1.35),
(5, 3, 1.3, 0.025),
(5, 4, 0.3, 3.0),
(5, 4, 1.3, -2.5),
(5, 5, 0.3, -10.0),
(5, 5, 1.3, 5.0),
];
for &(order, deriv, x, expected) in cases {
let positions = Array1::from_vec(vec![x]);
let result = bspline(&positions, Some(order), Some(deriv)).unwrap_or_else(|e| {
panic!("bspline(order={order}, deriv={deriv}, x={x}) failed: {e}")
});
assert!(
(result[0] - expected).abs() < 1e-9,
"order={order} deriv={deriv} x={x}: got {} expected {expected}",
result[0]
);
}
}
#[test]
fn test_bspline_is_symmetric_about_zero() {
for order in 1..=5usize {
let positions = Array1::from_vec(vec![0.3, -0.3, 1.3, -1.3]);
let result = bspline(&positions, Some(order), Some(0)).expect("Operation failed");
assert!((result[0] - result[1]).abs() < 1e-9, "order {order}");
assert!((result[2] - result[3]).abs() < 1e-9, "order {order}");
}
}
#[test]
fn test_bspline_rejects_invalid_order() {
let positions = Array1::from_vec(vec![0.0]);
assert!(bspline(&positions, Some(0), None).is_err());
assert!(bspline(&positions, Some(6), None).is_err());
}
#[test]
fn test_bspline_rejects_derivative_above_order() {
let positions = Array1::from_vec(vec![0.0]);
assert!(bspline(&positions, Some(2), Some(3)).is_err());
}
}