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