1use crate::basis::{bspline_basis, fourier_basis_with_period};
13use crate::helpers::simpsons_weights;
14use crate::matrix::FdMatrix;
15use nalgebra::DMatrix;
16use std::f64::consts::PI;
17
18#[non_exhaustive]
22#[derive(Debug, Clone, PartialEq)]
23pub enum BasisType {
24 Bspline { order: usize },
26 Fourier { period: f64 },
28}
29
30#[derive(Debug, Clone, PartialEq)]
32pub struct FdPar {
33 pub basis_type: BasisType,
35 pub nbasis: usize,
37 pub lambda: f64,
39 pub lfd_order: usize,
41 pub penalty_matrix: Vec<f64>,
43}
44
45#[derive(Debug, Clone, PartialEq)]
61#[non_exhaustive]
62#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
63pub struct SmoothPositiveResult {
64 pub fitted: FdMatrix,
66 pub log_coefficients: FdMatrix,
68 pub edf: f64,
70 pub gcv: f64,
72}
73
74#[derive(Debug, Clone, PartialEq)]
76#[non_exhaustive]
77pub struct SmoothBasisResult {
78 pub coefficients: FdMatrix,
80 pub fitted: FdMatrix,
82 pub edf: f64,
84 pub gcv: f64,
86 pub aic: f64,
88 pub bic: f64,
90 pub penalty_matrix: Vec<f64>,
92 pub nbasis: usize,
94}
95
96pub fn bspline_penalty_matrix(
113 argvals: &[f64],
114 nbasis: usize,
115 order: usize,
116 lfd_order: usize,
117) -> Vec<f64> {
118 if nbasis < 2 || order < 1 || lfd_order >= order || argvals.len() < 2 {
119 return vec![0.0; nbasis * nbasis];
120 }
121
122 let nknots = nbasis.saturating_sub(order).max(2);
123
124 let n_sub = 10;
126 let t_min = argvals[0];
127 let t_max = argvals[argvals.len() - 1];
128 let n_quad = (argvals.len() - 1) * n_sub + 1;
129 let quad_t: Vec<f64> = (0..n_quad)
130 .map(|i| t_min + (t_max - t_min) * i as f64 / (n_quad - 1) as f64)
131 .collect();
132
133 let basis_fine = bspline_basis(&quad_t, nknots, order);
135 let actual_nbasis = basis_fine.len() / n_quad;
136
137 let h = (t_max - t_min) / (n_quad - 1) as f64;
139 let deriv_basis = differentiate_basis_columns(&basis_fine, n_quad, actual_nbasis, h, lfd_order);
140
141 let weights = simpsons_weights(&quad_t);
143
144 integrate_symmetric_penalty(&deriv_basis, &weights, actual_nbasis, n_quad)
146}
147
148pub fn fourier_penalty_matrix(nbasis: usize, period: f64, lfd_order: usize) -> Vec<f64> {
160 let k = nbasis;
161 let mut penalty = vec![0.0; k * k];
162
163 let mut freq = 1;
169 let mut idx = 1;
170 while idx < k {
171 let omega = 2.0 * PI * f64::from(freq) / period;
172 let eigenval = omega.powi(2 * lfd_order as i32);
173
174 if idx < k {
176 penalty[idx + idx * k] = eigenval;
177 idx += 1;
178 }
179 if idx < k {
181 penalty[idx + idx * k] = eigenval;
182 idx += 1;
183 }
184 freq += 1;
185 }
186
187 penalty
188}
189
190pub fn smooth_basis(
205 data: &FdMatrix,
206 argvals: &[f64],
207 fdpar: &FdPar,
208) -> Result<SmoothBasisResult, crate::FdarError> {
209 let (n, m) = data.shape();
210 if n == 0 || m == 0 || argvals.len() != m || fdpar.nbasis < 2 {
211 return Err(crate::FdarError::InvalidDimension {
212 parameter: "data/argvals/fdpar",
213 expected: "n > 0, m > 0, argvals.len() == m, nbasis >= 2".to_string(),
214 actual: format!(
215 "n={}, m={}, argvals.len()={}, nbasis={}",
216 n,
217 m,
218 argvals.len(),
219 fdpar.nbasis
220 ),
221 });
222 }
223
224 let (basis_flat, actual_nbasis) = evaluate_basis(argvals, &fdpar.basis_type, fdpar.nbasis);
226 let k = actual_nbasis;
227
228 let b_mat = DMatrix::from_column_slice(m, k, &basis_flat);
229 let r_mat = DMatrix::from_column_slice(k, k, &fdpar.penalty_matrix);
230
231 let btb = b_mat.transpose() * &b_mat;
233 let ridge_eps = 1e-10;
234 let system: DMatrix<f64> =
235 &btb + fdpar.lambda * &r_mat + ridge_eps * DMatrix::<f64>::identity(k, k);
236
237 let system_inv =
239 invert_penalized_system(&system, k).ok_or_else(|| crate::FdarError::ComputationFailed {
240 operation: "matrix inversion",
241 detail: "failed to invert penalized system (Φ'Φ + λR); try increasing lambda or reducing the number of basis functions".to_string(),
242 })?;
243
244 let h_mat = &b_mat * &system_inv * b_mat.transpose();
246 let edf: f64 = (0..m).map(|i| h_mat[(i, i)]).sum();
247
248 let proj = &system_inv * b_mat.transpose();
250 let (all_coefs, all_fitted, total_rss) = project_all_curves(data, &b_mat, &proj, n, m, k);
251
252 let total_points = (n * m) as f64;
253 let gcv = compute_gcv(total_rss, total_points, edf, m);
254 let mse = total_rss / total_points;
255 let total_edf = n as f64 * edf;
257 let aic = total_points * mse.max(1e-300).ln() + 2.0 * total_edf;
258 let bic = total_points * mse.max(1e-300).ln() + total_points.ln() * total_edf;
259
260 Ok(SmoothBasisResult {
261 coefficients: all_coefs,
262 fitted: all_fitted,
263 edf,
264 gcv,
265 aic,
266 bic,
267 penalty_matrix: fdpar.penalty_matrix.clone(),
268 nbasis: k,
269 })
270}
271
272pub fn smooth_basis_gcv(
285 data: &FdMatrix,
286 argvals: &[f64],
287 basis_type: &BasisType,
288 nbasis: usize,
289 lfd_order: usize,
290 log_lambda_range: (f64, f64),
291 n_grid: usize,
292) -> Option<SmoothBasisResult> {
293 let m = argvals.len();
294 if m == 0 || nbasis < 2 || n_grid < 2 {
295 return None;
296 }
297
298 let penalty = match basis_type {
300 BasisType::Bspline { order } => bspline_penalty_matrix(argvals, nbasis, *order, lfd_order),
301 BasisType::Fourier { period } => fourier_penalty_matrix(nbasis, *period, lfd_order),
302 };
303
304 let (lo, hi) = log_lambda_range;
305 let mut best_gcv = f64::INFINITY;
306 let mut best_result: Option<SmoothBasisResult> = None;
307
308 for i in 0..n_grid {
309 let log_lam = lo + (hi - lo) * i as f64 / (n_grid - 1) as f64;
310 let lam = 10.0_f64.powf(log_lam);
311
312 let fdpar = FdPar {
313 basis_type: basis_type.clone(),
314 nbasis,
315 lambda: lam,
316 lfd_order,
317 penalty_matrix: penalty.clone(),
318 };
319
320 if let Ok(result) = smooth_basis(data, argvals, &fdpar) {
321 if result.gcv < best_gcv {
322 best_gcv = result.gcv;
323 best_result = Some(result);
324 }
325 }
326 }
327
328 best_result
329}
330
331pub fn smooth_basis_aic(
368 data: &FdMatrix,
369 argvals: &[f64],
370 basis_type: &BasisType,
371 nbasis: usize,
372 lfd_order: usize,
373 log_lambda_range: (f64, f64),
374 n_grid: usize,
375) -> Option<SmoothBasisResult> {
376 let m = argvals.len();
377 if m == 0 || nbasis < 2 || n_grid < 2 {
378 return None;
379 }
380
381 let penalty = match basis_type {
383 BasisType::Bspline { order } => bspline_penalty_matrix(argvals, nbasis, *order, lfd_order),
384 BasisType::Fourier { period } => fourier_penalty_matrix(nbasis, *period, lfd_order),
385 };
386
387 let (lo, hi) = log_lambda_range;
388 let mut best_aic = f64::INFINITY;
389 let mut best_result: Option<SmoothBasisResult> = None;
390
391 for i in 0..n_grid {
392 let log_lam = lo + (hi - lo) * i as f64 / (n_grid - 1) as f64;
393 let lam = 10.0_f64.powf(log_lam);
394
395 let fdpar = FdPar {
396 basis_type: basis_type.clone(),
397 nbasis,
398 lambda: lam,
399 lfd_order,
400 penalty_matrix: penalty.clone(),
401 };
402
403 if let Ok(result) = smooth_basis(data, argvals, &fdpar) {
404 if result.aic < best_aic {
407 best_aic = result.aic;
408 best_result = Some(result);
409 }
410 }
411 }
412
413 best_result
414}
415
416#[must_use = "expensive computation whose result should not be discarded"]
463pub fn smooth_positive(
464 data: &FdMatrix,
465 argvals: &[f64],
466 fdpar: &FdPar,
467) -> Result<SmoothPositiveResult, crate::FdarError> {
468 let (n, m) = data.shape();
469 if n == 0 || m == 0 || argvals.len() != m {
470 return Err(crate::FdarError::InvalidDimension {
471 parameter: "data/argvals",
472 expected: "n > 0, m > 0, argvals.len() == m".to_string(),
473 actual: format!("n={}, m={}, argvals.len()={}", n, m, argvals.len()),
474 });
475 }
476
477 for i in 0..n {
479 for j in 0..m {
480 if data[(i, j)] <= 0.0 {
481 return Err(crate::FdarError::InvalidParameter {
482 parameter: "data",
483 message: format!(
484 "smooth_positive requires strictly positive data (log-domain smoother); \
485 found value {} <= 0 at observation {}, evaluation point {}",
486 data[(i, j)],
487 i,
488 j
489 ),
490 });
491 }
492 }
493 }
494
495 let mut log_data = FdMatrix::zeros(n, m);
497 for i in 0..n {
498 for j in 0..m {
499 log_data[(i, j)] = data[(i, j)].ln();
500 }
501 }
502
503 let inner = smooth_basis(&log_data, argvals, fdpar)?;
505
506 let mut fitted = FdMatrix::zeros(n, m);
508 for i in 0..n {
509 for j in 0..m {
510 fitted[(i, j)] = inner.fitted[(i, j)].exp();
511 }
512 }
513
514 Ok(SmoothPositiveResult {
515 fitted,
516 log_coefficients: inner.coefficients,
517 edf: inner.edf,
518 gcv: inner.gcv,
519 })
520}
521
522#[derive(Debug, Clone, PartialEq)]
531#[non_exhaustive]
532#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
533pub struct SmoothMonotoneResult {
534 pub fitted: Vec<f64>,
536 pub beta0: f64,
538 pub beta1: f64,
540 pub w_coefficients: Vec<f64>,
542 pub iterations: usize,
544 pub converged: bool,
546}
547
548#[must_use = "expensive nonlinear least-squares computation whose result should not be discarded"]
593pub fn smooth_monotone(
594 data: &[f64],
595 argvals: &[f64],
596 nbasis: usize,
597 order: usize,
598 lambda: f64,
599 max_iter: usize,
600) -> Result<SmoothMonotoneResult, crate::FdarError> {
601 let m = data.len();
603 if m < 3 || argvals.len() != m {
604 return Err(crate::FdarError::InvalidDimension {
605 parameter: "data/argvals",
606 expected: "data.len() >= 3, argvals.len() == data.len()".to_string(),
607 actual: format!("data.len()={}, argvals.len()={}", m, argvals.len()),
608 });
609 }
610 if nbasis < 2 {
611 return Err(crate::FdarError::InvalidParameter {
612 parameter: "nbasis",
613 message: format!("smooth_monotone requires nbasis >= 2, got {}", nbasis),
614 });
615 }
616 if order < 1 {
617 return Err(crate::FdarError::InvalidParameter {
618 parameter: "order",
619 message: format!("smooth_monotone requires order >= 1, got {}", order),
620 });
621 }
622 if lambda < 0.0 {
623 return Err(crate::FdarError::InvalidParameter {
624 parameter: "lambda",
625 message: format!("smooth_monotone requires lambda >= 0, got {}", lambda),
626 });
627 }
628 if max_iter < 1 {
629 return Err(crate::FdarError::InvalidParameter {
630 parameter: "max_iter",
631 message: format!("smooth_monotone requires max_iter >= 1, got {}", max_iter),
632 });
633 }
634
635 let nknots = nbasis.saturating_sub(order).max(2);
639 let basis_flat = crate::basis::bspline_basis(argvals, nknots, order);
641 let actual_k = basis_flat.len() / m;
642
643 let penalty_col = bspline_penalty_matrix(argvals, nbasis, order, 2);
647 let actual_k_pm = (penalty_col.len() as f64).sqrt() as usize;
649 let k = actual_k.min(actual_k_pm);
651
652 let psi_at = |i: usize, j: usize| -> f64 {
654 if j < actual_k {
655 basis_flat[i + j * m]
656 } else {
657 0.0
658 }
659 };
660
661 let t_range = (argvals[m - 1] - argvals[0]).max(1e-12);
663 let beta0_init = data[0];
664 let beta1_init = (data[m - 1] - data[0]) / t_range;
666 let beta1_init = if beta1_init.abs() < 1e-12 {
668 1e-6
669 } else {
670 beta1_init
671 };
672 let mut beta0 = beta0_init;
673 let mut beta1 = beta1_init;
674 let mut alpha = vec![0.0_f64; k]; let big_p = 2 + k; let build_integrals = |alpha: &[f64]| -> (Vec<f64>, Vec<f64>) {
683 let mut exp_w = vec![0.0_f64; m];
684 for i in 0..m {
685 let w_i: f64 = (0..k).map(|j| alpha[j] * psi_at(i, j)).sum();
686 let w_clamped = w_i.clamp(-30.0, 30.0);
688 exp_w[i] = w_clamped.exp();
689 }
690 let mut w_int = vec![0.0_f64; m];
692 for i in 1..m {
693 let dt = argvals[i] - argvals[i - 1];
694 w_int[i] = w_int[i - 1] + 0.5 * (exp_w[i - 1] + exp_w[i]) * dt;
695 }
696 let mut iexp_psi = vec![0.0_f64; m * k];
698 for i in 1..m {
699 let dt = argvals[i] - argvals[i - 1];
700 for j in 0..k {
701 let integrand_prev = exp_w[i - 1] * psi_at(i - 1, j);
702 let integrand_cur = exp_w[i] * psi_at(i, j);
703 iexp_psi[i * k + j] =
704 iexp_psi[(i - 1) * k + j] + 0.5 * (integrand_prev + integrand_cur) * dt;
705 }
706 }
707 (w_int, iexp_psi)
708 };
709
710 let mut iterations = 0_usize;
712 let mut converged = false;
713
714 for iter in 0..max_iter {
715 let (w_int, iexp_psi) = build_integrals(&alpha);
716
717 let mut resid = vec![0.0_f64; m];
719 for i in 0..m {
720 let f_i = beta0 + beta1 * w_int[i];
721 resid[i] = data[i] - f_i;
722 }
723
724 let mut a_mat = vec![0.0_f64; big_p * big_p];
726 let mut g_vec = vec![0.0_f64; big_p];
727
728 for i in 0..m {
729 let j0 = 1.0_f64;
734 let j1 = w_int[i];
735
736 a_mat[0] += j0 * j0;
741 a_mat[1] += j0 * j1;
743 a_mat[big_p] += j1 * j0; a_mat[big_p + 1] += j1 * j1;
746
747 g_vec[0] += j0 * resid[i];
749 g_vec[1] += j1 * resid[i];
750
751 for aj in 0..k {
752 let jc = beta1 * iexp_psi[i * k + aj];
753 a_mat[2 + aj] += j0 * jc;
755 a_mat[(2 + aj) * big_p] += jc * j0;
756 a_mat[big_p + (2 + aj)] += j1 * jc;
758 a_mat[(2 + aj) * big_p + 1] += jc * j1;
759 for bj in 0..k {
761 let jd = beta1 * iexp_psi[i * k + bj];
762 a_mat[(2 + aj) * big_p + (2 + bj)] += jc * jd;
763 }
764 g_vec[2 + aj] += jc * resid[i];
766 }
767 }
768
769 for d in 0..big_p {
771 a_mat[d * big_p + d] += 1e-6 * (1.0 + a_mat[d * big_p + d].abs());
772 }
773
774 for a in 0..k {
778 for b in 0..k {
779 let r_ab = if a < actual_k_pm && b < actual_k_pm {
780 penalty_col[a + b * actual_k_pm]
781 } else {
782 0.0
783 };
784 a_mat[(2 + a) * big_p + (2 + b)] += lambda * r_ab;
785 }
786 }
787
788 let delta = crate::linalg::cholesky_solve(&a_mat, &g_vec, big_p)?;
790
791 beta0 += delta[0];
793 beta1 += delta[1];
794 for j in 0..k {
795 alpha[j] += delta[2 + j];
796 }
797
798 iterations = iter + 1;
799
800 let delta_norm: f64 = delta.iter().map(|d| d * d).sum::<f64>().sqrt();
802 if delta_norm < 1e-8 {
803 converged = true;
804 break;
805 }
806 }
807
808 let (w_int_final, _) = build_integrals(&alpha);
810 let fitted: Vec<f64> = (0..m).map(|i| beta0 + beta1 * w_int_final[i]).collect();
811
812 Ok(SmoothMonotoneResult {
813 fitted,
814 beta0,
815 beta1,
816 w_coefficients: alpha,
817 iterations,
818 converged,
819 })
820}
821
822#[non_exhaustive]
840#[derive(Debug, Clone, PartialEq)]
841pub struct SmoothBasisGcvConfig {
842 pub basis_type: BasisType,
844 pub nbasis: usize,
846 pub lfd_order: usize,
848 pub log_lambda_range: (f64, f64),
850 pub n_grid: usize,
852}
853
854impl Default for SmoothBasisGcvConfig {
855 fn default() -> Self {
856 Self {
857 basis_type: BasisType::Bspline { order: 4 },
858 nbasis: 15,
859 lfd_order: 2,
860 log_lambda_range: (-10.0, 2.0),
861 n_grid: 50,
862 }
863 }
864}
865
866#[must_use = "expensive computation whose result should not be discarded"]
881pub fn smooth_basis_gcv_with_config(
882 data: &FdMatrix,
883 argvals: &[f64],
884 config: &SmoothBasisGcvConfig,
885) -> Result<SmoothBasisResult, crate::FdarError> {
886 smooth_basis_gcv(
887 data,
888 argvals,
889 &config.basis_type,
890 config.nbasis,
891 config.lfd_order,
892 config.log_lambda_range,
893 config.n_grid,
894 )
895 .ok_or_else(|| crate::FdarError::ComputationFailed {
896 operation: "smooth_basis_gcv_with_config",
897 detail: "no valid smoothing result found in GCV lambda search".to_string(),
898 })
899}
900
901#[non_exhaustive]
917#[derive(Debug, Clone, PartialEq)]
918pub struct BasisNbasisCvConfig {
919 pub basis_type: BasisType,
921 pub nbasis_range: (usize, usize),
923 pub lambda: f64,
925 pub lfd_order: usize,
927 pub n_folds: usize,
929 pub criterion: BasisCriterion,
931}
932
933impl Default for BasisNbasisCvConfig {
934 fn default() -> Self {
935 Self {
936 basis_type: BasisType::Bspline { order: 4 },
937 nbasis_range: (5, 30),
938 lambda: 1e-4,
939 lfd_order: 2,
940 n_folds: 5,
941 criterion: BasisCriterion::Gcv,
942 }
943 }
944}
945
946#[must_use = "expensive computation whose result should not be discarded"]
964pub fn basis_nbasis_cv_with_config(
965 data: &FdMatrix,
966 argvals: &[f64],
967 config: &BasisNbasisCvConfig,
968) -> Result<BasisNbasisCvResult, crate::FdarError> {
969 let nbasis_range: Vec<usize> = (config.nbasis_range.0..=config.nbasis_range.1).collect();
970 basis_nbasis_cv(
971 data,
972 argvals,
973 &nbasis_range,
974 &config.basis_type,
975 config.criterion,
976 config.n_folds,
977 config.lambda,
978 )
979 .ok_or_else(|| crate::FdarError::ComputationFailed {
980 operation: "basis_nbasis_cv_with_config",
981 detail: "no valid result found in nbasis CV search".to_string(),
982 })
983}
984
985pub(crate) fn differentiate_basis_columns(
989 basis: &[f64],
990 n_quad: usize,
991 nbasis: usize,
992 h: f64,
993 lfd_order: usize,
994) -> Vec<f64> {
995 let mut deriv = basis.to_vec();
996 for _ in 0..lfd_order {
997 let mut new_deriv = vec![0.0; n_quad * nbasis];
998 for j in 0..nbasis {
999 let col: Vec<f64> = (0..n_quad).map(|i| deriv[i + j * n_quad]).collect();
1000 let grad = crate::helpers::gradient_uniform(&col, h);
1001 for i in 0..n_quad {
1002 new_deriv[i + j * n_quad] = grad[i];
1003 }
1004 }
1005 deriv = new_deriv;
1006 }
1007 deriv
1008}
1009
1010pub(crate) fn integrate_symmetric_penalty(
1012 deriv_basis: &[f64],
1013 weights: &[f64],
1014 k: usize,
1015 n_quad: usize,
1016) -> Vec<f64> {
1017 let mut penalty = vec![0.0; k * k];
1018 for j in 0..k {
1019 for l in j..k {
1020 let mut val = 0.0;
1021 for i in 0..n_quad {
1022 val += deriv_basis[i + j * n_quad] * deriv_basis[i + l * n_quad] * weights[i];
1023 }
1024 penalty[j + l * k] = val;
1025 penalty[l + j * k] = val;
1026 }
1027 }
1028 penalty
1029}
1030
1031fn evaluate_basis(argvals: &[f64], basis_type: &BasisType, nbasis: usize) -> (Vec<f64>, usize) {
1033 let m = argvals.len();
1034 match basis_type {
1035 BasisType::Bspline { order } => {
1036 let nknots = nbasis.saturating_sub(*order).max(2);
1037 let basis = bspline_basis(argvals, nknots, *order);
1038 let actual = basis.len() / m;
1039 (basis, actual)
1040 }
1041 BasisType::Fourier { period } => {
1042 let basis = fourier_basis_with_period(argvals, nbasis, *period);
1043 (basis, nbasis)
1044 }
1045 }
1046}
1047
1048fn invert_penalized_system(system: &DMatrix<f64>, k: usize) -> Option<DMatrix<f64>> {
1050 if let Some(chol) = system.clone().cholesky() {
1051 return Some(chol.inverse());
1052 }
1053 let svd = nalgebra::SVD::new(system.clone(), true, true);
1055 let u = svd.u.as_ref()?;
1056 let v_t = svd.v_t.as_ref()?;
1057 let max_sv: f64 = svd.singular_values.iter().copied().fold(0.0_f64, f64::max);
1058 let eps = 1e-10 * max_sv;
1059 let mut inv = DMatrix::<f64>::zeros(k, k);
1060 for ii in 0..k {
1061 for jj in 0..k {
1062 let mut sum = 0.0;
1063 for s in 0..k.min(svd.singular_values.len()) {
1064 if svd.singular_values[s] > eps {
1065 sum += v_t[(s, ii)] / svd.singular_values[s] * u[(jj, s)];
1066 }
1067 }
1068 inv[(ii, jj)] = sum;
1069 }
1070 }
1071 Some(inv)
1072}
1073
1074fn project_all_curves(
1076 data: &FdMatrix,
1077 b_mat: &DMatrix<f64>,
1078 proj: &DMatrix<f64>,
1079 n: usize,
1080 m: usize,
1081 k: usize,
1082) -> (FdMatrix, FdMatrix, f64) {
1083 let mut all_coefs = FdMatrix::zeros(n, k);
1084 let mut all_fitted = FdMatrix::zeros(n, m);
1085 let mut total_rss = 0.0;
1086
1087 for i in 0..n {
1088 let curve: Vec<f64> = (0..m).map(|j| data[(i, j)]).collect();
1089 let y_vec = nalgebra::DVector::from_vec(curve.clone());
1090 let coefs = proj * &y_vec;
1091
1092 for j in 0..k {
1093 all_coefs[(i, j)] = coefs[j];
1094 }
1095 let fitted = b_mat * &coefs;
1096 for j in 0..m {
1097 all_fitted[(i, j)] = fitted[j];
1098 let resid = curve[j] - fitted[j];
1099 total_rss += resid * resid;
1100 }
1101 }
1102
1103 (all_coefs, all_fitted, total_rss)
1104}
1105
1106fn compute_gcv(rss: f64, n_points: f64, edf: f64, m: usize) -> f64 {
1108 let gcv_denom = 1.0 - edf / m as f64;
1109 if gcv_denom.abs() > 1e-10 {
1110 (rss / n_points) / (gcv_denom * gcv_denom)
1111 } else {
1112 f64::INFINITY
1113 }
1114}
1115
1116#[non_exhaustive]
1120#[derive(Debug, Clone, Copy, PartialEq)]
1121pub enum BasisCriterion {
1122 Gcv,
1124 Cv,
1126 Aic,
1128 Bic,
1130}
1131
1132#[derive(Debug, Clone, PartialEq)]
1134#[non_exhaustive]
1135pub struct BasisNbasisCvResult {
1136 pub optimal_nbasis: usize,
1138 pub scores: Vec<f64>,
1140 pub nbasis_range: Vec<usize>,
1142 pub criterion: BasisCriterion,
1144}
1145
1146fn evaluate_nbasis_info_criterion(
1148 data: &FdMatrix,
1149 argvals: &[f64],
1150 nbasis_range: &[usize],
1151 basis_type: &BasisType,
1152 criterion: BasisCriterion,
1153 lambda: f64,
1154) -> Vec<f64> {
1155 let mut scores = Vec::with_capacity(nbasis_range.len());
1156 for &nb in nbasis_range {
1157 if nb < 2 {
1158 scores.push(f64::INFINITY);
1159 continue;
1160 }
1161 let penalty = match basis_type {
1162 BasisType::Bspline { order } => bspline_penalty_matrix(argvals, nb, *order, 2),
1163 BasisType::Fourier { period } => fourier_penalty_matrix(nb, *period, 2),
1164 };
1165 let fdpar = FdPar {
1166 basis_type: basis_type.clone(),
1167 nbasis: nb,
1168 lambda,
1169 lfd_order: 2,
1170 penalty_matrix: penalty,
1171 };
1172 match smooth_basis(data, argvals, &fdpar) {
1173 Ok(result) => {
1174 let score = match criterion {
1175 BasisCriterion::Gcv => result.gcv,
1176 BasisCriterion::Aic => result.aic,
1177 BasisCriterion::Bic => result.bic,
1178 BasisCriterion::Cv => unreachable!(),
1179 };
1180 scores.push(score);
1181 }
1182 Err(_) => scores.push(f64::INFINITY),
1183 }
1184 }
1185 scores
1186}
1187
1188fn evaluate_nbasis_cv(
1190 data: &FdMatrix,
1191 argvals: &[f64],
1192 nbasis_range: &[usize],
1193 basis_type: &BasisType,
1194 lambda: f64,
1195 n_folds: usize,
1196) -> Vec<f64> {
1197 let (n, m) = data.shape();
1198 let n_folds = n_folds.max(2).min(m);
1206 let point_folds = crate::cv::create_folds(m, n_folds, 42);
1207 let mut scores = Vec::with_capacity(nbasis_range.len());
1208
1209 for &nb in nbasis_range {
1210 if nb < 2 {
1211 scores.push(f64::INFINITY);
1212 continue;
1213 }
1214 let penalty = match basis_type {
1215 BasisType::Bspline { order } => bspline_penalty_matrix(argvals, nb, *order, 2),
1216 BasisType::Fourier { period } => fourier_penalty_matrix(nb, *period, 2),
1217 };
1218 let (basis_flat, actual_k) = evaluate_basis(argvals, basis_type, nb);
1219 let b_full = DMatrix::from_column_slice(m, actual_k, &basis_flat);
1220 let r_mat = DMatrix::from_column_slice(actual_k, actual_k, &penalty);
1221
1222 let mut total_se = 0.0;
1223 let mut count = 0usize;
1224
1225 for fold in 0..n_folds {
1226 let (train_pts, test_pts) = crate::cv::fold_indices(&point_folds, fold);
1227 if train_pts.is_empty() || test_pts.is_empty() {
1228 continue;
1229 }
1230 let b_train = b_full.select_rows(train_pts.iter());
1233 let b_test = b_full.select_rows(test_pts.iter());
1234 let btb = b_train.transpose() * &b_train;
1235 let ridge_eps = 1e-10;
1236 let system: DMatrix<f64> =
1237 &btb + lambda * &r_mat + ridge_eps * DMatrix::<f64>::identity(actual_k, actual_k);
1238 let Some(system_inv) = invert_penalized_system(&system, actual_k) else {
1239 continue;
1240 };
1241 let proj = &system_inv * b_train.transpose(); for i in 0..n {
1244 let y_train = nalgebra::DVector::from_iterator(
1245 train_pts.len(),
1246 train_pts.iter().map(|&j| data[(i, j)]),
1247 );
1248 let coefs = &proj * &y_train;
1249 let pred = &b_test * &coefs; for (t_idx, &j) in test_pts.iter().enumerate() {
1251 let err = data[(i, j)] - pred[t_idx];
1252 total_se += err * err;
1253 count += 1;
1254 }
1255 }
1256 }
1257
1258 if count > 0 {
1259 scores.push(total_se / count as f64);
1260 } else {
1261 scores.push(f64::INFINITY);
1262 }
1263 }
1264 scores
1265}
1266
1267pub fn basis_nbasis_cv(
1270 data: &FdMatrix,
1271 argvals: &[f64],
1272 nbasis_range: &[usize],
1273 basis_type: &BasisType,
1274 criterion: BasisCriterion,
1275 n_folds: usize,
1276 lambda: f64,
1277) -> Option<BasisNbasisCvResult> {
1278 let (n, m) = data.shape();
1279 if n == 0 || m == 0 || argvals.len() != m || nbasis_range.is_empty() {
1280 return None;
1281 }
1282
1283 let scores = match criterion {
1284 BasisCriterion::Gcv | BasisCriterion::Aic | BasisCriterion::Bic => {
1285 evaluate_nbasis_info_criterion(
1286 data,
1287 argvals,
1288 nbasis_range,
1289 basis_type,
1290 criterion,
1291 lambda,
1292 )
1293 }
1294 BasisCriterion::Cv => {
1295 evaluate_nbasis_cv(data, argvals, nbasis_range, basis_type, lambda, n_folds)
1296 }
1297 };
1298
1299 let (best_idx, _) = scores
1300 .iter()
1301 .enumerate()
1302 .min_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))?;
1303
1304 Some(BasisNbasisCvResult {
1305 optimal_nbasis: nbasis_range[best_idx],
1306 scores,
1307 nbasis_range: nbasis_range.to_vec(),
1308 criterion,
1309 })
1310}
1311
1312#[cfg(test)]
1313mod tests {
1314 use super::*;
1315 use crate::test_helpers::uniform_grid;
1316 use std::f64::consts::PI;
1317
1318 #[test]
1319 fn test_bspline_penalty_matrix_symmetric() {
1320 let t = uniform_grid(101);
1321 let penalty = bspline_penalty_matrix(&t, 15, 4, 2);
1322 let _k = 15; let actual_k = (penalty.len() as f64).sqrt() as usize;
1324 for i in 0..actual_k {
1325 for j in 0..actual_k {
1326 assert!(
1327 (penalty[i + j * actual_k] - penalty[j + i * actual_k]).abs() < 1e-10,
1328 "Penalty matrix not symmetric at ({}, {})",
1329 i,
1330 j
1331 );
1332 }
1333 }
1334 }
1335
1336 #[test]
1337 fn test_bspline_penalty_matrix_positive_semidefinite() {
1338 let t = uniform_grid(101);
1339 let penalty = bspline_penalty_matrix(&t, 10, 4, 2);
1340 let k = (penalty.len() as f64).sqrt() as usize;
1341 for i in 0..k {
1343 assert!(
1344 penalty[i + i * k] >= -1e-10,
1345 "Diagonal element {} is negative: {}",
1346 i,
1347 penalty[i + i * k]
1348 );
1349 }
1350 }
1351
1352 #[test]
1353 fn test_fourier_penalty_diagonal() {
1354 let penalty = fourier_penalty_matrix(7, 1.0, 2);
1355 for i in 0..7 {
1357 for j in 0..7 {
1358 if i != j {
1359 assert!(
1360 penalty[i + j * 7].abs() < 1e-10,
1361 "Off-diagonal ({},{}) = {}",
1362 i,
1363 j,
1364 penalty[i + j * 7]
1365 );
1366 }
1367 }
1368 }
1369 assert!(penalty[0].abs() < 1e-10);
1371 assert!(penalty[1 + 7] > 0.0);
1373 assert!(penalty[3 + 3 * 7] > penalty[1 + 7]);
1374 }
1375
1376 #[test]
1377 fn test_smooth_basis_bspline() {
1378 let m = 101;
1379 let n = 5;
1380 let t = uniform_grid(m);
1381
1382 let mut data = FdMatrix::zeros(n, m);
1384 for i in 0..n {
1385 for j in 0..m {
1386 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.1 * (i as f64 * 0.3 + j as f64 * 0.01);
1387 }
1388 }
1389
1390 let nbasis = 15;
1391 let penalty = bspline_penalty_matrix(&t, nbasis, 4, 2);
1392 let _actual_k = (penalty.len() as f64).sqrt() as usize;
1393
1394 let fdpar = FdPar {
1395 basis_type: BasisType::Bspline { order: 4 },
1396 nbasis,
1397 lambda: 1e-4,
1398 lfd_order: 2,
1399 penalty_matrix: penalty,
1400 };
1401
1402 let result = smooth_basis(&data, &t, &fdpar);
1403 assert!(result.is_ok(), "smooth_basis should succeed");
1404
1405 let res = result.unwrap();
1406 assert_eq!(res.fitted.shape(), (n, m));
1407 assert_eq!(res.coefficients.nrows(), n);
1408 assert!(res.edf > 0.0, "EDF should be positive");
1409 assert!(res.gcv > 0.0, "GCV should be positive");
1410 }
1411
1412 #[test]
1413 fn test_smooth_basis_fourier() {
1414 let m = 101;
1415 let n = 3;
1416 let t = uniform_grid(m);
1417
1418 let mut data = FdMatrix::zeros(n, m);
1419 for i in 0..n {
1420 for j in 0..m {
1421 data[(i, j)] = (2.0 * PI * t[j]).sin() + (4.0 * PI * t[j]).cos();
1422 }
1423 }
1424
1425 let nbasis = 7;
1426 let period = 1.0;
1427 let penalty = fourier_penalty_matrix(nbasis, period, 2);
1428
1429 let fdpar = FdPar {
1430 basis_type: BasisType::Fourier { period },
1431 nbasis,
1432 lambda: 1e-6,
1433 lfd_order: 2,
1434 penalty_matrix: penalty,
1435 };
1436
1437 let result = smooth_basis(&data, &t, &fdpar);
1438 assert!(result.is_ok());
1439
1440 let res = result.unwrap();
1441 for j in 0..m {
1443 let expected = (2.0 * PI * t[j]).sin() + (4.0 * PI * t[j]).cos();
1444 assert!(
1445 (res.fitted[(0, j)] - expected).abs() < 0.1,
1446 "Fourier fit poor at j={}: got {}, expected {}",
1447 j,
1448 res.fitted[(0, j)],
1449 expected
1450 );
1451 }
1452 }
1453
1454 #[test]
1455 fn test_smooth_basis_gcv_selects_reasonable_lambda() {
1456 let m = 101;
1457 let n = 5;
1458 let t = uniform_grid(m);
1459
1460 let mut data = FdMatrix::zeros(n, m);
1461 for i in 0..n {
1462 for j in 0..m {
1463 data[(i, j)] =
1464 (2.0 * PI * t[j]).sin() + 0.1 * ((i * 37 + j * 13) % 20) as f64 / 20.0;
1465 }
1466 }
1467
1468 let basis_type = BasisType::Bspline { order: 4 };
1469 let result = smooth_basis_gcv(&data, &t, &basis_type, 15, 2, (-8.0, 4.0), 25);
1470 assert!(result.is_some(), "GCV search should succeed");
1471 }
1472
1473 #[test]
1474 fn test_smooth_basis_aic_matches_brute_force_grid() {
1475 let m = 81;
1479 let n = 4;
1480 let t = uniform_grid(m);
1481 let mut data = FdMatrix::zeros(n, m);
1482 for i in 0..n {
1483 for j in 0..m {
1484 data[(i, j)] =
1485 (2.0 * PI * t[j]).sin() + 0.15 * ((i * 41 + j * 7) % 23) as f64 / 23.0;
1486 }
1487 }
1488
1489 let basis_type = BasisType::Bspline { order: 4 };
1490 let nbasis = 12;
1491 let lfd_order = 2;
1492 let range = (-8.0, 4.0);
1493 let n_grid = 25;
1494
1495 let penalty = bspline_penalty_matrix(&t, nbasis, 4, lfd_order);
1497 let (lo, hi) = range;
1498 let mut brute_best_aic = f64::INFINITY;
1499 for k in 0..n_grid {
1500 let log_lam = lo + (hi - lo) * k as f64 / (n_grid - 1) as f64;
1501 let lam = 10.0_f64.powf(log_lam);
1502 let fdpar = FdPar {
1503 basis_type: basis_type.clone(),
1504 nbasis,
1505 lambda: lam,
1506 lfd_order,
1507 penalty_matrix: penalty.clone(),
1508 };
1509 if let Ok(result) = smooth_basis(&data, &t, &fdpar) {
1510 if result.aic < brute_best_aic {
1511 brute_best_aic = result.aic;
1512 }
1513 }
1514 }
1515
1516 let selected =
1517 smooth_basis_aic(&data, &t, &basis_type, nbasis, lfd_order, range, n_grid).unwrap();
1518 assert!(
1519 (selected.aic - brute_best_aic).abs() < 1e-9,
1520 "selected aic={}, brute-force min aic={}",
1521 selected.aic,
1522 brute_best_aic
1523 );
1524 }
1525
1526 #[test]
1527 fn test_smooth_basis_aic_prefers_smoother_fit_than_smallest_lambda() {
1528 let m = 81;
1531 let n = 4;
1532 let t = uniform_grid(m);
1533 let mut data = FdMatrix::zeros(n, m);
1534 for i in 0..n {
1536 for j in 0..m {
1537 let noise = ((((i * 91 + j * 53) % 101) as f64) / 101.0) - 0.5;
1538 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.6 * noise;
1539 }
1540 }
1541
1542 let basis_type = BasisType::Bspline { order: 4 };
1543 let nbasis = 15;
1544 let lfd_order = 2;
1545 let range = (-8.0, 4.0);
1546 let n_grid = 25;
1547
1548 let penalty = bspline_penalty_matrix(&t, nbasis, 4, lfd_order);
1550 let smallest_lam = 10.0_f64.powf(range.0);
1551 let fdpar_small = FdPar {
1552 basis_type: basis_type.clone(),
1553 nbasis,
1554 lambda: smallest_lam,
1555 lfd_order,
1556 penalty_matrix: penalty.clone(),
1557 };
1558 let overfit = smooth_basis(&data, &t, &fdpar_small).unwrap();
1559
1560 let selected =
1561 smooth_basis_aic(&data, &t, &basis_type, nbasis, lfd_order, range, n_grid).unwrap();
1562
1563 assert!(
1564 selected.edf < overfit.edf,
1565 "AIC-selected edf ({}) should be smaller (smoother) than the smallest-lambda edf ({})",
1566 selected.edf,
1567 overfit.edf
1568 );
1569 }
1570
1571 #[test]
1572 fn test_smooth_basis_large_lambda_reduces_edf() {
1573 let m = 101;
1574 let n = 3;
1575 let t = uniform_grid(m);
1576
1577 let mut data = FdMatrix::zeros(n, m);
1578 for i in 0..n {
1579 for j in 0..m {
1580 data[(i, j)] = (2.0 * PI * t[j]).sin();
1581 }
1582 }
1583
1584 let nbasis = 15;
1585 let penalty = bspline_penalty_matrix(&t, nbasis, 4, 2);
1586 let _actual_k = (penalty.len() as f64).sqrt() as usize;
1587
1588 let fdpar_small = FdPar {
1589 basis_type: BasisType::Bspline { order: 4 },
1590 nbasis,
1591 lambda: 1e-8,
1592 lfd_order: 2,
1593 penalty_matrix: penalty.clone(),
1594 };
1595 let fdpar_large = FdPar {
1596 basis_type: BasisType::Bspline { order: 4 },
1597 nbasis,
1598 lambda: 1e2,
1599 lfd_order: 2,
1600 penalty_matrix: penalty,
1601 };
1602
1603 let res_small = smooth_basis(&data, &t, &fdpar_small).unwrap();
1604 let res_large = smooth_basis(&data, &t, &fdpar_large).unwrap();
1605
1606 assert!(
1607 res_large.edf < res_small.edf,
1608 "Larger lambda should reduce EDF: {} vs {}",
1609 res_large.edf,
1610 res_small.edf
1611 );
1612 }
1613
1614 #[test]
1617 fn test_basis_nbasis_cv_gcv() {
1618 let m = 101;
1619 let n = 5;
1620 let t = uniform_grid(m);
1621 let mut data = FdMatrix::zeros(n, m);
1622 for i in 0..n {
1623 for j in 0..m {
1624 data[(i, j)] =
1625 (2.0 * PI * t[j]).sin() + 0.1 * ((i * 37 + j * 13) % 20) as f64 / 20.0;
1626 }
1627 }
1628
1629 let nbasis_range: Vec<usize> = (4..=20).step_by(2).collect();
1630 let result = basis_nbasis_cv(
1631 &data,
1632 &t,
1633 &nbasis_range,
1634 &BasisType::Bspline { order: 4 },
1635 BasisCriterion::Gcv,
1636 5,
1637 1e-4,
1638 );
1639 assert!(result.is_some());
1640 let res = result.unwrap();
1641 assert!(nbasis_range.contains(&res.optimal_nbasis));
1642 assert_eq!(res.scores.len(), nbasis_range.len());
1643 assert_eq!(res.criterion, BasisCriterion::Gcv);
1644 }
1645
1646 #[test]
1647 fn test_basis_nbasis_cv_aic_bic() {
1648 let m = 51;
1649 let n = 5;
1650 let t = uniform_grid(m);
1651 let mut data = FdMatrix::zeros(n, m);
1652 for i in 0..n {
1653 for j in 0..m {
1654 data[(i, j)] = (2.0 * PI * t[j]).sin();
1655 }
1656 }
1657
1658 let nbasis_range: Vec<usize> = vec![5, 7, 9, 11];
1659 let aic_result = basis_nbasis_cv(
1660 &data,
1661 &t,
1662 &nbasis_range,
1663 &BasisType::Bspline { order: 4 },
1664 BasisCriterion::Aic,
1665 5,
1666 0.0,
1667 );
1668 let bic_result = basis_nbasis_cv(
1669 &data,
1670 &t,
1671 &nbasis_range,
1672 &BasisType::Bspline { order: 4 },
1673 BasisCriterion::Bic,
1674 5,
1675 0.0,
1676 );
1677 assert!(aic_result.is_some());
1678 assert!(bic_result.is_some());
1679 }
1680
1681 #[test]
1682 fn test_basis_nbasis_cv_kfold() {
1683 let m = 51;
1684 let n = 10;
1685 let t = uniform_grid(m);
1686 let mut data = FdMatrix::zeros(n, m);
1687 for i in 0..n {
1688 for j in 0..m {
1689 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.05 * ((i * 7 + j * 3) % 10) as f64;
1690 }
1691 }
1692
1693 let nbasis_range: Vec<usize> = vec![5, 7, 9];
1694 let result = basis_nbasis_cv(
1695 &data,
1696 &t,
1697 &nbasis_range,
1698 &BasisType::Bspline { order: 4 },
1699 BasisCriterion::Cv,
1700 5,
1701 1e-4,
1702 );
1703 assert!(result.is_some());
1704 let res = result.unwrap();
1705 assert!(nbasis_range.contains(&res.optimal_nbasis));
1706 assert_eq!(res.criterion, BasisCriterion::Cv);
1707 }
1708
1709 #[test]
1714 fn test_basis_nbasis_cv_penalizes_overfitting() {
1715 let m = 120;
1716 let n = 6;
1717 let t = uniform_grid(m);
1718 let mut data = FdMatrix::zeros(n, m);
1719 for i in 0..n {
1720 for j in 0..m {
1721 let noise = 0.2 * (((i * 31 + j * 17) % 13) as f64 / 13.0 - 0.5);
1723 data[(i, j)] = (2.0 * PI * t[j]).sin() + noise;
1724 }
1725 }
1726
1727 let nbasis_range: Vec<usize> = vec![5, 8, 12, 20, 30];
1728 let res = basis_nbasis_cv(
1729 &data,
1730 &t,
1731 &nbasis_range,
1732 &BasisType::Bspline { order: 4 },
1733 BasisCriterion::Cv,
1734 5,
1735 1e-6,
1736 )
1737 .unwrap();
1738
1739 assert_ne!(
1740 res.optimal_nbasis, 30,
1741 "CV must not always select the maximum n_basis (GH #33); scores={:?}",
1742 res.scores
1743 );
1744 let monotone_decreasing = res.scores.windows(2).all(|w| w[1] <= w[0] + 1e-12);
1745 assert!(
1746 !monotone_decreasing,
1747 "CV scores must not be monotone-decreasing in n_basis; scores={:?}",
1748 res.scores
1749 );
1750 }
1751
1752 fn make_test_data(n: usize, m: usize) -> (FdMatrix, Vec<f64>) {
1756 let t: Vec<f64> = (0..m).map(|i| i as f64 / (m - 1) as f64).collect();
1757 let mut data = FdMatrix::zeros(n, m);
1758 for i in 0..n {
1759 for j in 0..m {
1760 data[(i, j)] = (2.0 * PI * t[j]).sin()
1761 + 0.1 * (10.0 * t[j]).sin()
1762 + 0.05 * ((i * 37 + j * 13) % 20) as f64 / 20.0;
1763 }
1764 }
1765 (data, t)
1766 }
1767
1768 fn make_bspline_fdpar(argvals: &[f64], nbasis: usize, lambda: f64) -> FdPar {
1770 let penalty = bspline_penalty_matrix(argvals, nbasis, 4, 2);
1771 FdPar {
1772 basis_type: BasisType::Bspline { order: 4 },
1773 nbasis,
1774 lambda,
1775 lfd_order: 2,
1776 penalty_matrix: penalty,
1777 }
1778 }
1779
1780 fn make_fourier_fdpar(nbasis: usize, period: f64, lambda: f64) -> FdPar {
1782 let penalty = fourier_penalty_matrix(nbasis, period, 2);
1783 FdPar {
1784 basis_type: BasisType::Fourier { period },
1785 nbasis,
1786 lambda,
1787 lfd_order: 2,
1788 penalty_matrix: penalty,
1789 }
1790 }
1791
1792 #[test]
1795 fn test_basis_type_bspline_variant() {
1796 let bt = BasisType::Bspline { order: 4 };
1797 assert_eq!(bt, BasisType::Bspline { order: 4 });
1798 assert_ne!(bt, BasisType::Bspline { order: 3 });
1800 }
1801
1802 #[test]
1803 fn test_basis_type_fourier_variant() {
1804 let bt = BasisType::Fourier { period: 1.0 };
1805 assert_eq!(bt, BasisType::Fourier { period: 1.0 });
1806 assert_ne!(bt, BasisType::Fourier { period: 2.0 });
1807 }
1808
1809 #[test]
1810 fn test_basis_type_cross_variant_inequality() {
1811 let bspline = BasisType::Bspline { order: 4 };
1812 let fourier = BasisType::Fourier { period: 1.0 };
1813 assert_ne!(bspline, fourier);
1814 }
1815
1816 #[test]
1817 fn test_basis_type_clone_and_debug() {
1818 let bt = BasisType::Bspline { order: 4 };
1819 let cloned = bt.clone();
1820 assert_eq!(bt, cloned);
1821 let debug_str = format!("{:?}", bt);
1822 assert!(debug_str.contains("Bspline"));
1823 assert!(debug_str.contains("4"));
1824 }
1825
1826 #[test]
1829 fn test_fdpar_construction_and_fields() {
1830 let penalty = vec![1.0, 0.0, 0.0, 1.0];
1831 let fdpar = FdPar {
1832 basis_type: BasisType::Bspline { order: 4 },
1833 nbasis: 2,
1834 lambda: 0.01,
1835 lfd_order: 2,
1836 penalty_matrix: penalty.clone(),
1837 };
1838 assert_eq!(fdpar.nbasis, 2);
1839 assert!((fdpar.lambda - 0.01).abs() < 1e-15);
1840 assert_eq!(fdpar.lfd_order, 2);
1841 assert_eq!(fdpar.penalty_matrix.len(), 4);
1842 }
1843
1844 #[test]
1845 fn test_fdpar_clone_and_debug() {
1846 let t = uniform_grid(50);
1847 let fdpar = make_bspline_fdpar(&t, 8, 1e-3);
1848 let cloned = fdpar.clone();
1849 assert_eq!(fdpar, cloned);
1850 let debug_str = format!("{:?}", fdpar);
1851 assert!(debug_str.contains("FdPar"));
1852 }
1853
1854 #[test]
1857 fn test_basis_criterion_variants() {
1858 assert_eq!(BasisCriterion::Gcv, BasisCriterion::Gcv);
1859 assert_eq!(BasisCriterion::Cv, BasisCriterion::Cv);
1860 assert_eq!(BasisCriterion::Aic, BasisCriterion::Aic);
1861 assert_eq!(BasisCriterion::Bic, BasisCriterion::Bic);
1862 assert_ne!(BasisCriterion::Gcv, BasisCriterion::Aic);
1863 assert_ne!(BasisCriterion::Cv, BasisCriterion::Bic);
1864 }
1865
1866 #[test]
1867 fn test_basis_criterion_copy() {
1868 let c = BasisCriterion::Gcv;
1869 let copied = c; assert_eq!(c, copied);
1871 }
1872
1873 #[test]
1874 fn test_basis_criterion_debug() {
1875 let debug_str = format!("{:?}", BasisCriterion::Bic);
1876 assert!(debug_str.contains("Bic"));
1877 }
1878
1879 #[test]
1882 fn test_smooth_basis_result_all_fields() {
1883 let (data, t) = make_test_data(3, 50);
1884 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1885 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1886
1887 assert_eq!(res.coefficients.nrows(), 3);
1889 assert!(res.coefficients.ncols() > 0);
1890 assert_eq!(res.nbasis, res.coefficients.ncols());
1891 assert_eq!(res.fitted.shape(), (3, 50));
1893 assert!(res.edf > 0.0 && res.edf <= res.nbasis as f64);
1895 assert!(res.gcv.is_finite());
1897 assert!(res.aic.is_finite());
1898 assert!(res.bic.is_finite());
1899 let k = res.nbasis;
1901 assert_eq!(res.penalty_matrix.len(), k * k);
1902 }
1903
1904 #[test]
1905 fn test_smooth_basis_result_clone() {
1906 let (data, t) = make_test_data(2, 50);
1907 let fdpar = make_bspline_fdpar(&t, 8, 1e-3);
1908 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1909 let cloned = res.clone();
1910 assert_eq!(res, cloned);
1911 }
1912
1913 #[test]
1916 fn test_smooth_basis_bspline_coefficient_shape() {
1917 let (data, t) = make_test_data(4, 50);
1918 let nbasis = 12;
1919 let fdpar = make_bspline_fdpar(&t, nbasis, 1e-4);
1920 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1921 assert_eq!(res.coefficients.nrows(), 4);
1922 assert!(res.coefficients.ncols() >= 2);
1924 assert_eq!(res.nbasis, res.coefficients.ncols());
1925 }
1926
1927 #[test]
1928 fn test_smooth_basis_bspline_fitted_values_shape() {
1929 let m = 80;
1930 let n = 6;
1931 let (data, t) = make_test_data(n, m);
1932 let fdpar = make_bspline_fdpar(&t, 15, 1e-4);
1933 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1934 assert_eq!(res.fitted.shape(), (n, m));
1935 }
1936
1937 #[test]
1938 fn test_smooth_basis_bspline_zero_lambda_interpolates() {
1939 let m = 30;
1941 let n = 2;
1942 let (data, t) = make_test_data(n, m);
1943 let fdpar = make_bspline_fdpar(&t, 15, 0.0);
1944 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1945
1946 let mut max_resid = 0.0_f64;
1948 for i in 0..n {
1949 for j in 0..m {
1950 let resid = (data[(i, j)] - res.fitted[(i, j)]).abs();
1951 max_resid = max_resid.max(resid);
1952 }
1953 }
1954 assert!(
1955 max_resid < 0.5,
1956 "Zero-lambda B-spline should closely interpolate; max_resid = {}",
1957 max_resid
1958 );
1959 }
1960
1961 #[test]
1962 fn test_smooth_basis_bspline_large_lambda_oversmooths() {
1963 let m = 50;
1966 let n = 1;
1967 let (data, t) = make_test_data(n, m);
1968
1969 let fdpar_small = make_bspline_fdpar(&t, 15, 1e-6);
1970 let res_small = smooth_basis(&data, &t, &fdpar_small).unwrap();
1971
1972 let fdpar_large = make_bspline_fdpar(&t, 15, 1e6);
1973 let res_large = smooth_basis(&data, &t, &fdpar_large).unwrap();
1974
1975 let compute_variance = |fitted: &FdMatrix, row: usize, ncols: usize| -> f64 {
1976 let vals: Vec<f64> = (0..ncols).map(|j| fitted[(row, j)]).collect();
1977 let mean = vals.iter().sum::<f64>() / ncols as f64;
1978 vals.iter().map(|&v| (v - mean).powi(2)).sum::<f64>() / ncols as f64
1979 };
1980
1981 let var_small = compute_variance(&res_small.fitted, 0, m);
1982 let var_large = compute_variance(&res_large.fitted, 0, m);
1983 assert!(
1984 var_large < var_small,
1985 "Large lambda should yield lower variance fit: var_large={}, var_small={}",
1986 var_large,
1987 var_small
1988 );
1989 }
1990
1991 #[test]
1992 fn test_smooth_basis_bspline_penalty_effect_on_smoothness() {
1993 let m = 50;
1995 let n = 1;
1996 let (data, t) = make_test_data(n, m);
1997
1998 let fdpar_small = make_bspline_fdpar(&t, 15, 1e-8);
1999 let fdpar_large = make_bspline_fdpar(&t, 15, 1.0);
2000
2001 let res_small = smooth_basis(&data, &t, &fdpar_small).unwrap();
2002 let res_large = smooth_basis(&data, &t, &fdpar_large).unwrap();
2003
2004 let roughness = |fitted: &FdMatrix, row: usize, ncols: usize| -> f64 {
2006 (1..ncols - 1)
2007 .map(|j| {
2008 let d2 = fitted[(row, j + 1)] - 2.0 * fitted[(row, j)] + fitted[(row, j - 1)];
2009 d2 * d2
2010 })
2011 .sum::<f64>()
2012 };
2013
2014 let r_small = roughness(&res_small.fitted, 0, m);
2015 let r_large = roughness(&res_large.fitted, 0, m);
2016 assert!(
2017 r_large < r_small,
2018 "Larger lambda should produce smoother fit: roughness_large={}, roughness_small={}",
2019 r_large,
2020 r_small
2021 );
2022 }
2023
2024 #[test]
2025 fn test_smooth_basis_bspline_single_curve() {
2026 let m = 50;
2027 let (data, t) = make_test_data(1, m);
2028 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2029 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2030 assert_eq!(res.fitted.nrows(), 1);
2031 assert_eq!(res.fitted.ncols(), m);
2032 assert!(res.gcv.is_finite());
2033 }
2034
2035 #[test]
2036 fn test_smooth_basis_bspline_many_curves() {
2037 let m = 50;
2038 let n = 20;
2039 let (data, t) = make_test_data(n, m);
2040 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2041 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2042 assert_eq!(res.fitted.nrows(), n);
2043 assert_eq!(res.coefficients.nrows(), n);
2044 }
2045
2046 #[test]
2047 fn test_smooth_basis_bspline_minimal_nbasis() {
2048 let m = 50;
2050 let (data, t) = make_test_data(1, m);
2051 let fdpar = make_bspline_fdpar(&t, 2, 1e-4);
2052 let res = smooth_basis(&data, &t, &fdpar);
2053 assert!(res.is_ok());
2055 }
2056
2057 #[test]
2058 fn test_smooth_basis_bspline_different_orders() {
2059 let m = 50;
2060 let (data, t) = make_test_data(2, m);
2061 let penalty3 = bspline_penalty_matrix(&t, 10, 3, 2);
2063 let fdpar3 = FdPar {
2064 basis_type: BasisType::Bspline { order: 3 },
2065 nbasis: 10,
2066 lambda: 1e-4,
2067 lfd_order: 2,
2068 penalty_matrix: penalty3,
2069 };
2070 let res3 = smooth_basis(&data, &t, &fdpar3);
2071 assert!(res3.is_ok());
2072
2073 let penalty5 = bspline_penalty_matrix(&t, 10, 5, 2);
2075 let fdpar5 = FdPar {
2076 basis_type: BasisType::Bspline { order: 5 },
2077 nbasis: 10,
2078 lambda: 1e-4,
2079 lfd_order: 2,
2080 penalty_matrix: penalty5,
2081 };
2082 let res5 = smooth_basis(&data, &t, &fdpar5);
2083 assert!(res5.is_ok());
2084 }
2085
2086 #[test]
2089 fn test_smooth_basis_fourier_coefficient_shape() {
2090 let m = 50;
2091 let n = 3;
2092 let t = uniform_grid(m);
2093 let mut data = FdMatrix::zeros(n, m);
2094 for i in 0..n {
2095 for j in 0..m {
2096 data[(i, j)] = (2.0 * PI * t[j]).sin();
2097 }
2098 }
2099 let nbasis = 7;
2100 let fdpar = make_fourier_fdpar(nbasis, 1.0, 1e-6);
2101 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2102 assert_eq!(res.coefficients.nrows(), n);
2103 assert_eq!(res.coefficients.ncols(), nbasis);
2104 assert_eq!(res.nbasis, nbasis);
2105 }
2106
2107 #[test]
2108 fn test_smooth_basis_fourier_fits_pure_sine() {
2109 let m = 100;
2111 let t = uniform_grid(m);
2112 let mut data = FdMatrix::zeros(1, m);
2113 for j in 0..m {
2114 data[(0, j)] = (2.0 * PI * t[j]).sin();
2115 }
2116 let fdpar = make_fourier_fdpar(5, 1.0, 1e-8);
2117 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2118
2119 for j in 0..m {
2120 let expected = (2.0 * PI * t[j]).sin();
2121 assert!(
2122 (res.fitted[(0, j)] - expected).abs() < 0.05,
2123 "Fourier should fit pure sine; j={}, got={}, expected={}",
2124 j,
2125 res.fitted[(0, j)],
2126 expected
2127 );
2128 }
2129 }
2130
2131 #[test]
2132 fn test_smooth_basis_fourier_different_periods() {
2133 let m = 50;
2134 let t = uniform_grid(m);
2135 let mut data = FdMatrix::zeros(1, m);
2136 for j in 0..m {
2137 data[(0, j)] = (2.0 * PI * t[j]).sin();
2138 }
2139
2140 let fdpar1 = make_fourier_fdpar(7, 1.0, 1e-6);
2142 let res1 = smooth_basis(&data, &t, &fdpar1).unwrap();
2143
2144 let fdpar2 = make_fourier_fdpar(7, 2.0, 1e-6);
2146 let res2 = smooth_basis(&data, &t, &fdpar2).unwrap();
2147
2148 assert_eq!(res1.fitted.shape(), (1, m));
2150 assert_eq!(res2.fitted.shape(), (1, m));
2151 }
2152
2153 #[test]
2154 fn test_smooth_basis_fourier_zero_lambda() {
2155 let m = 50;
2156 let t = uniform_grid(m);
2157 let mut data = FdMatrix::zeros(1, m);
2158 for j in 0..m {
2159 data[(0, j)] = (2.0 * PI * t[j]).sin() + (4.0 * PI * t[j]).cos();
2160 }
2161 let fdpar = make_fourier_fdpar(9, 1.0, 0.0);
2162 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2163 assert_eq!(res.fitted.shape(), (1, m));
2164 assert!(res.edf > 1.0);
2166 }
2167
2168 #[test]
2169 fn test_smooth_basis_fourier_large_lambda() {
2170 let m = 50;
2171 let t = uniform_grid(m);
2172 let mut data = FdMatrix::zeros(1, m);
2173 for j in 0..m {
2174 data[(0, j)] = (2.0 * PI * t[j]).sin();
2175 }
2176 let fdpar = make_fourier_fdpar(9, 1.0, 1e6);
2177 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2178 assert!(
2180 res.edf < 5.0,
2181 "Large lambda should reduce EDF; edf={}",
2182 res.edf
2183 );
2184 }
2185
2186 #[test]
2189 fn test_smooth_basis_lambda_gradient_edf() {
2190 let m = 50;
2192 let (data, t) = make_test_data(3, m);
2193 let lambdas = [1e-8, 1e-4, 1e-2, 1.0, 1e2];
2194 let mut prev_edf = f64::INFINITY;
2195 for &lam in &lambdas {
2196 let fdpar = make_bspline_fdpar(&t, 12, lam);
2197 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2198 assert!(
2199 res.edf <= prev_edf + 0.01,
2200 "EDF should decrease: lambda={}, edf={}, prev_edf={}",
2201 lam,
2202 res.edf,
2203 prev_edf
2204 );
2205 prev_edf = res.edf;
2206 }
2207 }
2208
2209 #[test]
2210 fn test_smooth_basis_lambda_gradient_rss() {
2211 let m = 50;
2213 let n = 2;
2214 let (data, t) = make_test_data(n, m);
2215 let lambdas = [0.0, 1e-6, 1e-2, 1.0, 1e4];
2216 let mut prev_rss = -1.0;
2217 for &lam in &lambdas {
2218 let fdpar = make_bspline_fdpar(&t, 12, lam);
2219 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2220 let mut rss = 0.0;
2221 for i in 0..n {
2222 for j in 0..m {
2223 rss += (data[(i, j)] - res.fitted[(i, j)]).powi(2);
2224 }
2225 }
2226 assert!(
2227 rss >= prev_rss - 1e-8,
2228 "RSS should increase: lambda={}, rss={}, prev_rss={}",
2229 lam,
2230 rss,
2231 prev_rss
2232 );
2233 prev_rss = rss;
2234 }
2235 }
2236
2237 #[test]
2240 fn test_smooth_basis_empty_data_rows() {
2241 let t = uniform_grid(50);
2242 let data = FdMatrix::zeros(0, 50);
2243 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2244 let res = smooth_basis(&data, &t, &fdpar);
2245 assert!(res.is_err());
2246 }
2247
2248 #[test]
2249 fn test_smooth_basis_empty_data_cols() {
2250 let data = FdMatrix::zeros(5, 0);
2251 let fdpar = FdPar {
2252 basis_type: BasisType::Bspline { order: 4 },
2253 nbasis: 10,
2254 lambda: 1e-4,
2255 lfd_order: 2,
2256 penalty_matrix: vec![0.0; 100],
2257 };
2258 let res = smooth_basis(&data, &[], &fdpar);
2259 assert!(res.is_err());
2260 }
2261
2262 #[test]
2263 fn test_smooth_basis_mismatched_argvals() {
2264 let t = uniform_grid(50);
2265 let data = FdMatrix::zeros(3, 40); let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2267 let res = smooth_basis(&data, &t, &fdpar);
2268 assert!(res.is_err());
2269 }
2270
2271 #[test]
2272 fn test_smooth_basis_nbasis_too_small() {
2273 let t = uniform_grid(50);
2274 let data = FdMatrix::zeros(3, 50);
2275 let fdpar = FdPar {
2277 basis_type: BasisType::Bspline { order: 4 },
2278 nbasis: 1,
2279 lambda: 1e-4,
2280 lfd_order: 2,
2281 penalty_matrix: vec![0.0; 1],
2282 };
2283 let res = smooth_basis(&data, &t, &fdpar);
2284 assert!(res.is_err());
2285 }
2286
2287 #[test]
2288 fn test_smooth_basis_error_is_invalid_dimension() {
2289 let t = uniform_grid(50);
2290 let data = FdMatrix::zeros(0, 50);
2291 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2292 let err = smooth_basis(&data, &t, &fdpar).unwrap_err();
2293 match err {
2294 crate::FdarError::InvalidDimension { .. } => {} other => panic!("Expected InvalidDimension, got {:?}", other),
2296 }
2297 }
2298
2299 #[test]
2302 fn test_bspline_penalty_matrix_different_orders() {
2303 let t = uniform_grid(101);
2304 let p1 = bspline_penalty_matrix(&t, 10, 4, 1);
2306 let p2 = bspline_penalty_matrix(&t, 10, 4, 2);
2308 assert_eq!(p1.len(), p2.len());
2310 let diff: f64 = p1.iter().zip(p2.iter()).map(|(a, b)| (a - b).abs()).sum();
2312 assert!(
2313 diff > 1e-10,
2314 "Different lfd_orders should produce different penalties"
2315 );
2316 }
2317
2318 #[test]
2319 fn test_bspline_penalty_matrix_edge_cases() {
2320 let t = vec![0.0];
2322 let p = bspline_penalty_matrix(&t, 10, 4, 2);
2323 assert!(p.iter().all(|&v| v == 0.0));
2325
2326 let t2 = uniform_grid(50);
2328 let p2 = bspline_penalty_matrix(&t2, 1, 4, 2);
2329 assert!(p2.iter().all(|&v| v == 0.0));
2330
2331 let p3 = bspline_penalty_matrix(&t2, 10, 4, 4);
2333 assert!(p3.iter().all(|&v| v == 0.0));
2334 }
2335
2336 #[test]
2337 fn test_bspline_penalty_nonnegative_diagonal() {
2338 let t = uniform_grid(101);
2339 for nbasis in [5, 10, 20] {
2340 let p = bspline_penalty_matrix(&t, nbasis, 4, 2);
2341 let k = (p.len() as f64).sqrt() as usize;
2342 for i in 0..k {
2343 assert!(
2344 p[i + i * k] >= -1e-10,
2345 "Diagonal ({},{}) negative for nbasis={}: {}",
2346 i,
2347 i,
2348 nbasis,
2349 p[i + i * k]
2350 );
2351 }
2352 }
2353 }
2354
2355 #[test]
2356 fn test_fourier_penalty_increasing_with_frequency() {
2357 let penalty = fourier_penalty_matrix(11, 1.0, 2);
2358 let k = 11;
2359 assert!(penalty[0].abs() < 1e-15);
2361 let mut prev_eigenval = 0.0;
2363 for freq in 1..=5 {
2364 let idx_sin = 2 * freq - 1;
2365 let eigenval = penalty[idx_sin + idx_sin * k];
2366 assert!(
2367 eigenval > prev_eigenval,
2368 "Higher frequency should have larger penalty: freq={}, eigenval={}, prev={}",
2369 freq,
2370 eigenval,
2371 prev_eigenval
2372 );
2373 prev_eigenval = eigenval;
2374 let idx_cos = 2 * freq;
2376 if idx_cos < k {
2377 assert!(
2378 (penalty[idx_cos + idx_cos * k] - eigenval).abs() < 1e-10,
2379 "Sin and cos penalty should match at freq {}",
2380 freq
2381 );
2382 }
2383 }
2384 }
2385
2386 #[test]
2387 fn test_fourier_penalty_different_periods() {
2388 let p1 = fourier_penalty_matrix(7, 1.0, 2);
2389 let p2 = fourier_penalty_matrix(7, 2.0, 2);
2390 for i in 1..7 {
2392 assert!(
2393 p2[i + i * 7] < p1[i + i * 7] || (p1[i + i * 7] == 0.0 && p2[i + i * 7] == 0.0),
2394 "Longer period should have smaller penalties at i={}",
2395 i
2396 );
2397 }
2398 }
2399
2400 #[test]
2401 fn test_fourier_penalty_first_order() {
2402 let p = fourier_penalty_matrix(5, 1.0, 1);
2404 let omega1 = 2.0 * PI;
2406 let expected1 = omega1.powi(2);
2407 assert!(
2408 (p[1 + 5] - expected1).abs() < 1e-6,
2409 "First-order penalty eigenval: got {}, expected {}",
2410 p[1 + 5],
2411 expected1
2412 );
2413 }
2414
2415 #[test]
2416 fn test_fourier_penalty_zero_nbasis() {
2417 let p = fourier_penalty_matrix(0, 1.0, 2);
2418 assert!(p.is_empty());
2419 }
2420
2421 #[test]
2422 fn test_fourier_penalty_nbasis_one() {
2423 let p = fourier_penalty_matrix(1, 1.0, 2);
2424 assert_eq!(p.len(), 1);
2425 assert!(p[0].abs() < 1e-15); }
2427
2428 #[test]
2431 fn test_smooth_basis_gcv_returns_valid_result() {
2432 let (data, t) = make_test_data(5, 50);
2433 let bt = BasisType::Bspline { order: 4 };
2434 let result = smooth_basis_gcv(&data, &t, &bt, 12, 2, (-6.0, 2.0), 20);
2435 assert!(result.is_some());
2436 let res = result.unwrap();
2437 assert_eq!(res.fitted.shape(), (5, 50));
2438 assert!(res.gcv.is_finite());
2439 assert!(res.edf > 0.0);
2440 }
2441
2442 #[test]
2443 fn test_smooth_basis_gcv_fourier() {
2444 let m = 80;
2445 let t = uniform_grid(m);
2446 let mut data = FdMatrix::zeros(3, m);
2447 for i in 0..3 {
2448 for j in 0..m {
2449 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.5 * (4.0 * PI * t[j]).cos();
2450 }
2451 }
2452 let bt = BasisType::Fourier { period: 1.0 };
2453 let result = smooth_basis_gcv(&data, &t, &bt, 9, 2, (-8.0, 4.0), 25);
2454 assert!(result.is_some());
2455 let res = result.unwrap();
2456 assert_eq!(res.fitted.nrows(), 3);
2457 assert_eq!(res.nbasis, 9);
2458 }
2459
2460 #[test]
2461 fn test_smooth_basis_gcv_selects_finite_gcv() {
2462 let (data, t) = make_test_data(5, 60);
2463 let bt = BasisType::Bspline { order: 4 };
2464 let res = smooth_basis_gcv(&data, &t, &bt, 12, 2, (-6.0, 2.0), 15).unwrap();
2465 assert!(res.gcv.is_finite());
2466 assert!(res.gcv > 0.0);
2467 }
2468
2469 #[test]
2470 fn test_smooth_basis_gcv_empty_data() {
2471 let data = FdMatrix::zeros(0, 50);
2472 let t = uniform_grid(50);
2473 let bt = BasisType::Bspline { order: 4 };
2474 let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-6.0, 2.0), 10);
2475 assert!(result.is_none());
2477 }
2478
2479 #[test]
2480 fn test_smooth_basis_gcv_empty_argvals() {
2481 let data = FdMatrix::zeros(5, 0);
2482 let bt = BasisType::Bspline { order: 4 };
2483 let result = smooth_basis_gcv(&data, &[], &bt, 10, 2, (-6.0, 2.0), 10);
2484 assert!(result.is_none());
2485 }
2486
2487 #[test]
2488 fn test_smooth_basis_gcv_nbasis_too_small() {
2489 let (data, t) = make_test_data(5, 50);
2490 let bt = BasisType::Bspline { order: 4 };
2491 let result = smooth_basis_gcv(&data, &t, &bt, 1, 2, (-6.0, 2.0), 10);
2492 assert!(result.is_none());
2493 }
2494
2495 #[test]
2496 fn test_smooth_basis_gcv_ngrid_too_small() {
2497 let (data, t) = make_test_data(5, 50);
2498 let bt = BasisType::Bspline { order: 4 };
2499 let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-6.0, 2.0), 1);
2500 assert!(result.is_none());
2501 }
2502
2503 #[test]
2504 fn test_smooth_basis_gcv_narrow_range() {
2505 let (data, t) = make_test_data(3, 50);
2506 let bt = BasisType::Bspline { order: 4 };
2507 let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-3.0, -2.0), 5);
2509 assert!(result.is_some());
2510 }
2511
2512 #[test]
2513 fn test_smooth_basis_gcv_wide_range() {
2514 let (data, t) = make_test_data(3, 50);
2515 let bt = BasisType::Bspline { order: 4 };
2516 let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-12.0, 8.0), 30);
2518 assert!(result.is_some());
2519 }
2520
2521 #[test]
2524 fn test_basis_nbasis_cv_scores_length() {
2525 let (data, t) = make_test_data(5, 50);
2526 let nbasis_range: Vec<usize> = vec![4, 6, 8, 10, 12];
2527 let res = basis_nbasis_cv(
2528 &data,
2529 &t,
2530 &nbasis_range,
2531 &BasisType::Bspline { order: 4 },
2532 BasisCriterion::Gcv,
2533 5,
2534 1e-4,
2535 )
2536 .unwrap();
2537 assert_eq!(res.scores.len(), 5);
2538 assert_eq!(res.nbasis_range.len(), 5);
2539 assert_eq!(res.nbasis_range, nbasis_range);
2540 }
2541
2542 #[test]
2543 fn test_basis_nbasis_cv_optimal_within_range() {
2544 let (data, t) = make_test_data(8, 50);
2545 let nbasis_range: Vec<usize> = vec![5, 7, 9, 11, 13, 15];
2546 for criterion in [
2547 BasisCriterion::Gcv,
2548 BasisCriterion::Aic,
2549 BasisCriterion::Bic,
2550 ] {
2551 let res = basis_nbasis_cv(
2552 &data,
2553 &t,
2554 &nbasis_range,
2555 &BasisType::Bspline { order: 4 },
2556 criterion,
2557 5,
2558 1e-4,
2559 )
2560 .unwrap();
2561 assert!(
2562 nbasis_range.contains(&res.optimal_nbasis),
2563 "optimal_nbasis {} not in range for {:?}",
2564 res.optimal_nbasis,
2565 criterion
2566 );
2567 }
2568 }
2569
2570 #[test]
2571 fn test_basis_nbasis_cv_fourier_gcv() {
2572 let m = 80;
2573 let t = uniform_grid(m);
2574 let mut data = FdMatrix::zeros(5, m);
2575 for i in 0..5 {
2576 for j in 0..m {
2577 data[(i, j)] = (2.0 * PI * t[j]).sin()
2578 + 0.3 * (4.0 * PI * t[j]).cos()
2579 + 0.02 * ((i * 7 + j * 3) % 10) as f64;
2580 }
2581 }
2582 let nbasis_range: Vec<usize> = vec![5, 7, 9, 11];
2583 let res = basis_nbasis_cv(
2584 &data,
2585 &t,
2586 &nbasis_range,
2587 &BasisType::Fourier { period: 1.0 },
2588 BasisCriterion::Gcv,
2589 5,
2590 1e-4,
2591 )
2592 .unwrap();
2593 assert!(nbasis_range.contains(&res.optimal_nbasis));
2594 }
2595
2596 #[test]
2597 fn test_basis_nbasis_cv_fourier_cv() {
2598 let m = 60;
2599 let t = uniform_grid(m);
2600 let n = 10;
2601 let mut data = FdMatrix::zeros(n, m);
2602 for i in 0..n {
2603 for j in 0..m {
2604 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.02 * ((i * 11 + j) % 15) as f64;
2605 }
2606 }
2607 let nbasis_range: Vec<usize> = vec![5, 7, 9];
2608 let res = basis_nbasis_cv(
2609 &data,
2610 &t,
2611 &nbasis_range,
2612 &BasisType::Fourier { period: 1.0 },
2613 BasisCriterion::Cv,
2614 5,
2615 1e-4,
2616 )
2617 .unwrap();
2618 assert!(nbasis_range.contains(&res.optimal_nbasis));
2619 assert_eq!(res.criterion, BasisCriterion::Cv);
2620 }
2621
2622 #[test]
2623 fn test_basis_nbasis_cv_with_nbasis_below_minimum() {
2624 let (data, t) = make_test_data(5, 50);
2626 let nbasis_range: Vec<usize> = vec![1, 5, 10];
2627 let res = basis_nbasis_cv(
2628 &data,
2629 &t,
2630 &nbasis_range,
2631 &BasisType::Bspline { order: 4 },
2632 BasisCriterion::Gcv,
2633 5,
2634 1e-4,
2635 )
2636 .unwrap();
2637 assert!(
2639 res.optimal_nbasis >= 5,
2640 "Should skip invalid nbasis=1, got optimal={}",
2641 res.optimal_nbasis
2642 );
2643 assert!(res.scores[0].is_infinite());
2644 }
2645
2646 #[test]
2647 fn test_basis_nbasis_cv_empty_range() {
2648 let (data, t) = make_test_data(5, 50);
2649 let nbasis_range: Vec<usize> = vec![];
2650 let result = basis_nbasis_cv(
2651 &data,
2652 &t,
2653 &nbasis_range,
2654 &BasisType::Bspline { order: 4 },
2655 BasisCriterion::Gcv,
2656 5,
2657 1e-4,
2658 );
2659 assert!(result.is_none());
2660 }
2661
2662 #[test]
2663 fn test_basis_nbasis_cv_empty_data() {
2664 let data = FdMatrix::zeros(0, 50);
2665 let t = uniform_grid(50);
2666 let nbasis_range: Vec<usize> = vec![5, 10];
2667 let result = basis_nbasis_cv(
2668 &data,
2669 &t,
2670 &nbasis_range,
2671 &BasisType::Bspline { order: 4 },
2672 BasisCriterion::Gcv,
2673 5,
2674 1e-4,
2675 );
2676 assert!(result.is_none());
2677 }
2678
2679 #[test]
2680 fn test_basis_nbasis_cv_mismatched_argvals() {
2681 let data = FdMatrix::zeros(5, 50);
2682 let t = uniform_grid(40); let nbasis_range: Vec<usize> = vec![5, 10];
2684 let result = basis_nbasis_cv(
2685 &data,
2686 &t,
2687 &nbasis_range,
2688 &BasisType::Bspline { order: 4 },
2689 BasisCriterion::Gcv,
2690 5,
2691 1e-4,
2692 );
2693 assert!(result.is_none());
2694 }
2695
2696 #[test]
2697 fn test_basis_nbasis_cv_single_nbasis() {
2698 let (data, t) = make_test_data(5, 50);
2699 let nbasis_range: Vec<usize> = vec![10];
2700 let res = basis_nbasis_cv(
2701 &data,
2702 &t,
2703 &nbasis_range,
2704 &BasisType::Bspline { order: 4 },
2705 BasisCriterion::Gcv,
2706 5,
2707 1e-4,
2708 )
2709 .unwrap();
2710 assert_eq!(res.optimal_nbasis, 10);
2711 assert_eq!(res.scores.len(), 1);
2712 }
2713
2714 #[test]
2715 fn test_basis_nbasis_cv_bic_penalizes_more_than_aic() {
2716 let (data, t) = make_test_data(5, 80);
2719 let nbasis_range: Vec<usize> = (4..=20).step_by(2).collect();
2720
2721 let aic_res = basis_nbasis_cv(
2722 &data,
2723 &t,
2724 &nbasis_range,
2725 &BasisType::Bspline { order: 4 },
2726 BasisCriterion::Aic,
2727 5,
2728 1e-4,
2729 )
2730 .unwrap();
2731 let bic_res = basis_nbasis_cv(
2732 &data,
2733 &t,
2734 &nbasis_range,
2735 &BasisType::Bspline { order: 4 },
2736 BasisCriterion::Bic,
2737 5,
2738 1e-4,
2739 )
2740 .unwrap();
2741 assert!(
2744 bic_res.optimal_nbasis <= aic_res.optimal_nbasis + 4,
2745 "BIC selected {} vs AIC selected {} -- BIC should not select much more than AIC",
2746 bic_res.optimal_nbasis,
2747 aic_res.optimal_nbasis
2748 );
2749 }
2750
2751 #[test]
2754 fn test_smooth_basis_fitted_close_to_data() {
2755 let m = 50;
2757 let n = 3;
2758 let t = uniform_grid(m);
2759 let mut data = FdMatrix::zeros(n, m);
2760 for i in 0..n {
2761 for j in 0..m {
2762 data[(i, j)] = (2.0 * PI * t[j]).sin();
2763 }
2764 }
2765 let fdpar = make_bspline_fdpar(&t, 15, 1e-6);
2766 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2767
2768 let mut max_err = 0.0_f64;
2769 for i in 0..n {
2770 for j in 0..m {
2771 let err = (data[(i, j)] - res.fitted[(i, j)]).abs();
2772 max_err = max_err.max(err);
2773 }
2774 }
2775 assert!(
2776 max_err < 0.1,
2777 "Fitted should be close to smooth data; max_err={}",
2778 max_err
2779 );
2780 }
2781
2782 #[test]
2783 fn test_smooth_basis_constant_data() {
2784 let m = 50;
2786 let n = 2;
2787 let t = uniform_grid(m);
2788 let mut data = FdMatrix::zeros(n, m);
2789 for i in 0..n {
2790 for j in 0..m {
2791 data[(i, j)] = 3.15;
2792 }
2793 }
2794 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2795 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2796 for i in 0..n {
2797 for j in 0..m {
2798 assert!(
2799 (res.fitted[(i, j)] - 3.15).abs() < 0.01,
2800 "Constant data should be fit well at ({},{}): got {}",
2801 i,
2802 j,
2803 res.fitted[(i, j)]
2804 );
2805 }
2806 }
2807 }
2808
2809 #[test]
2810 fn test_smooth_basis_linear_data() {
2811 let m = 50;
2813 let t = uniform_grid(m);
2814 let mut data = FdMatrix::zeros(1, m);
2815 for j in 0..m {
2816 data[(0, j)] = 2.0 * t[j] + 1.0;
2817 }
2818 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2819 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2820 for j in 0..m {
2821 let expected = 2.0 * t[j] + 1.0;
2822 assert!(
2823 (res.fitted[(0, j)] - expected).abs() < 0.05,
2824 "Linear data should be fit well at j={}: got {}, expected {}",
2825 j,
2826 res.fitted[(0, j)],
2827 expected
2828 );
2829 }
2830 }
2831
2832 #[test]
2835 fn test_smooth_basis_edf_bounded() {
2836 let m = 50;
2837 let (data, t) = make_test_data(3, m);
2838 let fdpar = make_bspline_fdpar(&t, 12, 1e-4);
2839 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2840 assert!(
2842 res.edf > 0.0 && res.edf <= m as f64,
2843 "EDF should be in (0, {}]; got {}",
2844 m,
2845 res.edf
2846 );
2847 }
2848
2849 #[test]
2850 fn test_smooth_basis_gcv_aic_bic_all_finite() {
2851 let (data, t) = make_test_data(4, 60);
2852 let fdpar = make_bspline_fdpar(&t, 12, 1e-3);
2853 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2854 assert!(res.gcv.is_finite(), "GCV should be finite: {}", res.gcv);
2855 assert!(res.aic.is_finite(), "AIC should be finite: {}", res.aic);
2856 assert!(res.bic.is_finite(), "BIC should be finite: {}", res.bic);
2857 }
2858
2859 #[test]
2862 fn test_smooth_basis_penalty_matrix_in_result() {
2863 let (data, t) = make_test_data(3, 50);
2864 let nbasis = 10;
2865 let fdpar = make_bspline_fdpar(&t, nbasis, 1e-4);
2866 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2867 let k = res.nbasis;
2868 assert_eq!(
2869 res.penalty_matrix.len(),
2870 k * k,
2871 "Penalty matrix should be k*k = {}*{} = {}; got {}",
2872 k,
2873 k,
2874 k * k,
2875 res.penalty_matrix.len()
2876 );
2877 }
2878
2879 #[test]
2882 fn test_smooth_basis_identical_curves_same_coefficients() {
2883 let m = 50;
2884 let t = uniform_grid(m);
2885 let curve: Vec<f64> = (0..m).map(|j| (2.0 * PI * t[j]).sin()).collect();
2886 let n = 4;
2887 let mut data = FdMatrix::zeros(n, m);
2888 for i in 0..n {
2889 for j in 0..m {
2890 data[(i, j)] = curve[j];
2891 }
2892 }
2893 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2894 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2895
2896 let k = res.coefficients.ncols();
2898 for i in 1..n {
2899 for j in 0..k {
2900 assert!(
2901 (res.coefficients[(i, j)] - res.coefficients[(0, j)]).abs() < 1e-10,
2902 "Identical curves should have identical coefficients: curve {} col {} differs",
2903 i,
2904 j
2905 );
2906 }
2907 }
2908 }
2909
2910 #[test]
2913 fn test_basis_nbasis_cv_different_nfolds() {
2914 let (data, t) = make_test_data(12, 50);
2915 let nbasis_range: Vec<usize> = vec![5, 8, 11];
2916 for nfolds in [2, 3, 5, 10] {
2917 let res = basis_nbasis_cv(
2918 &data,
2919 &t,
2920 &nbasis_range,
2921 &BasisType::Bspline { order: 4 },
2922 BasisCriterion::Cv,
2923 nfolds,
2924 1e-4,
2925 );
2926 assert!(res.is_some(), "CV should succeed with nfolds={}", nfolds);
2927 let r = res.unwrap();
2928 assert!(nbasis_range.contains(&r.optimal_nbasis));
2929 }
2930 }
2931
2932 #[test]
2935 fn test_smooth_basis_many_basis_functions() {
2936 let m = 100;
2937 let (data, t) = make_test_data(2, m);
2938 let fdpar = make_bspline_fdpar(&t, 40, 1e-2);
2940 let res = smooth_basis(&data, &t, &fdpar);
2941 assert!(
2942 res.is_ok(),
2943 "Should handle many basis functions with penalty"
2944 );
2945 }
2946
2947 #[test]
2950 fn test_smooth_basis_bspline_vs_fourier_different_results() {
2951 let m = 50;
2952 let (data, t) = make_test_data(2, m);
2953 let fdpar_bs = make_bspline_fdpar(&t, 9, 1e-4);
2954 let fdpar_f = make_fourier_fdpar(9, 1.0, 1e-4);
2955 let res_bs = smooth_basis(&data, &t, &fdpar_bs).unwrap();
2956 let res_f = smooth_basis(&data, &t, &fdpar_f).unwrap();
2957 let diff: f64 = (0..m)
2959 .map(|j| (res_bs.fitted[(0, j)] - res_f.fitted[(0, j)]).abs())
2960 .sum();
2961 assert!(
2963 diff > 1e-10,
2964 "B-spline and Fourier fits should differ for the same data"
2965 );
2966 }
2967
2968 #[test]
2971 fn test_smooth_basis_gcv_positive_for_noisy_data() {
2972 let m = 50;
2973 let t = uniform_grid(m);
2974 let mut data = FdMatrix::zeros(1, m);
2975 for j in 0..m {
2976 data[(0, j)] = (2.0 * PI * t[j]).sin() + 0.5 * ((j * 37) % 20) as f64 / 20.0 - 0.25;
2978 }
2979 let fdpar = make_bspline_fdpar(&t, 10, 1e-3);
2980 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2981 assert!(res.gcv > 0.0, "GCV should be positive for noisy data");
2982 }
2983
2984 #[test]
2987 fn test_smooth_basis_different_lfd_orders() {
2988 let m = 50;
2989 let (data, t) = make_test_data(2, m);
2990
2991 let penalty1 = bspline_penalty_matrix(&t, 10, 4, 1);
2993 let fdpar1 = FdPar {
2994 basis_type: BasisType::Bspline { order: 4 },
2995 nbasis: 10,
2996 lambda: 1e-2,
2997 lfd_order: 1,
2998 penalty_matrix: penalty1,
2999 };
3000 let res1 = smooth_basis(&data, &t, &fdpar1);
3001 assert!(res1.is_ok());
3002
3003 let penalty2 = bspline_penalty_matrix(&t, 10, 4, 2);
3005 let fdpar2 = FdPar {
3006 basis_type: BasisType::Bspline { order: 4 },
3007 nbasis: 10,
3008 lambda: 1e-2,
3009 lfd_order: 2,
3010 penalty_matrix: penalty2,
3011 };
3012 let res2 = smooth_basis(&data, &t, &fdpar2);
3013 assert!(res2.is_ok());
3014
3015 let r1 = res1.unwrap();
3017 let r2 = res2.unwrap();
3018 let diff: f64 = (0..m)
3019 .map(|j| (r1.fitted[(0, j)] - r2.fitted[(0, j)]).abs())
3020 .sum();
3021 assert!(
3022 diff > 1e-10,
3023 "Different lfd_orders should produce different fits"
3024 );
3025 }
3026
3027 #[test]
3030 fn test_basis_nbasis_cv_result_fields() {
3031 let (data, t) = make_test_data(6, 50);
3032 let nbasis_range: Vec<usize> = vec![5, 7, 9, 11, 13];
3033 let res = basis_nbasis_cv(
3034 &data,
3035 &t,
3036 &nbasis_range,
3037 &BasisType::Bspline { order: 4 },
3038 BasisCriterion::Aic,
3039 5,
3040 1e-4,
3041 )
3042 .unwrap();
3043
3044 assert!(nbasis_range.contains(&res.optimal_nbasis));
3045 assert_eq!(res.scores.len(), nbasis_range.len());
3046 assert_eq!(res.nbasis_range, nbasis_range);
3047 assert_eq!(res.criterion, BasisCriterion::Aic);
3048 let min_score = res.scores.iter().copied().fold(f64::INFINITY, f64::min);
3050 let best_idx = res
3051 .scores
3052 .iter()
3053 .position(|&s| (s - min_score).abs() < 1e-15)
3054 .unwrap();
3055 assert_eq!(res.optimal_nbasis, nbasis_range[best_idx]);
3056 }
3057
3058 #[test]
3059 fn test_basis_nbasis_cv_result_clone() {
3060 let (data, t) = make_test_data(5, 50);
3061 let nbasis_range: Vec<usize> = vec![5, 10];
3062 let res = basis_nbasis_cv(
3063 &data,
3064 &t,
3065 &nbasis_range,
3066 &BasisType::Bspline { order: 4 },
3067 BasisCriterion::Gcv,
3068 5,
3069 1e-4,
3070 )
3071 .unwrap();
3072 let cloned = res.clone();
3073 assert_eq!(res, cloned);
3074 }
3075
3076 #[test]
3079 fn test_smooth_basis_nonuniform_argvals() {
3080 let m = 50;
3081 let t: Vec<f64> = (0..m)
3083 .map(|i| {
3084 let x = i as f64 / (m - 1) as f64;
3085 0.5 * (1.0 - (PI * x).cos())
3086 })
3087 .collect();
3088 let mut data = FdMatrix::zeros(2, m);
3089 for i in 0..2 {
3090 for j in 0..m {
3091 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.1 * i as f64;
3092 }
3093 }
3094 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
3095 let res = smooth_basis(&data, &t, &fdpar);
3096 assert!(res.is_ok(), "Should handle non-uniform argvals");
3097 let r = res.unwrap();
3098 assert_eq!(r.fitted.shape(), (2, m));
3099 }
3100
3101 #[test]
3104 fn test_smooth_basis_very_small_lambda() {
3105 let m = 50;
3106 let (data, t) = make_test_data(2, m);
3107 let fdpar = make_bspline_fdpar(&t, 10, 1e-15);
3108 let res = smooth_basis(&data, &t, &fdpar);
3109 assert!(res.is_ok(), "Should handle very small lambda");
3110 }
3111
3112 #[test]
3113 fn test_smooth_basis_very_large_lambda() {
3114 let m = 50;
3115 let (data, t) = make_test_data(2, m);
3116 let fdpar = make_bspline_fdpar(&t, 10, 1e10);
3117 let res = smooth_basis(&data, &t, &fdpar);
3118 assert!(res.is_ok(), "Should handle very large lambda");
3119 }
3120
3121 #[test]
3124 fn test_smooth_basis_multi_curve_vs_single_curve() {
3125 let m = 50;
3127 let n = 3;
3128 let (data, t) = make_test_data(n, m);
3129 let fdpar = make_bspline_fdpar(&t, 10, 1e-3);
3130
3131 let res_all = smooth_basis(&data, &t, &fdpar).unwrap();
3133
3134 for i in 0..n {
3136 let mut single = FdMatrix::zeros(1, m);
3137 for j in 0..m {
3138 single[(0, j)] = data[(i, j)];
3139 }
3140 let res_single = smooth_basis(&single, &t, &fdpar).unwrap();
3141 for j in 0..m {
3142 assert!(
3143 (res_all.fitted[(i, j)] - res_single.fitted[(0, j)]).abs() < 1e-10,
3144 "Multi-curve fit should match single-curve fit: curve {} point {}",
3145 i,
3146 j
3147 );
3148 }
3149 }
3150 }
3151
3152 #[test]
3155 fn test_basis_nbasis_cv_all_criteria_finite_scores() {
3156 let (data, t) = make_test_data(10, 60);
3157 let nbasis_range: Vec<usize> = vec![5, 7, 9, 11];
3158
3159 for criterion in [
3160 BasisCriterion::Gcv,
3161 BasisCriterion::Aic,
3162 BasisCriterion::Bic,
3163 BasisCriterion::Cv,
3164 ] {
3165 let res = basis_nbasis_cv(
3166 &data,
3167 &t,
3168 &nbasis_range,
3169 &BasisType::Bspline { order: 4 },
3170 criterion,
3171 5,
3172 1e-4,
3173 )
3174 .unwrap();
3175 let finite_count = res.scores.iter().filter(|s| s.is_finite()).count();
3177 assert!(
3178 finite_count > 0,
3179 "At least one score should be finite for {:?}",
3180 criterion
3181 );
3182 }
3183 }
3184
3185 #[test]
3188 fn test_smooth_basis_gcv_config_default() {
3189 let config = SmoothBasisGcvConfig::default();
3190 assert_eq!(config.basis_type, BasisType::Bspline { order: 4 });
3191 assert_eq!(config.nbasis, 15);
3192 assert_eq!(config.lfd_order, 2);
3193 assert_eq!(config.log_lambda_range, (-10.0, 2.0));
3194 assert_eq!(config.n_grid, 50);
3195 }
3196
3197 #[test]
3198 fn test_smooth_basis_gcv_config_clone_eq() {
3199 let config = SmoothBasisGcvConfig {
3200 nbasis: 20,
3201 ..SmoothBasisGcvConfig::default()
3202 };
3203 let cloned = config.clone();
3204 assert_eq!(config, cloned);
3205 }
3206
3207 #[test]
3208 fn test_smooth_basis_gcv_config_debug() {
3209 let config = SmoothBasisGcvConfig::default();
3210 let debug_str = format!("{:?}", config);
3211 assert!(debug_str.contains("SmoothBasisGcvConfig"));
3212 assert!(debug_str.contains("nbasis"));
3213 }
3214
3215 #[test]
3216 fn test_smooth_basis_gcv_config_partial_override() {
3217 let config = SmoothBasisGcvConfig {
3218 basis_type: BasisType::Fourier { period: 2.0 },
3219 n_grid: 100,
3220 ..SmoothBasisGcvConfig::default()
3221 };
3222 assert_eq!(config.basis_type, BasisType::Fourier { period: 2.0 });
3223 assert_eq!(config.n_grid, 100);
3224 assert_eq!(config.nbasis, 15);
3226 assert_eq!(config.lfd_order, 2);
3227 }
3228
3229 #[test]
3230 fn test_smooth_basis_gcv_with_config_default() {
3231 let (data, t) = make_test_data(5, 101);
3232 let config = SmoothBasisGcvConfig::default();
3233 let result = smooth_basis_gcv_with_config(&data, &t, &config);
3234 assert!(result.is_ok(), "GCV with default config should succeed");
3235 let res = result.unwrap();
3236 assert_eq!(res.fitted.shape(), (5, 101));
3237 assert!(res.edf > 0.0);
3238 assert!(res.gcv.is_finite());
3239 }
3240
3241 #[test]
3242 fn test_smooth_basis_gcv_with_config_custom() {
3243 let (data, t) = make_test_data(3, 50);
3244 let config = SmoothBasisGcvConfig {
3245 nbasis: 10,
3246 log_lambda_range: (-6.0, 0.0),
3247 n_grid: 15,
3248 ..SmoothBasisGcvConfig::default()
3249 };
3250 let result = smooth_basis_gcv_with_config(&data, &t, &config);
3251 assert!(result.is_ok());
3252 }
3253
3254 #[test]
3255 fn test_smooth_basis_gcv_with_config_matches_direct() {
3256 let (data, t) = make_test_data(3, 50);
3257 let config = SmoothBasisGcvConfig {
3258 nbasis: 10,
3259 log_lambda_range: (-6.0, 0.0),
3260 n_grid: 20,
3261 ..SmoothBasisGcvConfig::default()
3262 };
3263 let with_config = smooth_basis_gcv_with_config(&data, &t, &config).unwrap();
3264 let direct = smooth_basis_gcv(
3265 &data,
3266 &t,
3267 &config.basis_type,
3268 config.nbasis,
3269 config.lfd_order,
3270 config.log_lambda_range,
3271 config.n_grid,
3272 )
3273 .unwrap();
3274 assert_eq!(with_config.gcv, direct.gcv);
3275 assert_eq!(with_config.edf, direct.edf);
3276 assert_eq!(with_config.nbasis, direct.nbasis);
3277 }
3278
3279 #[test]
3280 fn test_smooth_basis_gcv_with_config_fourier() {
3281 let m = 100;
3282 let t = uniform_grid(m);
3283 let mut data = FdMatrix::zeros(2, m);
3284 for i in 0..2 {
3285 for j in 0..m {
3286 data[(i, j)] = (2.0 * PI * t[j]).sin() + (4.0 * PI * t[j]).cos();
3287 }
3288 }
3289 let config = SmoothBasisGcvConfig {
3290 basis_type: BasisType::Fourier { period: 1.0 },
3291 nbasis: 7,
3292 n_grid: 20,
3293 ..SmoothBasisGcvConfig::default()
3294 };
3295 let result = smooth_basis_gcv_with_config(&data, &t, &config);
3296 assert!(result.is_ok());
3297 }
3298
3299 #[test]
3302 fn test_basis_nbasis_cv_config_default() {
3303 let config = BasisNbasisCvConfig::default();
3304 assert_eq!(config.basis_type, BasisType::Bspline { order: 4 });
3305 assert_eq!(config.nbasis_range, (5, 30));
3306 assert!((config.lambda - 1e-4).abs() < 1e-15);
3307 assert_eq!(config.lfd_order, 2);
3308 assert_eq!(config.n_folds, 5);
3309 assert_eq!(config.criterion, BasisCriterion::Gcv);
3310 }
3311
3312 #[test]
3313 fn test_basis_nbasis_cv_config_clone_eq() {
3314 let config = BasisNbasisCvConfig {
3315 nbasis_range: (4, 15),
3316 ..BasisNbasisCvConfig::default()
3317 };
3318 let cloned = config.clone();
3319 assert_eq!(config, cloned);
3320 }
3321
3322 #[test]
3323 fn test_basis_nbasis_cv_config_debug() {
3324 let config = BasisNbasisCvConfig::default();
3325 let debug_str = format!("{:?}", config);
3326 assert!(debug_str.contains("BasisNbasisCvConfig"));
3327 assert!(debug_str.contains("nbasis_range"));
3328 }
3329
3330 #[test]
3331 fn test_basis_nbasis_cv_config_partial_override() {
3332 let config = BasisNbasisCvConfig {
3333 criterion: BasisCriterion::Aic,
3334 lambda: 1e-2,
3335 ..BasisNbasisCvConfig::default()
3336 };
3337 assert_eq!(config.criterion, BasisCriterion::Aic);
3338 assert!((config.lambda - 1e-2).abs() < 1e-15);
3339 assert_eq!(config.nbasis_range, (5, 30));
3341 assert_eq!(config.n_folds, 5);
3342 }
3343
3344 #[test]
3345 fn test_basis_nbasis_cv_with_config_default() {
3346 let (data, t) = make_test_data(5, 51);
3347 let config = BasisNbasisCvConfig {
3348 nbasis_range: (5, 12),
3349 ..BasisNbasisCvConfig::default()
3350 };
3351 let result = basis_nbasis_cv_with_config(&data, &t, &config);
3352 assert!(
3353 result.is_ok(),
3354 "nbasis CV with default config should succeed"
3355 );
3356 let res = result.unwrap();
3357 assert!(res.optimal_nbasis >= 5 && res.optimal_nbasis <= 12);
3358 assert_eq!(res.scores.len(), 8); assert_eq!(res.criterion, BasisCriterion::Gcv);
3360 }
3361
3362 #[test]
3363 fn test_basis_nbasis_cv_with_config_aic() {
3364 let (data, t) = make_test_data(5, 51);
3365 let config = BasisNbasisCvConfig {
3366 nbasis_range: (5, 10),
3367 criterion: BasisCriterion::Aic,
3368 ..BasisNbasisCvConfig::default()
3369 };
3370 let result = basis_nbasis_cv_with_config(&data, &t, &config);
3371 assert!(result.is_ok());
3372 assert_eq!(result.unwrap().criterion, BasisCriterion::Aic);
3373 }
3374
3375 #[test]
3376 fn test_basis_nbasis_cv_with_config_cv_folds() {
3377 let (data, t) = make_test_data(10, 51);
3378 let config = BasisNbasisCvConfig {
3379 nbasis_range: (5, 9),
3380 criterion: BasisCriterion::Cv,
3381 n_folds: 3,
3382 ..BasisNbasisCvConfig::default()
3383 };
3384 let result = basis_nbasis_cv_with_config(&data, &t, &config);
3385 assert!(result.is_ok());
3386 assert_eq!(result.unwrap().criterion, BasisCriterion::Cv);
3387 }
3388
3389 #[test]
3390 fn test_basis_nbasis_cv_with_config_matches_direct() {
3391 let (data, t) = make_test_data(5, 51);
3392 let config = BasisNbasisCvConfig {
3393 nbasis_range: (5, 10),
3394 criterion: BasisCriterion::Bic,
3395 lambda: 1e-3,
3396 ..BasisNbasisCvConfig::default()
3397 };
3398 let with_config = basis_nbasis_cv_with_config(&data, &t, &config).unwrap();
3399 let nbasis_range: Vec<usize> = (5..=10).collect();
3400 let direct = basis_nbasis_cv(
3401 &data,
3402 &t,
3403 &nbasis_range,
3404 &config.basis_type,
3405 config.criterion,
3406 config.n_folds,
3407 config.lambda,
3408 )
3409 .unwrap();
3410 assert_eq!(with_config.optimal_nbasis, direct.optimal_nbasis);
3411 assert_eq!(with_config.scores, direct.scores);
3412 assert_eq!(with_config.nbasis_range, direct.nbasis_range);
3413 }
3414
3415 #[test]
3416 fn test_basis_nbasis_cv_with_config_nbasis_range_expansion() {
3417 let (data, t) = make_test_data(5, 51);
3418 let config = BasisNbasisCvConfig {
3419 nbasis_range: (7, 7), ..BasisNbasisCvConfig::default()
3421 };
3422 let result = basis_nbasis_cv_with_config(&data, &t, &config);
3423 assert!(result.is_ok());
3424 let res = result.unwrap();
3425 assert_eq!(res.optimal_nbasis, 7);
3426 assert_eq!(res.scores.len(), 1);
3427 }
3428
3429 fn make_positive_fdpar(argvals: &[f64]) -> FdPar {
3433 let nbasis = 10;
3434 let penalty = bspline_penalty_matrix(argvals, nbasis, 4, 2);
3435 FdPar {
3436 basis_type: BasisType::Bspline { order: 4 },
3437 nbasis,
3438 lambda: 1e-3,
3439 lfd_order: 2,
3440 penalty_matrix: penalty,
3441 }
3442 }
3443
3444 #[test]
3445 fn test_smooth_positive_is_positive() {
3446 let m = 41;
3448 let t = uniform_grid(m);
3449 let mut data = FdMatrix::zeros(1, m);
3450 for j in 0..m {
3451 let wiggle = 0.05 * ((j * 7) % 13) as f64 / 13.0;
3453 data[(0, j)] = 2.0 + (2.0 * PI * t[j]).sin() + wiggle;
3454 }
3455
3456 let fdpar = make_positive_fdpar(&t);
3457 let result = smooth_positive(&data, &t, &fdpar);
3458 assert!(
3459 result.is_ok(),
3460 "smooth_positive should succeed on positive data"
3461 );
3462
3463 let res = result.unwrap();
3464 assert_eq!(res.fitted.shape(), (1, m));
3465 assert_eq!(res.log_coefficients.nrows(), 1);
3466
3467 for j in 0..m {
3468 let v = res.fitted[(0, j)];
3469 assert!(v > 0.0, "fitted value at j={} is not positive: {}", j, v);
3470 assert!(
3471 v.is_finite(),
3472 "fitted value at j={} is not finite: {}",
3473 j,
3474 v
3475 );
3476 }
3477 }
3478
3479 #[test]
3480 fn test_smooth_positive_recovers_curve() {
3481 let m = 41;
3483 let t = uniform_grid(m);
3484 let mut data = FdMatrix::zeros(1, m);
3485 let mut truth = vec![0.0f64; m];
3486 for j in 0..m {
3487 truth[j] = 2.0 + (2.0 * PI * t[j]).sin();
3488 data[(0, j)] = truth[j]; }
3490
3491 let nbasis = 10;
3492 let penalty = bspline_penalty_matrix(&t, nbasis, 4, 2);
3493 let fdpar = FdPar {
3494 basis_type: BasisType::Bspline { order: 4 },
3495 nbasis,
3496 lambda: 1e-6, lfd_order: 2,
3498 penalty_matrix: penalty,
3499 };
3500
3501 let res = smooth_positive(&data, &t, &fdpar).unwrap();
3502 let mae: f64 = (0..m)
3503 .map(|j| (res.fitted[(0, j)] - truth[j]).abs())
3504 .sum::<f64>()
3505 / m as f64;
3506
3507 assert!(
3508 mae < 0.2,
3509 "Mean absolute error too large: {}; smooth_positive should recover positive curve",
3510 mae
3511 );
3512 assert!(res.edf > 0.0, "EDF should be positive");
3514 assert!(res.gcv.is_finite(), "GCV should be finite");
3515 }
3516
3517 #[test]
3518 fn test_smooth_positive_rejects_nonpositive() {
3519 let m = 41;
3521 let t = uniform_grid(m);
3522 let mut data = FdMatrix::zeros(1, m);
3523 for j in 0..m {
3524 data[(0, j)] = 2.0 + (2.0 * PI * t[j]).sin();
3525 }
3526 data[(0, 10)] = 0.0;
3528
3529 let fdpar = make_positive_fdpar(&t);
3530 let result = smooth_positive(&data, &t, &fdpar);
3531 assert!(
3532 result.is_err(),
3533 "smooth_positive must reject data with a zero element"
3534 );
3535
3536 match result.unwrap_err() {
3537 crate::FdarError::InvalidParameter { parameter, .. } => {
3538 assert_eq!(parameter, "data");
3539 }
3540 other => panic!("Expected InvalidParameter, got {:?}", other),
3541 }
3542 }
3543
3544 #[test]
3545 fn test_smooth_positive_rejects_negative() {
3546 let m = 41;
3548 let t = uniform_grid(m);
3549 let mut data = FdMatrix::zeros(1, m);
3550 for j in 0..m {
3551 data[(0, j)] = 2.0 + (2.0 * PI * t[j]).sin();
3552 }
3553 data[(0, 20)] = -0.5;
3554
3555 let fdpar = make_positive_fdpar(&t);
3556 let result = smooth_positive(&data, &t, &fdpar);
3557 assert!(
3558 result.is_err(),
3559 "smooth_positive must reject data with a negative element"
3560 );
3561
3562 match result.unwrap_err() {
3563 crate::FdarError::InvalidParameter { parameter, .. } => {
3564 assert_eq!(parameter, "data");
3565 }
3566 other => panic!("Expected InvalidParameter, got {:?}", other),
3567 }
3568 }
3569
3570 #[test]
3573 fn test_smooth_monotone_is_monotone() {
3574 let m = 41_usize;
3576 let t = uniform_grid(m); let y: Vec<f64> = t.iter().map(|&ti| ti * ti).collect();
3578
3579 let result = smooth_monotone(&y, &t, 8, 4, 1e-3, 50)
3580 .expect("smooth_monotone should succeed on t² data");
3581
3582 for (i, &v) in result.fitted.iter().enumerate() {
3584 assert!(v.is_finite(), "fitted[{}] = {} is not finite", i, v);
3585 }
3586
3587 for i in 1..m {
3589 assert!(
3590 result.fitted[i] >= result.fitted[i - 1] - 1e-9,
3591 "Monotonicity violated at i={}: fitted[{}]={} < fitted[{}]={}",
3592 i,
3593 i,
3594 result.fitted[i],
3595 i - 1,
3596 result.fitted[i - 1]
3597 );
3598 }
3599
3600 assert!(
3602 result.beta1 > 0.0,
3603 "beta1 should be positive for increasing data, got {}",
3604 result.beta1
3605 );
3606 }
3607
3608 #[test]
3611 fn test_smooth_monotone_recovers_increasing() {
3612 let m = 51_usize;
3618 let t = uniform_grid(m);
3619 let y: Vec<f64> = t
3620 .iter()
3621 .map(|&ti| 1.0 / (1.0 + (-8.0 * (ti - 0.5)).exp()))
3622 .collect();
3623
3624 let result = smooth_monotone(&y, &t, 10, 4, 1e-4, 100)
3625 .expect("smooth_monotone should succeed on logistic data");
3626
3627 let mae: f64 = y
3631 .iter()
3632 .zip(result.fitted.iter())
3633 .map(|(&yi, &fi)| (yi - fi).abs())
3634 .sum::<f64>()
3635 / m as f64;
3636 assert!(
3637 mae < 0.15,
3638 "Mean absolute error {} too large for logistic recovery (tolerance 0.15, iterations={})",
3639 mae,
3640 result.iterations
3641 );
3642
3643 for i in 1..m {
3645 assert!(
3646 result.fitted[i] >= result.fitted[i - 1] - 1e-9,
3647 "Monotonicity violated at i={}: fitted[{}]={} < fitted[{}]={}",
3648 i,
3649 result.fitted[i],
3650 i,
3651 result.fitted[i - 1],
3652 i - 1
3653 );
3654 }
3655 }
3656
3657 #[test]
3658 fn test_smooth_monotone_decreasing() {
3659 let m = 41_usize;
3661 let t = uniform_grid(m);
3662 let y: Vec<f64> = t.iter().map(|&ti| 1.0 - ti).collect();
3663
3664 let result = smooth_monotone(&y, &t, 8, 4, 1e-3, 50)
3665 .expect("smooth_monotone should succeed on decreasing data");
3666
3667 assert!(
3669 result.beta1 < 0.0,
3670 "beta1 should be negative for decreasing data, got {}",
3671 result.beta1
3672 );
3673
3674 for i in 1..m {
3676 assert!(
3677 result.fitted[i] <= result.fitted[i - 1] + 1e-9,
3678 "Nonincreasing violated at i={}: fitted[{}]={} > fitted[{}]={}",
3679 i,
3680 result.fitted[i],
3681 i,
3682 result.fitted[i - 1],
3683 i - 1
3684 );
3685 }
3686 }
3687
3688 #[test]
3689 fn test_smooth_monotone_bounded_iterations() {
3690 let m = 51_usize;
3692 let t = uniform_grid(m);
3693 let y: Vec<f64> = t
3694 .iter()
3695 .enumerate()
3696 .map(|(i, &ti)| {
3697 let noise = 0.05 * ((i * 17 + 3) % 11) as f64 / 10.0 - 0.025;
3698 ti + noise })
3700 .collect();
3701
3702 let result = smooth_monotone(&y, &t, 8, 4, 1e-2, 50)
3703 .expect("smooth_monotone should succeed on noisy data");
3704
3705 assert!(
3707 result.iterations <= 50,
3708 "iterations ({}) must be <= max_iter (50)",
3709 result.iterations
3710 );
3711
3712 if result.beta1 >= 0.0 {
3716 for i in 1..m {
3717 assert!(
3718 result.fitted[i] >= result.fitted[i - 1] - 1e-9,
3719 "Nondecreasing violated (beta1={}) at i={}",
3720 result.beta1,
3721 i
3722 );
3723 }
3724 } else {
3725 for i in 1..m {
3726 assert!(
3727 result.fitted[i] <= result.fitted[i - 1] + 1e-9,
3728 "Nonincreasing violated (beta1={}) at i={}",
3729 result.beta1,
3730 i
3731 );
3732 }
3733 }
3734 }
3735
3736 #[test]
3739 fn test_smooth_monotone_errors_on_short_input() {
3740 let data = vec![0.0, 1.0];
3742 let argvals = vec![0.0, 1.0];
3743 let result = smooth_monotone(&data, &argvals, 4, 4, 1e-3, 50);
3744 assert!(
3745 result.is_err(),
3746 "smooth_monotone should fail for data.len() == 2"
3747 );
3748 match result.unwrap_err() {
3749 crate::FdarError::InvalidDimension { parameter, .. } => {
3750 assert!(
3751 parameter.contains("data") || parameter.contains("argvals"),
3752 "Expected data/argvals dimension error, got param={}",
3753 parameter
3754 );
3755 }
3756 other => panic!("Expected InvalidDimension, got {:?}", other),
3757 }
3758 }
3759
3760 #[test]
3761 fn test_smooth_monotone_errors_on_argvals_mismatch() {
3762 let m = 10_usize;
3764 let data: Vec<f64> = (0..m).map(|i| i as f64).collect();
3765 let argvals: Vec<f64> = (0..m - 1).map(|i| i as f64).collect();
3766 let result = smooth_monotone(&data, &argvals, 4, 4, 1e-3, 50);
3767 assert!(
3768 result.is_err(),
3769 "smooth_monotone should fail for argvals length mismatch"
3770 );
3771 match result.unwrap_err() {
3772 crate::FdarError::InvalidDimension { .. } => {}
3773 other => panic!("Expected InvalidDimension, got {:?}", other),
3774 }
3775 }
3776
3777 #[test]
3778 fn test_smooth_monotone_errors_on_bad_params() {
3779 let m = 10_usize;
3780 let t = uniform_grid(m);
3781 let data: Vec<f64> = t.clone();
3782
3783 let result = smooth_monotone(&data, &t, 1, 4, 1e-3, 50);
3785 assert!(
3786 result.is_err(),
3787 "smooth_monotone should fail for nbasis == 1"
3788 );
3789 match result.unwrap_err() {
3790 crate::FdarError::InvalidParameter { parameter, .. } => {
3791 assert_eq!(parameter, "nbasis");
3792 }
3793 other => panic!("Expected InvalidParameter(nbasis), got {:?}", other),
3794 }
3795
3796 let result = smooth_monotone(&data, &t, 8, 4, 1e-3, 0);
3798 assert!(
3799 result.is_err(),
3800 "smooth_monotone should fail for max_iter == 0"
3801 );
3802 match result.unwrap_err() {
3803 crate::FdarError::InvalidParameter { parameter, .. } => {
3804 assert_eq!(parameter, "max_iter");
3805 }
3806 other => panic!("Expected InvalidParameter(max_iter), got {:?}", other),
3807 }
3808 }
3809}