use crate::basis::basis_system::BasisSystem;
use crate::error::FdarError;
use crate::helpers::simpsons_weights;
use crate::smooth_basis::{differentiate_basis_columns, integrate_symmetric_penalty};
pub fn polygonal_basis(argvals: &[f64], knots: &[f64]) -> Result<BasisSystem, FdarError> {
let n = argvals.len();
if n < 2 {
return Err(FdarError::InvalidDimension {
parameter: "argvals",
expected: ">= 2".to_string(),
actual: n.to_string(),
});
}
let nbasis = knots.len();
if nbasis < 2 {
return Err(FdarError::InvalidParameter {
parameter: "knots",
message: "must have length >= 2".to_string(),
});
}
if !knots.windows(2).all(|w| w[1] > w[0]) {
return Err(FdarError::InvalidParameter {
parameter: "knots",
message: "knots must be strictly increasing (no duplicate or out-of-order values)"
.to_string(),
});
}
let mut eval_matrix = vec![0.0_f64; n * nbasis];
for (ti, &t) in argvals.iter().enumerate() {
for j in 0..nbasis {
eval_matrix[ti + j * n] = hat_function(t, knots, j);
}
}
let lfd_order = 1_usize;
let penalty_matrix = polygonal_penalty_numeric(argvals, knots, nbasis, lfd_order);
Ok(BasisSystem {
eval_matrix,
penalty_matrix,
nbasis,
n_eval: n,
lfd_order,
})
}
fn hat_function(t: f64, knots: &[f64], j: usize) -> f64 {
let nk = knots.len();
let at_last_knot = j == nk - 1;
let left = if j > 0 {
let in_left = if at_last_knot {
t >= knots[j - 1] && t <= knots[j]
} else {
t >= knots[j - 1] && t < knots[j]
};
if in_left {
(t - knots[j - 1]) / (knots[j] - knots[j - 1])
} else {
0.0
}
} else {
0.0
};
let right = if j + 1 < nk && t >= knots[j] && t <= knots[j + 1] {
(knots[j + 1] - t) / (knots[j + 1] - knots[j])
} else {
0.0
};
left + right
}
fn polygonal_penalty_numeric(
argvals: &[f64],
knots: &[f64],
nbasis: usize,
lfd_order: usize,
) -> Vec<f64> {
if argvals.len() < 2 {
return vec![0.0; nbasis * nbasis];
}
let t_min = knots[0];
let t_max = knots[knots.len() - 1];
let n_sub = 10;
let n_quad = (knots.len() - 1) * n_sub + 1;
let quad_t: Vec<f64> = (0..n_quad)
.map(|i| t_min + (t_max - t_min) * i as f64 / (n_quad - 1) as f64)
.collect();
let mut basis_fine = vec![0.0_f64; n_quad * nbasis];
for (ti, &t) in quad_t.iter().enumerate() {
for j in 0..nbasis {
basis_fine[ti + j * n_quad] = hat_function(t, knots, j);
}
}
let h = (t_max - t_min) / (n_quad - 1) as f64;
let deriv_basis = differentiate_basis_columns(&basis_fine, n_quad, nbasis, h, lfd_order);
let weights = simpsons_weights(&quad_t);
integrate_symmetric_penalty(&deriv_basis, &weights, nbasis, n_quad)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn poly_invalid_argvals_too_short() {
let result = polygonal_basis(&[0.5], &[0.0, 0.5, 1.0]);
assert!(matches!(result, Err(FdarError::InvalidDimension { .. })));
}
#[test]
fn poly_invalid_knots_too_few() {
let result = polygonal_basis(&[0.0, 1.0], &[0.5]);
assert!(matches!(result, Err(FdarError::InvalidParameter { .. })));
}
#[test]
fn poly_invalid_duplicate_knots() {
let result = polygonal_basis(&[0.0, 0.5, 1.0], &[0.0, 0.5, 0.5]);
assert!(matches!(result, Err(FdarError::InvalidParameter { .. })));
}
#[test]
fn poly_invalid_nonmonotone_knots() {
let result = polygonal_basis(&[0.0, 0.5, 1.0], &[0.0, 1.0, 0.5]);
assert!(matches!(result, Err(FdarError::InvalidParameter { .. })));
}
#[test]
fn poly_hat_peaks_at_knot() {
let knots = vec![0.0, 0.5, 1.0];
let argvals = knots.clone();
let bs = polygonal_basis(&argvals, &knots).unwrap();
let n = bs.n_eval;
for j in 0..bs.nbasis {
let val = bs.eval_matrix[j + j * n];
assert!(
(val - 1.0).abs() < 1e-12,
"B_{j}(knots[{j}])={val}, expected 1.0"
);
}
}
#[test]
fn poly_midpoint_values() {
let knots = vec![0.0, 0.5, 1.0];
let argvals = vec![0.0, 0.25, 0.5, 0.75, 1.0];
let bs = polygonal_basis(&argvals, &knots).unwrap();
let n = bs.n_eval; let b0 = bs.eval_matrix[1]; let b1 = bs.eval_matrix[1 + n]; let b2 = bs.eval_matrix[1 + 2 * n]; assert!((b0 - 0.5).abs() < 1e-12, "B₀(0.25)={b0}");
assert!((b1 - 0.5).abs() < 1e-12, "B₁(0.25)={b1}");
assert!(b2.abs() < 1e-12, "B₂(0.25)={b2}");
}
#[test]
fn poly_eval_at_interior_knot() {
let knots = vec![0.0, 0.5, 1.0];
let argvals = vec![0.0, 0.25, 0.5, 0.75, 1.0];
let bs = polygonal_basis(&argvals, &knots).unwrap();
let n = bs.n_eval;
let b0 = bs.eval_matrix[2];
let b1 = bs.eval_matrix[2 + n];
let b2 = bs.eval_matrix[2 + 2 * n];
assert!(b0.abs() < 1e-12, "B₀(0.5)={b0}");
assert!((b1 - 1.0).abs() < 1e-12, "B₁(0.5)={b1}");
assert!(b2.abs() < 1e-12, "B₂(0.5)={b2}");
}
#[test]
fn poly_partition_of_unity() {
let knots = vec![0.0, 0.5, 1.0];
let argvals: Vec<f64> = (0..=20).map(|i| i as f64 / 20.0).collect();
let bs = polygonal_basis(&argvals, &knots).unwrap();
let n = bs.n_eval;
for ti in 0..n {
let sum: f64 = (0..bs.nbasis).map(|j| bs.eval_matrix[ti + j * n]).sum();
assert!(
(sum - 1.0).abs() < 1e-12,
"partition-of-unity violated at ti={ti}: sum={sum}"
);
}
}
#[test]
fn poly_shape_invariants() {
let knots = vec![0.0, 0.25, 0.5, 0.75, 1.0];
let argvals: Vec<f64> = (0..=10).map(|i| i as f64 / 10.0).collect();
let bs = polygonal_basis(&argvals, &knots).unwrap();
assert_eq!(bs.eval_matrix.len(), 11 * 5);
assert_eq!(bs.penalty_matrix.len(), 5 * 5);
assert_eq!(bs.nbasis, 5);
assert_eq!(bs.n_eval, 11);
assert_eq!(bs.lfd_order, 1);
}
#[test]
fn poly_penalty_symmetric() {
let knots = vec![0.0, 0.5, 1.0];
let argvals: Vec<f64> = (0..=10).map(|i| i as f64 / 10.0).collect();
let bs = polygonal_basis(&argvals, &knots).unwrap();
let k = bs.nbasis;
for j in 0..k {
for l in 0..k {
let pjl = bs.penalty_matrix[j + l * k];
let plj = bs.penalty_matrix[l + j * k];
assert!(
(pjl - plj).abs() < 1e-10,
"P[{j},{l}]={pjl} != P[{l},{j}]={plj}"
);
}
}
}
#[test]
fn poly_penalty_diagonal_nonneg() {
let knots = vec![0.0, 0.5, 1.0];
let argvals: Vec<f64> = (0..=10).map(|i| i as f64 / 10.0).collect();
let bs = polygonal_basis(&argvals, &knots).unwrap();
let k = bs.nbasis;
for j in 0..k {
let diag = bs.penalty_matrix[j + j * k];
assert!(diag >= -1e-10, "P[{j},{j}]={diag} is negative");
}
}
#[test]
fn poly_basis_system_derives() {
let knots = vec![0.0, 1.0];
let bs = polygonal_basis(&knots, &knots).unwrap();
let bs2 = bs.clone();
assert_eq!(bs, bs2);
let _ = format!("{bs:?}");
}
}