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 exponential_basis(argvals: &[f64], rates: &[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 = rates.len();
if nbasis < 1 {
return Err(FdarError::InvalidParameter {
parameter: "rates",
message: "must be non-empty".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] = (rates[j] * t).exp();
}
}
let lfd_order = 2_usize;
let penalty_matrix = exponential_penalty_numeric(argvals, rates, nbasis, lfd_order);
Ok(BasisSystem {
eval_matrix,
penalty_matrix,
nbasis,
n_eval: n,
lfd_order,
})
}
fn exponential_penalty_numeric(
argvals: &[f64],
rates: &[f64],
nbasis: usize,
lfd_order: usize,
) -> Vec<f64> {
if argvals.len() < 2 {
return vec![0.0; nbasis * nbasis];
}
let t_min = argvals[0];
let t_max = argvals[argvals.len() - 1];
let n_sub = 10;
let n_quad = (argvals.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] = (rates[j] * t).exp();
}
}
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 exp_invalid_argvals_too_short() {
let result = exponential_basis(&[0.5], &[0.0, -1.0]);
assert!(matches!(result, Err(FdarError::InvalidDimension { .. })));
}
#[test]
fn exp_invalid_empty_rates() {
let result = exponential_basis(&[0.0, 1.0], &[]);
assert!(matches!(result, Err(FdarError::InvalidParameter { .. })));
}
#[test]
fn exp_eval_at_zero_is_one() {
let t = vec![0.0, 1.0];
let rates = [0.0, -1.0, -5.0, 2.0];
let bs = exponential_basis(&t, &rates).unwrap();
let n = bs.n_eval;
for j in 0..bs.nbasis {
let val = bs.eval_matrix[j * n]; assert!((val - 1.0).abs() < 1e-12, "B_{j}(0)={val}, expected 1.0");
}
}
#[test]
fn exp_rate_zero_is_constant_one() {
let t: Vec<f64> = (0..5).map(|i| i as f64 * 0.5).collect();
let bs = exponential_basis(&t, &[0.0]).unwrap();
let n = bs.n_eval;
for ti in 0..n {
let val = bs.eval_matrix[ti]; assert!((val - 1.0).abs() < 1e-12, "B_0(t_{ti})={val}, expected 1.0");
}
}
#[test]
fn exp_eval_closed_form() {
let t = vec![0.0, 1.0];
let bs = exponential_basis(&t, &[0.0, -1.0]).unwrap();
let n = bs.n_eval; assert!((bs.eval_matrix[0] - 1.0).abs() < 1e-12, "B_0(0)");
assert!((bs.eval_matrix[1] - 1.0).abs() < 1e-12, "B_0(1)");
assert!((bs.eval_matrix[n] - 1.0).abs() < 1e-12, "B_1(0)");
assert!(
(bs.eval_matrix[n + 1] - (-1.0_f64).exp()).abs() < 1e-12,
"B_1(1)"
);
}
#[test]
fn exp_shape_invariants() {
let t: Vec<f64> = (0..5).map(|i| i as f64 / 4.0).collect();
let rates = [0.0, -1.0, -2.0];
let bs = exponential_basis(&t, &rates).unwrap();
assert_eq!(bs.eval_matrix.len(), 5 * 3);
assert_eq!(bs.penalty_matrix.len(), 3 * 3);
assert_eq!(bs.nbasis, 3);
assert_eq!(bs.n_eval, 5);
assert_eq!(bs.lfd_order, 2);
}
#[test]
fn exp_penalty_symmetric() {
let t: Vec<f64> = (0..10).map(|i| i as f64 / 9.0).collect();
let rates = [0.0, -1.0, -3.0];
let bs = exponential_basis(&t, &rates).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 exp_penalty_diagonal_nonneg() {
let t: Vec<f64> = (0..10).map(|i| i as f64 / 9.0).collect();
let bs = exponential_basis(&t, &[0.0, -1.0, 1.0]).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 exp_basis_system_derives() {
let t = vec![0.0, 1.0];
let bs = exponential_basis(&t, &[0.0]).unwrap();
let bs2 = bs.clone();
assert_eq!(bs, bs2);
let _ = format!("{bs:?}");
}
}