use nalgebra::DMatrix;
use std::f64;
const MAX_DIMS: usize = 64;
fn evaluate_poly1(s: f64, c: &[Vec<Vec<f64>>], ci: usize, cj: usize, dx: i32) -> f64 {
let mut res = 0.0;
let mut z = 1.0;
let k = c.len();
if dx < 0 {
for _ in 0..(-dx) {
z *= s;
}
}
for kp in 0..k {
let prefactor = if dx == 0 {
1.0
} else if dx > 0 {
if kp < (dx as usize) {
continue; } else {
let mut pref = 1.0;
for k_val in (kp - (dx as usize) + 1)..=kp {
pref *= k_val as f64;
}
pref
}
} else {
let mut pref = 1.0;
for k_val in (kp + 1)..=(kp + (-dx as usize)) {
pref /= k_val as f64;
}
pref
};
res += c[k - kp - 1][ci][cj] * z * prefactor;
if kp < k - 1 && (kp as i32) >= dx {
z *= s;
}
}
res
}
fn find_interval_ascending(x: &[f64], xval: f64, prev_interval: usize, extrapolate: bool) -> i32 {
let n = x.len();
if xval.is_nan() {
return -1;
}
if xval < x[0] {
return if extrapolate { 0 } else { -1 };
}
if xval > x[n - 1] {
return if extrapolate { (n - 2) as i32 } else { -1 };
}
if xval == x[n - 1] {
return (n - 2) as i32;
}
let mut low = if prev_interval < n - 1 {
prev_interval
} else {
0
};
let mut high = n - 1;
if low < n - 1 && x[low] <= xval && xval < x[low + 1] {
return low as i32;
}
while high - low > 1 {
let mid = (high + low) / 2;
if xval < x[mid] {
high = mid;
} else {
low = mid;
}
}
low as i32
}
fn find_interval_descending(x: &[f64], xval: f64, prev_interval: usize, extrapolate: bool) -> i32 {
let n = x.len();
if xval.is_nan() {
return -1;
}
if xval > x[0] {
return if extrapolate { 0 } else { -1 };
}
if xval <= x[n - 1] {
return if extrapolate { (n - 2) as i32 } else { -1 };
}
let mut low = if prev_interval < n - 1 {
prev_interval
} else {
0
};
let mut high = n - 1;
if low < n - 1 && x[low] >= xval && xval > x[low + 1] {
return low as i32;
}
while high - low > 1 {
let mid = (high + low) / 2;
if xval > x[mid] {
high = mid;
} else {
low = mid;
}
}
low as i32
}
pub fn evaluate_nd(
c: &[Vec<Vec<f64>>],
xs: &[Vec<f64>],
ks: &[usize],
xp: &[Vec<f64>],
dx: &[i32],
extrapolate: bool,
out: &mut [Vec<f64>],
) -> Result<(), String> {
let ndim = xs.len();
if ndim > MAX_DIMS {
return Err(format!("Too many dimensions (maximum: {})", MAX_DIMS));
}
if dx.len() != ndim {
return Err("dx has incompatible shape".to_string());
}
if xp.len() > 0 && xp[0].len() != ndim {
return Err("xp has incompatible shape".to_string());
}
if out.len() != xp.len() {
return Err("out and xp have incompatible shapes".to_string());
}
if out.len() > 0 && out[0].len() != c[0][0].len() {
return Err("out and c have incompatible shapes".to_string());
}
for &d in dx {
if d < 0 {
return Err("Order of derivative cannot be negative".to_string());
}
}
for (_i, x_dim) in xs.iter().enumerate() {
if x_dim.len() < 2 {
return Err("each dimension must have >= 2 points".to_string());
}
}
let mut strides = vec![0; ndim];
let mut ntot = 1;
for i in (0..ndim).rev() {
strides[i] = ntot;
ntot *= xs[i].len() - 1;
}
if c[0].len() != ntot {
return Err("xs and c have incompatible shapes".to_string());
}
let mut kstrides = vec![0; ndim];
let mut ktot = 1;
for i in 0..ndim {
kstrides[i] = ktot;
ktot *= ks[i];
}
if c.len() != ktot {
return Err("ks and c have incompatible shapes".to_string());
}
let mut c2 = vec![vec![vec![0.0; 1]; 1]; c.len()];
let mut intervals = vec![0; ndim];
for (ip, xp_point) in xp.iter().enumerate() {
let mut out_of_range = false;
for k in 0..ndim {
let xval = xp_point[k];
let interval = find_interval_ascending(&xs[k], xval, intervals[k], extrapolate);
if interval < 0 {
out_of_range = true;
break;
} else {
intervals[k] = interval as usize;
}
}
if out_of_range {
for jp in 0..out[ip].len() {
out[ip][jp] = f64::NAN;
}
continue;
}
let mut pos = 0;
for k in 0..ndim {
pos += intervals[k] * strides[k];
}
for jp in 0..out[ip].len() {
for i in 0..c.len() {
c2[i][0][0] = c[i][pos][jp];
}
for k in (0..ndim).rev() {
let xval = xp_point[k] - xs[k][intervals[k]];
let mut kpos = 0;
for koutpos in 0..kstrides[k] {
let slice_start = kpos;
let _slice_end = kpos + ks[k];
let mut c_slice = vec![vec![vec![0.0]]];
c_slice.resize(ks[k], vec![vec![0.0]]);
for idx in 0..ks[k] {
c_slice[idx][0][0] = c2[slice_start + idx][0][0];
}
c2[koutpos][0][0] = evaluate_poly1(xval, &c_slice, 0, 0, dx[k]);
kpos += ks[k];
}
}
out[ip][jp] = c2[0][0][0];
}
}
Ok(())
}
pub fn evaluate(
c: &[Vec<Vec<f64>>],
x: &[f64],
xp: &[f64],
dx: i32,
extrapolate: bool,
out: &mut [Vec<f64>],
) -> Result<(), String> {
if dx < 0 {
return Err("Order of derivative cannot be negative".to_string());
}
if out.len() != xp.len() {
return Err("out and xp have incompatible shapes".to_string());
}
if out.len() > 0 && out[0].len() != c[0][0].len() {
return Err("out and c have incompatible shapes".to_string());
}
if c.len() > 0 && c[0].len() != x.len() - 1 {
return Err("x and c have incompatible shapes".to_string());
}
let mut interval = 0;
let ascending = x[x.len() - 1] >= x[0];
for (ip, &xval) in xp.iter().enumerate() {
let i = if ascending {
find_interval_ascending(x, xval, interval, extrapolate)
} else {
find_interval_descending(x, xval, interval, extrapolate)
};
if i < 0 {
for jp in 0..out[ip].len() {
out[ip][jp] = f64::NAN;
}
continue;
} else {
interval = i as usize;
}
for jp in 0..out[ip].len() {
out[ip][jp] = evaluate_poly1(xval - x[interval], c, interval, jp, dx);
}
}
Ok(())
}
#[derive(Clone, Debug)]
pub struct PPoly {
pub c: Vec<Vec<Vec<f64>>>,
pub x: Vec<f64>,
pub extrapolate: Extrapolate,
pub axis: usize,
}
#[derive(Clone, Debug)]
pub enum Extrapolate {
Bool(bool),
Periodic,
}
impl PPoly {
pub fn call(
&self,
x_eval: &[f64], x_shape: &[usize], nu: Option<i32>,
extrapolate: Option<Extrapolate>,
) -> DMatrix<f64> {
let extrapolate = extrapolate.unwrap_or_else(|| self.extrapolate.clone());
let mut x_flat: Vec<f64> = x_eval.to_vec();
let extrap_bool = match &extrapolate {
Extrapolate::Bool(b) => *b,
Extrapolate::Periodic => {
let x0 = self.x[0];
let x1 = self.x[self.x.len() - 1];
let period = x1 - x0;
for x in &mut x_flat {
*x = x0 + (*x - x0).rem_euclid(period);
}
false
}
};
let dx = nu.unwrap_or(0);
let r = x_flat.len();
let trailing_dims = self.c[0][0].len();
let mut out = vec![vec![0.0; trailing_dims]; r];
let _ = evaluate(&self.c, &self.x, &x_flat, dx, extrap_bool, &mut out);
let mut full_shape = x_shape.to_vec();
full_shape.push(trailing_dims);
let flat_out: Vec<f64> = out.into_iter().flat_map(|v| v.into_iter()).collect();
if self.axis == 0 {
return DMatrix::from_row_slice(r, trailing_dims, &flat_out);
}
if self.axis == 1 && x_shape.len() == 1 {
return DMatrix::from_row_slice(x_shape[0], trailing_dims, &flat_out);
}
DMatrix::from_row_slice(r, trailing_dims, &flat_out)
}
}
#[cfg(test)]
mod tests_PPoly {
use super::*;
use approx::assert_relative_eq;
fn make_ppoly_linear() -> PPoly {
let c = vec![
vec![vec![2.0]], vec![vec![1.0]], ];
let x = vec![0.0, 1.0];
PPoly {
c,
x,
extrapolate: Extrapolate::Bool(true),
axis: 0,
}
}
fn make_ppoly_quadratic() -> PPoly {
let c = vec![
vec![vec![1.0], vec![1.0]], vec![vec![-2.0], vec![-2.0]], vec![vec![1.0], vec![1.0]], ];
let x = vec![0.0, 1.0, 2.0];
PPoly {
c,
x,
extrapolate: Extrapolate::Bool(true),
axis: 0,
}
}
fn make_ppoly_periodic() -> PPoly {
let c = vec![
vec![vec![1.0]], vec![vec![0.0]], ];
let x = vec![0.0, 1.0];
PPoly {
c,
x,
extrapolate: Extrapolate::Periodic,
axis: 0,
}
}
fn make_ppoly_multidim() -> PPoly {
let c = vec![
vec![vec![1.0, 1.0]], vec![vec![0.0, 1.0]], ];
let x = vec![0.0, 1.0];
PPoly {
c,
x,
extrapolate: Extrapolate::Bool(true),
axis: 0,
}
}
#[test]
fn test_ppoly_linear_basic() {
let ppoly = make_ppoly_linear();
let x_eval = vec![0.0, 0.5, 1.0, 1.5];
let x_shape = &[4];
let result = ppoly.call(&x_eval, x_shape, Some(0), None);
let expected = vec![1.0, 2.0, 3.0, 4.0];
for i in 0..x_eval.len() {
assert_relative_eq!(result[(i, 0)], expected[i], epsilon = 1e-10);
}
}
#[test]
fn test_ppoly_quadratic_basic() {
let ppoly = make_ppoly_quadratic();
let x_eval = vec![0.0, 0.5, 1.0, 1.5, 2.0];
let x_shape = &[5];
let result = ppoly.call(&x_eval, x_shape, Some(0), None);
println!("result {:?}", result);
let expected = vec![1.0, 0.25, 1.0, 0.25, 0.0];
for i in 0..x_eval.len() {
assert_relative_eq!(result[(i, 0)], expected[i], epsilon = 1e-10);
}
}
#[test]
fn test_ppoly_periodic_extrapolation() {
let ppoly = make_ppoly_periodic();
let x_eval = vec![-0.5, 0.0, 0.5, 1.0, 1.5, 2.0];
let x_shape = &[6];
let result = ppoly.call(&x_eval, x_shape, Some(0), None);
let expected = vec![0.5, 0.0, 0.5, 0.0, 0.5, 0.0];
for i in 0..x_eval.len() {
assert_relative_eq!(result[(i, 0)], expected[i], epsilon = 1e-10);
}
}
#[test]
fn test_ppoly_multidim_trailing_dims() {
let ppoly = make_ppoly_multidim();
let x_eval = vec![0.0, 0.5, 1.0];
let x_shape = &[3];
let result = ppoly.call(&x_eval, x_shape, Some(0), None);
let expected = vec![vec![0.0, 1.0], vec![0.5, 1.5], vec![1.0, 2.0]];
for i in 0..x_eval.len() {
for j in 0..2 {
assert_relative_eq!(result[(i, j)], expected[i][j], epsilon = 1e-10);
}
}
}
#[test]
fn test_ppoly_axis_permutation() {
let ppoly = make_ppoly_multidim();
let x_eval = vec![0.0, 0.5, 1.0];
let x_shape = &[3];
let mut ppoly_axis = ppoly.clone();
ppoly_axis.axis = 1;
let result = ppoly_axis.call(&x_eval, x_shape, Some(0), None);
let expected = vec![vec![0.0, 1.0], vec![0.5, 1.5], vec![1.0, 2.0]];
for i in 0..x_eval.len() {
for j in 0..2 {
assert_relative_eq!(result[(i, j)], expected[i][j], epsilon = 1e-10);
}
}
}
#[test]
fn test_ppoly_derivative() {
let c = vec![
vec![vec![1.0]], vec![vec![0.0]], vec![vec![0.0]], ];
let x = vec![0.0, 1.0];
let ppoly = PPoly {
c,
x,
extrapolate: Extrapolate::Bool(true),
axis: 0,
};
let x_eval = vec![0.0, 0.5, 1.0];
let x_shape = &[3];
let result = ppoly.call(&x_eval, x_shape, Some(1), None);
let expected = vec![0.0, 1.0, 2.0];
for i in 0..x_eval.len() {
assert_relative_eq!(result[(i, 0)], expected[i], epsilon = 1e-10);
}
}
#[test]
fn test_ppoly_out_of_bounds_nan() {
let ppoly = make_ppoly_linear();
let mut ppoly_no_extrap = ppoly.clone();
ppoly_no_extrap.extrapolate = Extrapolate::Bool(false);
let x_eval = vec![-1.0, 0.0, 0.5, 1.0, 2.0];
let x_shape = &[5];
let result: DMatrix<f64> = ppoly_no_extrap.call(&x_eval, x_shape, Some(0), None);
println!("{:?}", result);
assert!(result[(0, 0)].is_nan());
assert!(!result[(1, 0)].is_nan());
assert!(!result[(2, 0)].is_nan());
assert!(!result[(3, 0)].is_nan());
assert!(result[(4, 0)].is_nan());
}
#[test]
fn test_ppoly_multi_dimensional_x_shape() {
let ppoly = make_ppoly_linear();
let x_eval = vec![0.0, 0.5, 1.0, 1.5];
let x_shape = &[2, 2];
let result = ppoly.call(&x_eval, x_shape, Some(0), None);
let expected = vec![1.0, 2.0, 3.0, 4.0];
for i in 0..x_eval.len() {
assert_relative_eq!(result[(i, 0)], expected[i], epsilon = 1e-10);
}
}
#[test]
fn test_evaluate_poly1_basic() {
let c = vec![
vec![vec![2.0]], vec![vec![3.0]], vec![vec![1.0]], ];
let result = evaluate_poly1(2.0, &c, 0, 0, 0);
let expected = 2.0 * 4.0 + 3.0 * 2.0 + 1.0; assert_relative_eq!(result, expected, epsilon = 1e-10);
let result_dx1 = evaluate_poly1(2.0, &c, 0, 0, 1);
let expected_dx1 = 4.0 * 2.0 + 3.0; assert_relative_eq!(result_dx1, expected_dx1, epsilon = 1e-10);
let result_dx2 = evaluate_poly1(2.0, &c, 0, 0, 2);
let expected_dx2 = 4.0;
assert_relative_eq!(result_dx2, expected_dx2, epsilon = 1e-10);
}
#[test]
fn test_evaluate_nd_1d() {
let c = vec![
vec![vec![1.0], vec![1.0]], vec![vec![2.0], vec![2.0]], vec![vec![3.0], vec![3.0]], ];
let xs = vec![vec![0.0, 1.0, 2.0]]; let ks = vec![3]; let xp = vec![vec![0.5], vec![1.5]]; let dx = vec![0];
let mut out = vec![vec![0.0]; 2];
let result = evaluate_nd(&c, &xs, &ks, &xp, &dx, true, &mut out);
assert!(result.is_ok());
assert_relative_eq!(out[0][0], 4.25, epsilon = 1e-10);
assert_relative_eq!(out[1][0], 4.25, epsilon = 1e-10);
}
#[test]
fn test_find_interval_ascending() {
let x = vec![0.0, 1.0, 2.0, 3.0];
assert_eq!(find_interval_ascending(&x, 0.5, 0, true), 0);
assert_eq!(find_interval_ascending(&x, 1.5, 0, true), 1);
assert_eq!(find_interval_ascending(&x, 2.5, 0, true), 2);
assert_eq!(find_interval_ascending(&x, 0.0, 0, true), 0);
assert_eq!(find_interval_ascending(&x, 3.0, 0, true), 2);
assert_eq!(find_interval_ascending(&x, -0.5, 0, true), 0);
assert_eq!(find_interval_ascending(&x, 3.5, 0, true), 2);
assert_eq!(find_interval_ascending(&x, -0.5, 0, false), -1);
assert_eq!(find_interval_ascending(&x, 3.5, 0, false), -1);
assert_eq!(find_interval_ascending(&x, f64::NAN, 0, true), -1);
}
#[test]
fn test_evaluate_1d_piecewise() {
let c = vec![
vec![vec![1.0], vec![2.0]], vec![vec![1.0], vec![1.0]], vec![vec![1.0], vec![0.5]], ];
let x = vec![0.0, 1.0, 2.0]; let xp = vec![0.5, 1.5]; let dx = 0;
let mut out = vec![vec![0.0]; 2];
let result = evaluate(&c, &x, &xp, dx, true, &mut out);
assert!(result.is_ok());
assert_relative_eq!(out[0][0], 1.75, epsilon = 1e-10);
assert_relative_eq!(out[1][0], 1.5, epsilon = 1e-10);
}
#[test]
fn test_evaluate_with_derivatives() {
let c = vec![
vec![vec![1.0]], vec![vec![2.0]], vec![vec![3.0]], vec![vec![4.0]], ];
let x = vec![0.0, 2.0]; let xp = vec![1.0];
let mut out0 = vec![vec![0.0]; 1];
let result = evaluate(&c, &x, &xp, 0, true, &mut out0);
assert!(result.is_ok());
assert_relative_eq!(out0[0][0], 10.0, epsilon = 1e-10);
let mut out1 = vec![vec![0.0]; 1];
let result = evaluate(&c, &x, &xp, 1, true, &mut out1);
assert!(result.is_ok());
assert_relative_eq!(out1[0][0], 10.0, epsilon = 1e-10);
let mut out2 = vec![vec![0.0]; 1];
let result = evaluate(&c, &x, &xp, 2, true, &mut out2);
assert!(result.is_ok());
assert_relative_eq!(out2[0][0], 10.0, epsilon = 1e-10);
}
#[test]
fn test_find_interval_descending() {
let x = vec![3.0, 2.0, 1.0, 0.0];
assert_eq!(find_interval_descending(&x, 2.5, 0, true), 0);
assert_eq!(find_interval_descending(&x, 1.5, 0, true), 1);
assert_eq!(find_interval_descending(&x, 0.5, 0, true), 2);
assert_eq!(find_interval_descending(&x, 3.0, 0, true), 0);
assert_eq!(find_interval_descending(&x, 0.0, 0, true), 2);
assert_eq!(find_interval_descending(&x, 3.5, 0, true), 0);
assert_eq!(find_interval_descending(&x, -0.5, 0, true), 2);
assert_eq!(find_interval_descending(&x, 3.5, 0, false), -1);
assert_eq!(find_interval_descending(&x, -0.5, 0, false), -1);
assert_eq!(find_interval_descending(&x, f64::NAN, 0, true), -1);
}
#[test]
fn test_ppoly_axis1_linear() {
let c = vec![
vec![vec![3.0]], vec![vec![2.0]], ];
let x = vec![0.0, 1.0];
let ppoly = PPoly {
c,
x,
extrapolate: Extrapolate::Bool(true),
axis: 1,
};
let x_eval = vec![0.0, 0.5, 1.0];
let x_shape = &[3];
let result = ppoly.call(&x_eval, x_shape, Some(0), None);
println!("result {:?}", result);
let expected = vec![2.0, 3.5, 5.0];
for i in 0..x_eval.len() {
assert_relative_eq!(result[(i, 0)], expected[i], epsilon = 1e-10);
}
}
#[test]
fn test_ppoly_axis1_quadratic_and_derivative() {
let c = vec![
vec![vec![2.0]], vec![vec![3.0]], vec![vec![1.0]], ];
let x = vec![0.0, 2.0];
let ppoly = PPoly {
c,
x,
extrapolate: Extrapolate::Bool(true),
axis: 1,
};
let x_eval = vec![0.0, 1.0, 2.0];
let x_shape = &[3];
let result = ppoly.call(&x_eval, x_shape, Some(0), None);
println!("result {:?}", result);
let expected = vec![1.0, 6.0, 15.0];
for i in 0..x_eval.len() {
assert_relative_eq!(result[(i, 0)], expected[i], epsilon = 1e-10);
}
let result_deriv = ppoly.call(&x_eval, x_shape, Some(1), None);
let expected_deriv = vec![3.0, 7.0, 11.0];
for i in 0..x_eval.len() {
assert_relative_eq!(result_deriv[(i, 0)], expected_deriv[i], epsilon = 1e-10);
}
}
#[test]
fn test_ppoly_axis1_multidim() {
let c = vec![
vec![vec![0.0, 1.0]], vec![vec![1.0, 0.0]], ];
let x = vec![0.0, 1.0];
let ppoly = PPoly {
c,
x,
extrapolate: Extrapolate::Bool(true),
axis: 1,
};
let x_eval = vec![0.0, 0.5, 1.0];
let x_shape = &[3];
let result = ppoly.call(&x_eval, x_shape, Some(0), None);
let expected = vec![
vec![1.0, 0.0], vec![1.0, 0.5], vec![1.0, 1.0], ];
println!("result {:?}", result);
for i in 0..x_eval.len() {
for j in 0..2 {
assert_relative_eq!(result[(i, j)], expected[i][j], epsilon = 1e-10);
}
}
}
#[test]
fn test_ppoly_axis1_known_points_and_derivative() {
let c = vec![
vec![vec![1.0]], vec![vec![-2.0]], vec![vec![1.0]], vec![vec![-1.0]], ];
let x = vec![0.0, 2.0];
let ppoly = PPoly {
c,
x,
extrapolate: Extrapolate::Bool(true),
axis: 1,
};
let x_eval = vec![0.0, 1.0, 2.0];
let x_shape = &[3];
let expected = vec![-1.0, -1.0, 1.0];
let result: DMatrix<f64> = ppoly.call(&x_eval, x_shape, Some(0), None);
for i in 0..x_eval.len() {
assert_relative_eq!(result[(i, 0)], expected[i], epsilon = 1e-10);
}
let expected_deriv = vec![1.0, 0.0, 5.0];
let result_deriv = ppoly.call(&x_eval, x_shape, Some(1), None);
for i in 0..x_eval.len() {
assert_relative_eq!(result_deriv[(i, 0)], expected_deriv[i], epsilon = 1e-10);
}
}
#[test]
fn test_ppoly_high_order_polynomial() {
let c = vec![
vec![vec![1.0]], vec![vec![-3.0]], vec![vec![2.0]], vec![vec![-1.0]], vec![vec![4.0]], vec![vec![-1.0]], ];
let x = vec![0.0, 2.0];
let ppoly = PPoly {
c,
x,
extrapolate: Extrapolate::Bool(true),
axis: 0,
};
let x_eval = vec![0.0, 0.5, 1.0, 1.5, 2.0];
let x_shape = &[5];
let result = ppoly.call(&x_eval, x_shape, Some(0), None);
let expected = vec![
-1.0, 0.84375, 2.0, 1.90625, 3.0, ];
for i in 0..x_eval.len() {
assert_relative_eq!(result[(i, 0)], expected[i], epsilon = 1e-5);
}
}
#[test]
fn test_ppoly_discontinuous_piecewise() {
let c = vec![
vec![vec![1.0], vec![-1.0], vec![0.0]],
vec![vec![0.0], vec![4.0], vec![2.0]],
vec![vec![1.0], vec![-2.0], vec![-3.0]],
];
let x = vec![0.0, 1.0, 2.0, 3.0];
let ppoly = PPoly {
c,
x,
extrapolate: Extrapolate::Bool(true),
axis: 0,
};
let x_eval = vec![0.5, 1.0, 1.5, 2.0, 2.5];
let x_shape = &[5];
let result = ppoly.call(&x_eval, x_shape, Some(0), None);
let expected = vec![
1.25, -2.0, -0.25, -3.0, -2.0, ];
for i in 0..x_eval.len() {
assert_relative_eq!(result[(i, 0)], expected[i], epsilon = 1e-10);
}
}
#[test]
fn test_ppoly_multiple_derivatives() {
let c = vec![
vec![vec![1.0]], vec![vec![-2.0]], vec![vec![1.0]], vec![vec![-3.0]], vec![vec![5.0]], ];
let x = vec![0.0, 3.0];
let ppoly = PPoly {
c,
x,
extrapolate: Extrapolate::Bool(true),
axis: 0,
};
let x_eval = vec![1.0, 2.0];
let x_shape = &[2];
let result0 = ppoly.call(&x_eval, x_shape, Some(0), None);
let expected0 = vec![2.0, 3.0]; for i in 0..x_eval.len() {
assert_relative_eq!(result0[(i, 0)], expected0[i], epsilon = 1e-10);
}
let result1 = ppoly.call(&x_eval, x_shape, Some(1), None);
let expected1 = vec![-3.0, 9.0]; for i in 0..x_eval.len() {
assert_relative_eq!(result1[(i, 0)], expected1[i], epsilon = 1e-10);
}
let result2 = ppoly.call(&x_eval, x_shape, Some(2), None);
let expected2 = vec![2.0, 26.0]; for i in 0..x_eval.len() {
assert_relative_eq!(result2[(i, 0)], expected2[i], epsilon = 1e-10);
}
}
#[test]
fn test_ppoly_stress_many_intervals() {
let n_intervals = 100;
let mut c = vec![vec![vec![]; n_intervals]; 2]; let mut x = vec![0.0];
for i in 0..n_intervals {
x.push((i + 1) as f64);
if i % 2 == 0 {
c[0][i] = vec![1.0]; c[1][i] = vec![0.0]; } else {
c[0][i] = vec![-1.0]; c[1][i] = vec![2.0]; }
}
let ppoly = PPoly {
c,
x,
extrapolate: Extrapolate::Bool(true),
axis: 0,
};
let x_eval = vec![0.5, 1.5, 2.5, 50.5, 99.5];
let x_shape = &[5];
let result = ppoly.call(&x_eval, x_shape, Some(0), None);
let expected = vec![0.5, 1.5, 0.5, 0.5, 1.5];
for i in 0..x_eval.len() {
assert_relative_eq!(result[(i, 0)], expected[i], epsilon = 1e-10);
}
}
#[test]
fn test_ppoly_periodic_complex() {
let c = vec![
vec![vec![-1.0 / 6.0]], vec![vec![0.0]], vec![vec![1.0]], vec![vec![0.0]], ];
let pi_half = std::f64::consts::PI / 2.0;
let x = vec![0.0, pi_half];
let ppoly = PPoly {
c,
x,
extrapolate: Extrapolate::Periodic,
axis: 0,
};
let x_eval = vec![
0.0,
pi_half / 2.0,
pi_half,
pi_half + 0.5,
2.0 * pi_half + 0.5,
];
let x_shape = &[5];
let result = ppoly.call(&x_eval, x_shape, Some(0), None);
let mid_val = pi_half / 2.0 - (pi_half / 2.0).powi(3) / 6.0;
let expected = vec![
0.0,
mid_val,
0.0,
0.5 - 0.5_f64.powi(3) / 6.0,
0.5 - 0.5_f64.powi(3) / 6.0,
];
for i in 0..x_eval.len() {
assert_relative_eq!(result[(i, 0)], expected[i], epsilon = 1e-10);
}
}
#[test]
fn test_ppoly_multidim_complex() {
let c = vec![
vec![vec![1.0, 0.0, -1.0]], vec![vec![0.0, 2.0, 3.0]], vec![vec![0.0, 1.0, 0.0]], ];
let x = vec![0.0, 2.0];
let ppoly = PPoly {
c,
x,
extrapolate: Extrapolate::Bool(true),
axis: 0,
};
let x_eval = vec![0.0, 1.0, 2.0];
let x_shape = &[3];
let result = ppoly.call(&x_eval, x_shape, Some(0), None);
let expected = vec![
vec![0.0, 1.0, 0.0], vec![1.0, 3.0, 2.0], vec![4.0, 5.0, 2.0], ];
for i in 0..x_eval.len() {
for j in 0..3 {
assert_relative_eq!(result[(i, j)], expected[i][j], epsilon = 1e-10);
}
}
}
#[test]
fn test_ppoly_edge_case_single_point() {
let ppoly = make_ppoly_quadratic();
let x_eval = vec![1.0];
let x_shape = &[1];
let result = ppoly.call(&x_eval, x_shape, Some(0), None);
assert_relative_eq!(result[(0, 0)], 1.0, epsilon = 1e-10);
}
#[test]
fn test_ppoly_edge_case_boundary_values() {
let ppoly = make_ppoly_quadratic();
let x_eval = vec![0.0, 1.0, 2.0]; let x_shape = &[3];
let result = ppoly.call(&x_eval, x_shape, Some(0), None);
let expected = vec![1.0, 1.0, 0.0]; for i in 0..x_eval.len() {
assert_relative_eq!(result[(i, 0)], expected[i], epsilon = 1e-10);
}
}
#[test]
fn test_ppoly_numerical_stability() {
let c = vec![
vec![vec![1e-15]], vec![vec![1e15]], vec![vec![1.0]], ];
let x = vec![0.0, 1.0];
let ppoly = PPoly {
c,
x,
extrapolate: Extrapolate::Bool(true),
axis: 0,
};
let x_eval = vec![0.5];
let x_shape = &[1];
let result = ppoly.call(&x_eval, x_shape, Some(0), None);
let expected = 1e-15 * 0.25 + 1e15 * 0.5 + 1.0;
assert_relative_eq!(result[(0, 0)], expected, epsilon = 1e-10);
}
#[test]
fn test_evaluate_nd_2d() {
let c = vec![
vec![vec![1.0]],
vec![vec![1.0]],
vec![vec![1.0]],
vec![vec![1.0]],
];
let xs = vec![vec![0.0, 1.0], vec![0.0, 1.0]]; let ks = vec![2, 2]; let xp = vec![vec![0.5, 0.5]]; let dx = vec![0, 0];
let mut out = vec![vec![0.0]; 1];
let result = evaluate_nd(&c, &xs, &ks, &xp, &dx, true, &mut out);
assert!(result.is_ok());
assert_relative_eq!(out[0][0], 2.25, epsilon = 1e-10);
}
#[test]
fn test_error_handling_comprehensive() {
let c = vec![vec![vec![1.0]]];
let xs = vec![vec![0.0, 1.0]];
let ks = vec![1];
let xp = vec![vec![0.5]];
let dx = vec![0];
let mut out = vec![vec![0.0]];
let dx_neg = vec![-1];
let result = evaluate_nd(&c, &xs, &ks, &xp, &dx_neg, true, &mut out);
assert!(result.is_err());
assert!(result.unwrap_err().contains("negative"));
let dx_wrong = vec![0, 0]; let result = evaluate_nd(&c, &xs, &ks, &xp, &dx_wrong, true, &mut out);
assert!(result.is_err());
assert!(result.unwrap_err().contains("incompatible"));
let xs_short = vec![vec![0.0]]; let result = evaluate_nd(&c, &xs_short, &ks, &xp, &dx, true, &mut out);
assert!(result.is_err());
assert!(result.unwrap_err().contains("2 points"));
}
}