use crate::basis::basis_system::BasisSystem;
use crate::error::FdarError;
pub fn monomial_basis(argvals: &[f64], nbasis: usize) -> Result<BasisSystem, FdarError> {
let n = argvals.len();
if n < 2 {
return Err(FdarError::InvalidDimension {
parameter: "argvals",
expected: ">= 2".to_string(),
actual: n.to_string(),
});
}
if nbasis < 1 {
return Err(FdarError::InvalidParameter {
parameter: "nbasis",
message: "must be >= 1".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] = t.powi(j as i32);
}
}
let lfd_order = 2_usize;
let a = argvals[0];
let b = argvals[n - 1];
let penalty_matrix = monomial_penalty_analytic(nbasis, lfd_order, a, b);
Ok(BasisSystem {
eval_matrix,
penalty_matrix,
nbasis,
n_eval: n,
lfd_order,
})
}
fn falling_factorial(e: f64, d: usize) -> f64 {
if d == 0 {
return 1.0;
}
let mut acc = 1.0_f64;
for k in 0..d {
acc *= e - k as f64;
}
acc
}
fn gram_entry(ei: f64, ej: f64, d: usize, a: f64, b: f64) -> f64 {
let ci = falling_factorial(ei, d);
let cj = falling_factorial(ej, d);
if ci.abs() < 1e-15 || cj.abs() < 1e-15 {
return 0.0;
}
let p = ei + ej - 2.0 * d as f64 + 1.0;
if p.abs() < 1e-15 {
if a <= 0.0 {
debug_assert!(
false,
"gram_entry: improper integral t^(-1) encountered (a={a}, b={b}); \
penalty result would be wrong if lfd_order < 2"
);
return 0.0;
}
ci * cj * (b.ln() - a.ln())
} else {
ci * cj * (b.powf(p) - a.powf(p)) / p
}
}
fn monomial_penalty_analytic(nbasis: usize, lfd_order: usize, a: f64, b: f64) -> Vec<f64> {
let mut penalty = vec![0.0_f64; nbasis * nbasis];
for j in 0..nbasis {
for k in j..nbasis {
let val = gram_entry(j as f64, k as f64, lfd_order, a, b);
penalty[j + k * nbasis] = val;
penalty[k + j * nbasis] = val; }
}
penalty
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn monomial_invalid_argvals_too_short() {
let result = monomial_basis(&[0.5], 2);
assert!(matches!(result, Err(FdarError::InvalidDimension { .. })));
}
#[test]
fn monomial_invalid_nbasis_zero() {
let result = monomial_basis(&[0.0, 1.0], 0);
assert!(matches!(result, Err(FdarError::InvalidParameter { .. })));
}
#[test]
fn monomial_eval_matrix_closed_form() {
let t = vec![0.0, 1.0, 2.0];
let bs = monomial_basis(&t, 3).unwrap();
let n = bs.n_eval; assert!((bs.eval_matrix[0] - 1.0).abs() < 1e-12, "B₀(0)");
assert!((bs.eval_matrix[1] - 1.0).abs() < 1e-12, "B₀(1)");
assert!((bs.eval_matrix[2] - 1.0).abs() < 1e-12, "B₀(2)");
assert!((bs.eval_matrix[n] - 0.0).abs() < 1e-12, "B₁(0)");
assert!((bs.eval_matrix[n + 1] - 1.0).abs() < 1e-12, "B₁(1)");
assert!((bs.eval_matrix[n + 2] - 2.0).abs() < 1e-12, "B₁(2)");
assert!((bs.eval_matrix[2 * n] - 0.0).abs() < 1e-12, "B₂(0)");
assert!((bs.eval_matrix[2 * n + 1] - 1.0).abs() < 1e-12, "B₂(1)");
assert!((bs.eval_matrix[2 * n + 2] - 4.0).abs() < 1e-12, "B₂(2)");
}
#[test]
fn monomial_eval_matrix_shape() {
let t = vec![0.0, 0.5, 1.0];
let bs = monomial_basis(&t, 4).unwrap();
assert_eq!(bs.eval_matrix.len(), 3 * 4);
assert_eq!(bs.penalty_matrix.len(), 4 * 4);
assert_eq!(bs.nbasis, 4);
assert_eq!(bs.n_eval, 3);
assert_eq!(bs.lfd_order, 2);
}
#[test]
fn monomial_penalty_p22_standard_domain() {
let t = vec![0.0, 0.5, 1.0];
let bs = monomial_basis(&t, 3).unwrap();
let k = 3;
let p22 = bs.penalty_matrix[2 + 2 * k];
assert!((p22 - 4.0).abs() < 1e-9, "P[2,2] = {p22}, expected 4.0");
}
#[test]
fn monomial_penalty_symmetry() {
let t: Vec<f64> = (0..10).map(|i| i as f64 / 9.0).collect();
let bs = monomial_basis(&t, 5).unwrap();
let k = 5;
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-12,
"P[{j},{l}]={pjl} != P[{l},{j}]={plj}"
);
}
}
}
#[test]
fn monomial_penalty_diagonal_psd() {
let t: Vec<f64> = (0..10).map(|i| i as f64 / 9.0).collect();
let bs = monomial_basis(&t, 6).unwrap();
let k = 6;
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 monomial_penalty_low_exponents_zero() {
let t = vec![0.0, 0.5, 1.0];
let bs = monomial_basis(&t, 3).unwrap();
let k = bs.nbasis; assert!(bs.penalty_matrix[0].abs() < 1e-12, "P[0,0] should be 0");
assert!(bs.penalty_matrix[1 + k].abs() < 1e-12, "P[1,1] should be 0");
}
#[test]
fn basis_system_derives() {
let t = vec![0.0, 1.0];
let bs = monomial_basis(&t, 2).unwrap();
let bs2 = bs.clone();
assert_eq!(bs, bs2);
let _ = format!("{bs:?}"); }
}