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)]
46#[non_exhaustive]
47pub struct SmoothBasisResult {
48 pub coefficients: FdMatrix,
50 pub fitted: FdMatrix,
52 pub edf: f64,
54 pub gcv: f64,
56 pub aic: f64,
58 pub bic: f64,
60 pub penalty_matrix: Vec<f64>,
62 pub nbasis: usize,
64}
65
66pub fn bspline_penalty_matrix(
83 argvals: &[f64],
84 nbasis: usize,
85 order: usize,
86 lfd_order: usize,
87) -> Vec<f64> {
88 if nbasis < 2 || order < 1 || lfd_order >= order || argvals.len() < 2 {
89 return vec![0.0; nbasis * nbasis];
90 }
91
92 let nknots = nbasis.saturating_sub(order).max(2);
93
94 let n_sub = 10;
96 let t_min = argvals[0];
97 let t_max = argvals[argvals.len() - 1];
98 let n_quad = (argvals.len() - 1) * n_sub + 1;
99 let quad_t: Vec<f64> = (0..n_quad)
100 .map(|i| t_min + (t_max - t_min) * i as f64 / (n_quad - 1) as f64)
101 .collect();
102
103 let basis_fine = bspline_basis(&quad_t, nknots, order);
105 let actual_nbasis = basis_fine.len() / n_quad;
106
107 let h = (t_max - t_min) / (n_quad - 1) as f64;
109 let deriv_basis = differentiate_basis_columns(&basis_fine, n_quad, actual_nbasis, h, lfd_order);
110
111 let weights = simpsons_weights(&quad_t);
113
114 integrate_symmetric_penalty(&deriv_basis, &weights, actual_nbasis, n_quad)
116}
117
118pub fn fourier_penalty_matrix(nbasis: usize, period: f64, lfd_order: usize) -> Vec<f64> {
130 let k = nbasis;
131 let mut penalty = vec![0.0; k * k];
132
133 let mut freq = 1;
139 let mut idx = 1;
140 while idx < k {
141 let omega = 2.0 * PI * f64::from(freq) / period;
142 let eigenval = omega.powi(2 * lfd_order as i32);
143
144 if idx < k {
146 penalty[idx + idx * k] = eigenval;
147 idx += 1;
148 }
149 if idx < k {
151 penalty[idx + idx * k] = eigenval;
152 idx += 1;
153 }
154 freq += 1;
155 }
156
157 penalty
158}
159
160pub fn smooth_basis(
175 data: &FdMatrix,
176 argvals: &[f64],
177 fdpar: &FdPar,
178) -> Result<SmoothBasisResult, crate::FdarError> {
179 let (n, m) = data.shape();
180 if n == 0 || m == 0 || argvals.len() != m || fdpar.nbasis < 2 {
181 return Err(crate::FdarError::InvalidDimension {
182 parameter: "data/argvals/fdpar",
183 expected: "n > 0, m > 0, argvals.len() == m, nbasis >= 2".to_string(),
184 actual: format!(
185 "n={}, m={}, argvals.len()={}, nbasis={}",
186 n,
187 m,
188 argvals.len(),
189 fdpar.nbasis
190 ),
191 });
192 }
193
194 let (basis_flat, actual_nbasis) = evaluate_basis(argvals, &fdpar.basis_type, fdpar.nbasis);
196 let k = actual_nbasis;
197
198 let b_mat = DMatrix::from_column_slice(m, k, &basis_flat);
199 let r_mat = DMatrix::from_column_slice(k, k, &fdpar.penalty_matrix);
200
201 let btb = b_mat.transpose() * &b_mat;
203 let ridge_eps = 1e-10;
204 let system: DMatrix<f64> =
205 &btb + fdpar.lambda * &r_mat + ridge_eps * DMatrix::<f64>::identity(k, k);
206
207 let system_inv =
209 invert_penalized_system(&system, k).ok_or_else(|| crate::FdarError::ComputationFailed {
210 operation: "matrix inversion",
211 detail: "failed to invert penalized system (Φ'Φ + λR); try increasing lambda or reducing the number of basis functions".to_string(),
212 })?;
213
214 let h_mat = &b_mat * &system_inv * b_mat.transpose();
216 let edf: f64 = (0..m).map(|i| h_mat[(i, i)]).sum();
217
218 let proj = &system_inv * b_mat.transpose();
220 let (all_coefs, all_fitted, total_rss) = project_all_curves(data, &b_mat, &proj, n, m, k);
221
222 let total_points = (n * m) as f64;
223 let gcv = compute_gcv(total_rss, total_points, edf, m);
224 let mse = total_rss / total_points;
225 let total_edf = n as f64 * edf;
227 let aic = total_points * mse.max(1e-300).ln() + 2.0 * total_edf;
228 let bic = total_points * mse.max(1e-300).ln() + total_points.ln() * total_edf;
229
230 Ok(SmoothBasisResult {
231 coefficients: all_coefs,
232 fitted: all_fitted,
233 edf,
234 gcv,
235 aic,
236 bic,
237 penalty_matrix: fdpar.penalty_matrix.clone(),
238 nbasis: k,
239 })
240}
241
242pub fn smooth_basis_gcv(
255 data: &FdMatrix,
256 argvals: &[f64],
257 basis_type: &BasisType,
258 nbasis: usize,
259 lfd_order: usize,
260 log_lambda_range: (f64, f64),
261 n_grid: usize,
262) -> Option<SmoothBasisResult> {
263 let m = argvals.len();
264 if m == 0 || nbasis < 2 || n_grid < 2 {
265 return None;
266 }
267
268 let penalty = match basis_type {
270 BasisType::Bspline { order } => bspline_penalty_matrix(argvals, nbasis, *order, lfd_order),
271 BasisType::Fourier { period } => fourier_penalty_matrix(nbasis, *period, lfd_order),
272 };
273
274 let (lo, hi) = log_lambda_range;
275 let mut best_gcv = f64::INFINITY;
276 let mut best_result: Option<SmoothBasisResult> = None;
277
278 for i in 0..n_grid {
279 let log_lam = lo + (hi - lo) * i as f64 / (n_grid - 1) as f64;
280 let lam = 10.0_f64.powf(log_lam);
281
282 let fdpar = FdPar {
283 basis_type: basis_type.clone(),
284 nbasis,
285 lambda: lam,
286 lfd_order,
287 penalty_matrix: penalty.clone(),
288 };
289
290 if let Ok(result) = smooth_basis(data, argvals, &fdpar) {
291 if result.gcv < best_gcv {
292 best_gcv = result.gcv;
293 best_result = Some(result);
294 }
295 }
296 }
297
298 best_result
299}
300
301pub fn smooth_basis_aic(
338 data: &FdMatrix,
339 argvals: &[f64],
340 basis_type: &BasisType,
341 nbasis: usize,
342 lfd_order: usize,
343 log_lambda_range: (f64, f64),
344 n_grid: usize,
345) -> Option<SmoothBasisResult> {
346 let m = argvals.len();
347 if m == 0 || nbasis < 2 || n_grid < 2 {
348 return None;
349 }
350
351 let penalty = match basis_type {
353 BasisType::Bspline { order } => bspline_penalty_matrix(argvals, nbasis, *order, lfd_order),
354 BasisType::Fourier { period } => fourier_penalty_matrix(nbasis, *period, lfd_order),
355 };
356
357 let (lo, hi) = log_lambda_range;
358 let mut best_aic = f64::INFINITY;
359 let mut best_result: Option<SmoothBasisResult> = None;
360
361 for i in 0..n_grid {
362 let log_lam = lo + (hi - lo) * i as f64 / (n_grid - 1) as f64;
363 let lam = 10.0_f64.powf(log_lam);
364
365 let fdpar = FdPar {
366 basis_type: basis_type.clone(),
367 nbasis,
368 lambda: lam,
369 lfd_order,
370 penalty_matrix: penalty.clone(),
371 };
372
373 if let Ok(result) = smooth_basis(data, argvals, &fdpar) {
374 if result.aic < best_aic {
377 best_aic = result.aic;
378 best_result = Some(result);
379 }
380 }
381 }
382
383 best_result
384}
385
386#[derive(Debug, Clone, PartialEq)]
404pub struct SmoothBasisGcvConfig {
405 pub basis_type: BasisType,
407 pub nbasis: usize,
409 pub lfd_order: usize,
411 pub log_lambda_range: (f64, f64),
413 pub n_grid: usize,
415}
416
417impl Default for SmoothBasisGcvConfig {
418 fn default() -> Self {
419 Self {
420 basis_type: BasisType::Bspline { order: 4 },
421 nbasis: 15,
422 lfd_order: 2,
423 log_lambda_range: (-10.0, 2.0),
424 n_grid: 50,
425 }
426 }
427}
428
429#[must_use = "expensive computation whose result should not be discarded"]
444pub fn smooth_basis_gcv_with_config(
445 data: &FdMatrix,
446 argvals: &[f64],
447 config: &SmoothBasisGcvConfig,
448) -> Result<SmoothBasisResult, crate::FdarError> {
449 smooth_basis_gcv(
450 data,
451 argvals,
452 &config.basis_type,
453 config.nbasis,
454 config.lfd_order,
455 config.log_lambda_range,
456 config.n_grid,
457 )
458 .ok_or_else(|| crate::FdarError::ComputationFailed {
459 operation: "smooth_basis_gcv_with_config",
460 detail: "no valid smoothing result found in GCV lambda search".to_string(),
461 })
462}
463
464#[derive(Debug, Clone, PartialEq)]
480pub struct BasisNbasisCvConfig {
481 pub basis_type: BasisType,
483 pub nbasis_range: (usize, usize),
485 pub lambda: f64,
487 pub lfd_order: usize,
489 pub n_folds: usize,
491 pub criterion: BasisCriterion,
493}
494
495impl Default for BasisNbasisCvConfig {
496 fn default() -> Self {
497 Self {
498 basis_type: BasisType::Bspline { order: 4 },
499 nbasis_range: (5, 30),
500 lambda: 1e-4,
501 lfd_order: 2,
502 n_folds: 5,
503 criterion: BasisCriterion::Gcv,
504 }
505 }
506}
507
508#[must_use = "expensive computation whose result should not be discarded"]
526pub fn basis_nbasis_cv_with_config(
527 data: &FdMatrix,
528 argvals: &[f64],
529 config: &BasisNbasisCvConfig,
530) -> Result<BasisNbasisCvResult, crate::FdarError> {
531 let nbasis_range: Vec<usize> = (config.nbasis_range.0..=config.nbasis_range.1).collect();
532 basis_nbasis_cv(
533 data,
534 argvals,
535 &nbasis_range,
536 &config.basis_type,
537 config.criterion,
538 config.n_folds,
539 config.lambda,
540 )
541 .ok_or_else(|| crate::FdarError::ComputationFailed {
542 operation: "basis_nbasis_cv_with_config",
543 detail: "no valid result found in nbasis CV search".to_string(),
544 })
545}
546
547fn differentiate_basis_columns(
551 basis: &[f64],
552 n_quad: usize,
553 nbasis: usize,
554 h: f64,
555 lfd_order: usize,
556) -> Vec<f64> {
557 let mut deriv = basis.to_vec();
558 for _ in 0..lfd_order {
559 let mut new_deriv = vec![0.0; n_quad * nbasis];
560 for j in 0..nbasis {
561 let col: Vec<f64> = (0..n_quad).map(|i| deriv[i + j * n_quad]).collect();
562 let grad = crate::helpers::gradient_uniform(&col, h);
563 for i in 0..n_quad {
564 new_deriv[i + j * n_quad] = grad[i];
565 }
566 }
567 deriv = new_deriv;
568 }
569 deriv
570}
571
572fn integrate_symmetric_penalty(
574 deriv_basis: &[f64],
575 weights: &[f64],
576 k: usize,
577 n_quad: usize,
578) -> Vec<f64> {
579 let mut penalty = vec![0.0; k * k];
580 for j in 0..k {
581 for l in j..k {
582 let mut val = 0.0;
583 for i in 0..n_quad {
584 val += deriv_basis[i + j * n_quad] * deriv_basis[i + l * n_quad] * weights[i];
585 }
586 penalty[j + l * k] = val;
587 penalty[l + j * k] = val;
588 }
589 }
590 penalty
591}
592
593fn evaluate_basis(argvals: &[f64], basis_type: &BasisType, nbasis: usize) -> (Vec<f64>, usize) {
595 let m = argvals.len();
596 match basis_type {
597 BasisType::Bspline { order } => {
598 let nknots = nbasis.saturating_sub(*order).max(2);
599 let basis = bspline_basis(argvals, nknots, *order);
600 let actual = basis.len() / m;
601 (basis, actual)
602 }
603 BasisType::Fourier { period } => {
604 let basis = fourier_basis_with_period(argvals, nbasis, *period);
605 (basis, nbasis)
606 }
607 }
608}
609
610fn invert_penalized_system(system: &DMatrix<f64>, k: usize) -> Option<DMatrix<f64>> {
612 if let Some(chol) = system.clone().cholesky() {
613 return Some(chol.inverse());
614 }
615 let svd = nalgebra::SVD::new(system.clone(), true, true);
617 let u = svd.u.as_ref()?;
618 let v_t = svd.v_t.as_ref()?;
619 let max_sv: f64 = svd.singular_values.iter().copied().fold(0.0_f64, f64::max);
620 let eps = 1e-10 * max_sv;
621 let mut inv = DMatrix::<f64>::zeros(k, k);
622 for ii in 0..k {
623 for jj in 0..k {
624 let mut sum = 0.0;
625 for s in 0..k.min(svd.singular_values.len()) {
626 if svd.singular_values[s] > eps {
627 sum += v_t[(s, ii)] / svd.singular_values[s] * u[(jj, s)];
628 }
629 }
630 inv[(ii, jj)] = sum;
631 }
632 }
633 Some(inv)
634}
635
636fn project_all_curves(
638 data: &FdMatrix,
639 b_mat: &DMatrix<f64>,
640 proj: &DMatrix<f64>,
641 n: usize,
642 m: usize,
643 k: usize,
644) -> (FdMatrix, FdMatrix, f64) {
645 let mut all_coefs = FdMatrix::zeros(n, k);
646 let mut all_fitted = FdMatrix::zeros(n, m);
647 let mut total_rss = 0.0;
648
649 for i in 0..n {
650 let curve: Vec<f64> = (0..m).map(|j| data[(i, j)]).collect();
651 let y_vec = nalgebra::DVector::from_vec(curve.clone());
652 let coefs = proj * &y_vec;
653
654 for j in 0..k {
655 all_coefs[(i, j)] = coefs[j];
656 }
657 let fitted = b_mat * &coefs;
658 for j in 0..m {
659 all_fitted[(i, j)] = fitted[j];
660 let resid = curve[j] - fitted[j];
661 total_rss += resid * resid;
662 }
663 }
664
665 (all_coefs, all_fitted, total_rss)
666}
667
668fn compute_gcv(rss: f64, n_points: f64, edf: f64, m: usize) -> f64 {
670 let gcv_denom = 1.0 - edf / m as f64;
671 if gcv_denom.abs() > 1e-10 {
672 (rss / n_points) / (gcv_denom * gcv_denom)
673 } else {
674 f64::INFINITY
675 }
676}
677
678#[derive(Debug, Clone, Copy, PartialEq)]
682pub enum BasisCriterion {
683 Gcv,
685 Cv,
687 Aic,
689 Bic,
691}
692
693#[derive(Debug, Clone, PartialEq)]
695#[non_exhaustive]
696pub struct BasisNbasisCvResult {
697 pub optimal_nbasis: usize,
699 pub scores: Vec<f64>,
701 pub nbasis_range: Vec<usize>,
703 pub criterion: BasisCriterion,
705}
706
707fn evaluate_nbasis_info_criterion(
709 data: &FdMatrix,
710 argvals: &[f64],
711 nbasis_range: &[usize],
712 basis_type: &BasisType,
713 criterion: BasisCriterion,
714 lambda: f64,
715) -> Vec<f64> {
716 let mut scores = Vec::with_capacity(nbasis_range.len());
717 for &nb in nbasis_range {
718 if nb < 2 {
719 scores.push(f64::INFINITY);
720 continue;
721 }
722 let penalty = match basis_type {
723 BasisType::Bspline { order } => bspline_penalty_matrix(argvals, nb, *order, 2),
724 BasisType::Fourier { period } => fourier_penalty_matrix(nb, *period, 2),
725 };
726 let fdpar = FdPar {
727 basis_type: basis_type.clone(),
728 nbasis: nb,
729 lambda,
730 lfd_order: 2,
731 penalty_matrix: penalty,
732 };
733 match smooth_basis(data, argvals, &fdpar) {
734 Ok(result) => {
735 let score = match criterion {
736 BasisCriterion::Gcv => result.gcv,
737 BasisCriterion::Aic => result.aic,
738 BasisCriterion::Bic => result.bic,
739 BasisCriterion::Cv => unreachable!(),
740 };
741 scores.push(score);
742 }
743 Err(_) => scores.push(f64::INFINITY),
744 }
745 }
746 scores
747}
748
749fn evaluate_nbasis_cv(
751 data: &FdMatrix,
752 argvals: &[f64],
753 nbasis_range: &[usize],
754 basis_type: &BasisType,
755 lambda: f64,
756 n_folds: usize,
757) -> Vec<f64> {
758 let (n, m) = data.shape();
759 let n_folds = n_folds.max(2).min(m);
767 let point_folds = crate::cv::create_folds(m, n_folds, 42);
768 let mut scores = Vec::with_capacity(nbasis_range.len());
769
770 for &nb in nbasis_range {
771 if nb < 2 {
772 scores.push(f64::INFINITY);
773 continue;
774 }
775 let penalty = match basis_type {
776 BasisType::Bspline { order } => bspline_penalty_matrix(argvals, nb, *order, 2),
777 BasisType::Fourier { period } => fourier_penalty_matrix(nb, *period, 2),
778 };
779 let (basis_flat, actual_k) = evaluate_basis(argvals, basis_type, nb);
780 let b_full = DMatrix::from_column_slice(m, actual_k, &basis_flat);
781 let r_mat = DMatrix::from_column_slice(actual_k, actual_k, &penalty);
782
783 let mut total_se = 0.0;
784 let mut count = 0usize;
785
786 for fold in 0..n_folds {
787 let (train_pts, test_pts) = crate::cv::fold_indices(&point_folds, fold);
788 if train_pts.is_empty() || test_pts.is_empty() {
789 continue;
790 }
791 let b_train = b_full.select_rows(train_pts.iter());
794 let b_test = b_full.select_rows(test_pts.iter());
795 let btb = b_train.transpose() * &b_train;
796 let ridge_eps = 1e-10;
797 let system: DMatrix<f64> =
798 &btb + lambda * &r_mat + ridge_eps * DMatrix::<f64>::identity(actual_k, actual_k);
799 let Some(system_inv) = invert_penalized_system(&system, actual_k) else {
800 continue;
801 };
802 let proj = &system_inv * b_train.transpose(); for i in 0..n {
805 let y_train = nalgebra::DVector::from_iterator(
806 train_pts.len(),
807 train_pts.iter().map(|&j| data[(i, j)]),
808 );
809 let coefs = &proj * &y_train;
810 let pred = &b_test * &coefs; for (t_idx, &j) in test_pts.iter().enumerate() {
812 let err = data[(i, j)] - pred[t_idx];
813 total_se += err * err;
814 count += 1;
815 }
816 }
817 }
818
819 if count > 0 {
820 scores.push(total_se / count as f64);
821 } else {
822 scores.push(f64::INFINITY);
823 }
824 }
825 scores
826}
827
828pub fn basis_nbasis_cv(
831 data: &FdMatrix,
832 argvals: &[f64],
833 nbasis_range: &[usize],
834 basis_type: &BasisType,
835 criterion: BasisCriterion,
836 n_folds: usize,
837 lambda: f64,
838) -> Option<BasisNbasisCvResult> {
839 let (n, m) = data.shape();
840 if n == 0 || m == 0 || argvals.len() != m || nbasis_range.is_empty() {
841 return None;
842 }
843
844 let scores = match criterion {
845 BasisCriterion::Gcv | BasisCriterion::Aic | BasisCriterion::Bic => {
846 evaluate_nbasis_info_criterion(
847 data,
848 argvals,
849 nbasis_range,
850 basis_type,
851 criterion,
852 lambda,
853 )
854 }
855 BasisCriterion::Cv => {
856 evaluate_nbasis_cv(data, argvals, nbasis_range, basis_type, lambda, n_folds)
857 }
858 };
859
860 let (best_idx, _) = scores
861 .iter()
862 .enumerate()
863 .min_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))?;
864
865 Some(BasisNbasisCvResult {
866 optimal_nbasis: nbasis_range[best_idx],
867 scores,
868 nbasis_range: nbasis_range.to_vec(),
869 criterion,
870 })
871}
872
873#[cfg(test)]
874mod tests {
875 use super::*;
876 use crate::test_helpers::uniform_grid;
877 use std::f64::consts::PI;
878
879 #[test]
880 fn test_bspline_penalty_matrix_symmetric() {
881 let t = uniform_grid(101);
882 let penalty = bspline_penalty_matrix(&t, 15, 4, 2);
883 let _k = 15; let actual_k = (penalty.len() as f64).sqrt() as usize;
885 for i in 0..actual_k {
886 for j in 0..actual_k {
887 assert!(
888 (penalty[i + j * actual_k] - penalty[j + i * actual_k]).abs() < 1e-10,
889 "Penalty matrix not symmetric at ({}, {})",
890 i,
891 j
892 );
893 }
894 }
895 }
896
897 #[test]
898 fn test_bspline_penalty_matrix_positive_semidefinite() {
899 let t = uniform_grid(101);
900 let penalty = bspline_penalty_matrix(&t, 10, 4, 2);
901 let k = (penalty.len() as f64).sqrt() as usize;
902 for i in 0..k {
904 assert!(
905 penalty[i + i * k] >= -1e-10,
906 "Diagonal element {} is negative: {}",
907 i,
908 penalty[i + i * k]
909 );
910 }
911 }
912
913 #[test]
914 fn test_fourier_penalty_diagonal() {
915 let penalty = fourier_penalty_matrix(7, 1.0, 2);
916 for i in 0..7 {
918 for j in 0..7 {
919 if i != j {
920 assert!(
921 penalty[i + j * 7].abs() < 1e-10,
922 "Off-diagonal ({},{}) = {}",
923 i,
924 j,
925 penalty[i + j * 7]
926 );
927 }
928 }
929 }
930 assert!(penalty[0].abs() < 1e-10);
932 assert!(penalty[1 + 7] > 0.0);
934 assert!(penalty[3 + 3 * 7] > penalty[1 + 7]);
935 }
936
937 #[test]
938 fn test_smooth_basis_bspline() {
939 let m = 101;
940 let n = 5;
941 let t = uniform_grid(m);
942
943 let mut data = FdMatrix::zeros(n, m);
945 for i in 0..n {
946 for j in 0..m {
947 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.1 * (i as f64 * 0.3 + j as f64 * 0.01);
948 }
949 }
950
951 let nbasis = 15;
952 let penalty = bspline_penalty_matrix(&t, nbasis, 4, 2);
953 let _actual_k = (penalty.len() as f64).sqrt() as usize;
954
955 let fdpar = FdPar {
956 basis_type: BasisType::Bspline { order: 4 },
957 nbasis,
958 lambda: 1e-4,
959 lfd_order: 2,
960 penalty_matrix: penalty,
961 };
962
963 let result = smooth_basis(&data, &t, &fdpar);
964 assert!(result.is_ok(), "smooth_basis should succeed");
965
966 let res = result.unwrap();
967 assert_eq!(res.fitted.shape(), (n, m));
968 assert_eq!(res.coefficients.nrows(), n);
969 assert!(res.edf > 0.0, "EDF should be positive");
970 assert!(res.gcv > 0.0, "GCV should be positive");
971 }
972
973 #[test]
974 fn test_smooth_basis_fourier() {
975 let m = 101;
976 let n = 3;
977 let t = uniform_grid(m);
978
979 let mut data = FdMatrix::zeros(n, m);
980 for i in 0..n {
981 for j in 0..m {
982 data[(i, j)] = (2.0 * PI * t[j]).sin() + (4.0 * PI * t[j]).cos();
983 }
984 }
985
986 let nbasis = 7;
987 let period = 1.0;
988 let penalty = fourier_penalty_matrix(nbasis, period, 2);
989
990 let fdpar = FdPar {
991 basis_type: BasisType::Fourier { period },
992 nbasis,
993 lambda: 1e-6,
994 lfd_order: 2,
995 penalty_matrix: penalty,
996 };
997
998 let result = smooth_basis(&data, &t, &fdpar);
999 assert!(result.is_ok());
1000
1001 let res = result.unwrap();
1002 for j in 0..m {
1004 let expected = (2.0 * PI * t[j]).sin() + (4.0 * PI * t[j]).cos();
1005 assert!(
1006 (res.fitted[(0, j)] - expected).abs() < 0.1,
1007 "Fourier fit poor at j={}: got {}, expected {}",
1008 j,
1009 res.fitted[(0, j)],
1010 expected
1011 );
1012 }
1013 }
1014
1015 #[test]
1016 fn test_smooth_basis_gcv_selects_reasonable_lambda() {
1017 let m = 101;
1018 let n = 5;
1019 let t = uniform_grid(m);
1020
1021 let mut data = FdMatrix::zeros(n, m);
1022 for i in 0..n {
1023 for j in 0..m {
1024 data[(i, j)] =
1025 (2.0 * PI * t[j]).sin() + 0.1 * ((i * 37 + j * 13) % 20) as f64 / 20.0;
1026 }
1027 }
1028
1029 let basis_type = BasisType::Bspline { order: 4 };
1030 let result = smooth_basis_gcv(&data, &t, &basis_type, 15, 2, (-8.0, 4.0), 25);
1031 assert!(result.is_some(), "GCV search should succeed");
1032 }
1033
1034 #[test]
1035 fn test_smooth_basis_aic_matches_brute_force_grid() {
1036 let m = 81;
1040 let n = 4;
1041 let t = uniform_grid(m);
1042 let mut data = FdMatrix::zeros(n, m);
1043 for i in 0..n {
1044 for j in 0..m {
1045 data[(i, j)] =
1046 (2.0 * PI * t[j]).sin() + 0.15 * ((i * 41 + j * 7) % 23) as f64 / 23.0;
1047 }
1048 }
1049
1050 let basis_type = BasisType::Bspline { order: 4 };
1051 let nbasis = 12;
1052 let lfd_order = 2;
1053 let range = (-8.0, 4.0);
1054 let n_grid = 25;
1055
1056 let penalty = bspline_penalty_matrix(&t, nbasis, 4, lfd_order);
1058 let (lo, hi) = range;
1059 let mut brute_best_aic = f64::INFINITY;
1060 for k in 0..n_grid {
1061 let log_lam = lo + (hi - lo) * k as f64 / (n_grid - 1) as f64;
1062 let lam = 10.0_f64.powf(log_lam);
1063 let fdpar = FdPar {
1064 basis_type: basis_type.clone(),
1065 nbasis,
1066 lambda: lam,
1067 lfd_order,
1068 penalty_matrix: penalty.clone(),
1069 };
1070 if let Ok(result) = smooth_basis(&data, &t, &fdpar) {
1071 if result.aic < brute_best_aic {
1072 brute_best_aic = result.aic;
1073 }
1074 }
1075 }
1076
1077 let selected =
1078 smooth_basis_aic(&data, &t, &basis_type, nbasis, lfd_order, range, n_grid).unwrap();
1079 assert!(
1080 (selected.aic - brute_best_aic).abs() < 1e-9,
1081 "selected aic={}, brute-force min aic={}",
1082 selected.aic,
1083 brute_best_aic
1084 );
1085 }
1086
1087 #[test]
1088 fn test_smooth_basis_aic_prefers_smoother_fit_than_smallest_lambda() {
1089 let m = 81;
1092 let n = 4;
1093 let t = uniform_grid(m);
1094 let mut data = FdMatrix::zeros(n, m);
1095 for i in 0..n {
1097 for j in 0..m {
1098 let noise = ((((i * 91 + j * 53) % 101) as f64) / 101.0) - 0.5;
1099 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.6 * noise;
1100 }
1101 }
1102
1103 let basis_type = BasisType::Bspline { order: 4 };
1104 let nbasis = 15;
1105 let lfd_order = 2;
1106 let range = (-8.0, 4.0);
1107 let n_grid = 25;
1108
1109 let penalty = bspline_penalty_matrix(&t, nbasis, 4, lfd_order);
1111 let smallest_lam = 10.0_f64.powf(range.0);
1112 let fdpar_small = FdPar {
1113 basis_type: basis_type.clone(),
1114 nbasis,
1115 lambda: smallest_lam,
1116 lfd_order,
1117 penalty_matrix: penalty.clone(),
1118 };
1119 let overfit = smooth_basis(&data, &t, &fdpar_small).unwrap();
1120
1121 let selected =
1122 smooth_basis_aic(&data, &t, &basis_type, nbasis, lfd_order, range, n_grid).unwrap();
1123
1124 assert!(
1125 selected.edf < overfit.edf,
1126 "AIC-selected edf ({}) should be smaller (smoother) than the smallest-lambda edf ({})",
1127 selected.edf,
1128 overfit.edf
1129 );
1130 }
1131
1132 #[test]
1133 fn test_smooth_basis_large_lambda_reduces_edf() {
1134 let m = 101;
1135 let n = 3;
1136 let t = uniform_grid(m);
1137
1138 let mut data = FdMatrix::zeros(n, m);
1139 for i in 0..n {
1140 for j in 0..m {
1141 data[(i, j)] = (2.0 * PI * t[j]).sin();
1142 }
1143 }
1144
1145 let nbasis = 15;
1146 let penalty = bspline_penalty_matrix(&t, nbasis, 4, 2);
1147 let _actual_k = (penalty.len() as f64).sqrt() as usize;
1148
1149 let fdpar_small = FdPar {
1150 basis_type: BasisType::Bspline { order: 4 },
1151 nbasis,
1152 lambda: 1e-8,
1153 lfd_order: 2,
1154 penalty_matrix: penalty.clone(),
1155 };
1156 let fdpar_large = FdPar {
1157 basis_type: BasisType::Bspline { order: 4 },
1158 nbasis,
1159 lambda: 1e2,
1160 lfd_order: 2,
1161 penalty_matrix: penalty,
1162 };
1163
1164 let res_small = smooth_basis(&data, &t, &fdpar_small).unwrap();
1165 let res_large = smooth_basis(&data, &t, &fdpar_large).unwrap();
1166
1167 assert!(
1168 res_large.edf < res_small.edf,
1169 "Larger lambda should reduce EDF: {} vs {}",
1170 res_large.edf,
1171 res_small.edf
1172 );
1173 }
1174
1175 #[test]
1178 fn test_basis_nbasis_cv_gcv() {
1179 let m = 101;
1180 let n = 5;
1181 let t = uniform_grid(m);
1182 let mut data = FdMatrix::zeros(n, m);
1183 for i in 0..n {
1184 for j in 0..m {
1185 data[(i, j)] =
1186 (2.0 * PI * t[j]).sin() + 0.1 * ((i * 37 + j * 13) % 20) as f64 / 20.0;
1187 }
1188 }
1189
1190 let nbasis_range: Vec<usize> = (4..=20).step_by(2).collect();
1191 let result = basis_nbasis_cv(
1192 &data,
1193 &t,
1194 &nbasis_range,
1195 &BasisType::Bspline { order: 4 },
1196 BasisCriterion::Gcv,
1197 5,
1198 1e-4,
1199 );
1200 assert!(result.is_some());
1201 let res = result.unwrap();
1202 assert!(nbasis_range.contains(&res.optimal_nbasis));
1203 assert_eq!(res.scores.len(), nbasis_range.len());
1204 assert_eq!(res.criterion, BasisCriterion::Gcv);
1205 }
1206
1207 #[test]
1208 fn test_basis_nbasis_cv_aic_bic() {
1209 let m = 51;
1210 let n = 5;
1211 let t = uniform_grid(m);
1212 let mut data = FdMatrix::zeros(n, m);
1213 for i in 0..n {
1214 for j in 0..m {
1215 data[(i, j)] = (2.0 * PI * t[j]).sin();
1216 }
1217 }
1218
1219 let nbasis_range: Vec<usize> = vec![5, 7, 9, 11];
1220 let aic_result = basis_nbasis_cv(
1221 &data,
1222 &t,
1223 &nbasis_range,
1224 &BasisType::Bspline { order: 4 },
1225 BasisCriterion::Aic,
1226 5,
1227 0.0,
1228 );
1229 let bic_result = basis_nbasis_cv(
1230 &data,
1231 &t,
1232 &nbasis_range,
1233 &BasisType::Bspline { order: 4 },
1234 BasisCriterion::Bic,
1235 5,
1236 0.0,
1237 );
1238 assert!(aic_result.is_some());
1239 assert!(bic_result.is_some());
1240 }
1241
1242 #[test]
1243 fn test_basis_nbasis_cv_kfold() {
1244 let m = 51;
1245 let n = 10;
1246 let t = uniform_grid(m);
1247 let mut data = FdMatrix::zeros(n, m);
1248 for i in 0..n {
1249 for j in 0..m {
1250 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.05 * ((i * 7 + j * 3) % 10) as f64;
1251 }
1252 }
1253
1254 let nbasis_range: Vec<usize> = vec![5, 7, 9];
1255 let result = basis_nbasis_cv(
1256 &data,
1257 &t,
1258 &nbasis_range,
1259 &BasisType::Bspline { order: 4 },
1260 BasisCriterion::Cv,
1261 5,
1262 1e-4,
1263 );
1264 assert!(result.is_some());
1265 let res = result.unwrap();
1266 assert!(nbasis_range.contains(&res.optimal_nbasis));
1267 assert_eq!(res.criterion, BasisCriterion::Cv);
1268 }
1269
1270 #[test]
1275 fn test_basis_nbasis_cv_penalizes_overfitting() {
1276 let m = 120;
1277 let n = 6;
1278 let t = uniform_grid(m);
1279 let mut data = FdMatrix::zeros(n, m);
1280 for i in 0..n {
1281 for j in 0..m {
1282 let noise = 0.2 * (((i * 31 + j * 17) % 13) as f64 / 13.0 - 0.5);
1284 data[(i, j)] = (2.0 * PI * t[j]).sin() + noise;
1285 }
1286 }
1287
1288 let nbasis_range: Vec<usize> = vec![5, 8, 12, 20, 30];
1289 let res = basis_nbasis_cv(
1290 &data,
1291 &t,
1292 &nbasis_range,
1293 &BasisType::Bspline { order: 4 },
1294 BasisCriterion::Cv,
1295 5,
1296 1e-6,
1297 )
1298 .unwrap();
1299
1300 assert_ne!(
1301 res.optimal_nbasis, 30,
1302 "CV must not always select the maximum n_basis (GH #33); scores={:?}",
1303 res.scores
1304 );
1305 let monotone_decreasing = res.scores.windows(2).all(|w| w[1] <= w[0] + 1e-12);
1306 assert!(
1307 !monotone_decreasing,
1308 "CV scores must not be monotone-decreasing in n_basis; scores={:?}",
1309 res.scores
1310 );
1311 }
1312
1313 fn make_test_data(n: usize, m: usize) -> (FdMatrix, Vec<f64>) {
1317 let t: Vec<f64> = (0..m).map(|i| i as f64 / (m - 1) as f64).collect();
1318 let mut data = FdMatrix::zeros(n, m);
1319 for i in 0..n {
1320 for j in 0..m {
1321 data[(i, j)] = (2.0 * PI * t[j]).sin()
1322 + 0.1 * (10.0 * t[j]).sin()
1323 + 0.05 * ((i * 37 + j * 13) % 20) as f64 / 20.0;
1324 }
1325 }
1326 (data, t)
1327 }
1328
1329 fn make_bspline_fdpar(argvals: &[f64], nbasis: usize, lambda: f64) -> FdPar {
1331 let penalty = bspline_penalty_matrix(argvals, nbasis, 4, 2);
1332 FdPar {
1333 basis_type: BasisType::Bspline { order: 4 },
1334 nbasis,
1335 lambda,
1336 lfd_order: 2,
1337 penalty_matrix: penalty,
1338 }
1339 }
1340
1341 fn make_fourier_fdpar(nbasis: usize, period: f64, lambda: f64) -> FdPar {
1343 let penalty = fourier_penalty_matrix(nbasis, period, 2);
1344 FdPar {
1345 basis_type: BasisType::Fourier { period },
1346 nbasis,
1347 lambda,
1348 lfd_order: 2,
1349 penalty_matrix: penalty,
1350 }
1351 }
1352
1353 #[test]
1356 fn test_basis_type_bspline_variant() {
1357 let bt = BasisType::Bspline { order: 4 };
1358 assert_eq!(bt, BasisType::Bspline { order: 4 });
1359 assert_ne!(bt, BasisType::Bspline { order: 3 });
1361 }
1362
1363 #[test]
1364 fn test_basis_type_fourier_variant() {
1365 let bt = BasisType::Fourier { period: 1.0 };
1366 assert_eq!(bt, BasisType::Fourier { period: 1.0 });
1367 assert_ne!(bt, BasisType::Fourier { period: 2.0 });
1368 }
1369
1370 #[test]
1371 fn test_basis_type_cross_variant_inequality() {
1372 let bspline = BasisType::Bspline { order: 4 };
1373 let fourier = BasisType::Fourier { period: 1.0 };
1374 assert_ne!(bspline, fourier);
1375 }
1376
1377 #[test]
1378 fn test_basis_type_clone_and_debug() {
1379 let bt = BasisType::Bspline { order: 4 };
1380 let cloned = bt.clone();
1381 assert_eq!(bt, cloned);
1382 let debug_str = format!("{:?}", bt);
1383 assert!(debug_str.contains("Bspline"));
1384 assert!(debug_str.contains("4"));
1385 }
1386
1387 #[test]
1390 fn test_fdpar_construction_and_fields() {
1391 let penalty = vec![1.0, 0.0, 0.0, 1.0];
1392 let fdpar = FdPar {
1393 basis_type: BasisType::Bspline { order: 4 },
1394 nbasis: 2,
1395 lambda: 0.01,
1396 lfd_order: 2,
1397 penalty_matrix: penalty.clone(),
1398 };
1399 assert_eq!(fdpar.nbasis, 2);
1400 assert!((fdpar.lambda - 0.01).abs() < 1e-15);
1401 assert_eq!(fdpar.lfd_order, 2);
1402 assert_eq!(fdpar.penalty_matrix.len(), 4);
1403 }
1404
1405 #[test]
1406 fn test_fdpar_clone_and_debug() {
1407 let t = uniform_grid(50);
1408 let fdpar = make_bspline_fdpar(&t, 8, 1e-3);
1409 let cloned = fdpar.clone();
1410 assert_eq!(fdpar, cloned);
1411 let debug_str = format!("{:?}", fdpar);
1412 assert!(debug_str.contains("FdPar"));
1413 }
1414
1415 #[test]
1418 fn test_basis_criterion_variants() {
1419 assert_eq!(BasisCriterion::Gcv, BasisCriterion::Gcv);
1420 assert_eq!(BasisCriterion::Cv, BasisCriterion::Cv);
1421 assert_eq!(BasisCriterion::Aic, BasisCriterion::Aic);
1422 assert_eq!(BasisCriterion::Bic, BasisCriterion::Bic);
1423 assert_ne!(BasisCriterion::Gcv, BasisCriterion::Aic);
1424 assert_ne!(BasisCriterion::Cv, BasisCriterion::Bic);
1425 }
1426
1427 #[test]
1428 fn test_basis_criterion_copy() {
1429 let c = BasisCriterion::Gcv;
1430 let copied = c; assert_eq!(c, copied);
1432 }
1433
1434 #[test]
1435 fn test_basis_criterion_debug() {
1436 let debug_str = format!("{:?}", BasisCriterion::Bic);
1437 assert!(debug_str.contains("Bic"));
1438 }
1439
1440 #[test]
1443 fn test_smooth_basis_result_all_fields() {
1444 let (data, t) = make_test_data(3, 50);
1445 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1446 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1447
1448 assert_eq!(res.coefficients.nrows(), 3);
1450 assert!(res.coefficients.ncols() > 0);
1451 assert_eq!(res.nbasis, res.coefficients.ncols());
1452 assert_eq!(res.fitted.shape(), (3, 50));
1454 assert!(res.edf > 0.0 && res.edf <= res.nbasis as f64);
1456 assert!(res.gcv.is_finite());
1458 assert!(res.aic.is_finite());
1459 assert!(res.bic.is_finite());
1460 let k = res.nbasis;
1462 assert_eq!(res.penalty_matrix.len(), k * k);
1463 }
1464
1465 #[test]
1466 fn test_smooth_basis_result_clone() {
1467 let (data, t) = make_test_data(2, 50);
1468 let fdpar = make_bspline_fdpar(&t, 8, 1e-3);
1469 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1470 let cloned = res.clone();
1471 assert_eq!(res, cloned);
1472 }
1473
1474 #[test]
1477 fn test_smooth_basis_bspline_coefficient_shape() {
1478 let (data, t) = make_test_data(4, 50);
1479 let nbasis = 12;
1480 let fdpar = make_bspline_fdpar(&t, nbasis, 1e-4);
1481 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1482 assert_eq!(res.coefficients.nrows(), 4);
1483 assert!(res.coefficients.ncols() >= 2);
1485 assert_eq!(res.nbasis, res.coefficients.ncols());
1486 }
1487
1488 #[test]
1489 fn test_smooth_basis_bspline_fitted_values_shape() {
1490 let m = 80;
1491 let n = 6;
1492 let (data, t) = make_test_data(n, m);
1493 let fdpar = make_bspline_fdpar(&t, 15, 1e-4);
1494 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1495 assert_eq!(res.fitted.shape(), (n, m));
1496 }
1497
1498 #[test]
1499 fn test_smooth_basis_bspline_zero_lambda_interpolates() {
1500 let m = 30;
1502 let n = 2;
1503 let (data, t) = make_test_data(n, m);
1504 let fdpar = make_bspline_fdpar(&t, 15, 0.0);
1505 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1506
1507 let mut max_resid = 0.0_f64;
1509 for i in 0..n {
1510 for j in 0..m {
1511 let resid = (data[(i, j)] - res.fitted[(i, j)]).abs();
1512 max_resid = max_resid.max(resid);
1513 }
1514 }
1515 assert!(
1516 max_resid < 0.5,
1517 "Zero-lambda B-spline should closely interpolate; max_resid = {}",
1518 max_resid
1519 );
1520 }
1521
1522 #[test]
1523 fn test_smooth_basis_bspline_large_lambda_oversmooths() {
1524 let m = 50;
1527 let n = 1;
1528 let (data, t) = make_test_data(n, m);
1529
1530 let fdpar_small = make_bspline_fdpar(&t, 15, 1e-6);
1531 let res_small = smooth_basis(&data, &t, &fdpar_small).unwrap();
1532
1533 let fdpar_large = make_bspline_fdpar(&t, 15, 1e6);
1534 let res_large = smooth_basis(&data, &t, &fdpar_large).unwrap();
1535
1536 let compute_variance = |fitted: &FdMatrix, row: usize, ncols: usize| -> f64 {
1537 let vals: Vec<f64> = (0..ncols).map(|j| fitted[(row, j)]).collect();
1538 let mean = vals.iter().sum::<f64>() / ncols as f64;
1539 vals.iter().map(|&v| (v - mean).powi(2)).sum::<f64>() / ncols as f64
1540 };
1541
1542 let var_small = compute_variance(&res_small.fitted, 0, m);
1543 let var_large = compute_variance(&res_large.fitted, 0, m);
1544 assert!(
1545 var_large < var_small,
1546 "Large lambda should yield lower variance fit: var_large={}, var_small={}",
1547 var_large,
1548 var_small
1549 );
1550 }
1551
1552 #[test]
1553 fn test_smooth_basis_bspline_penalty_effect_on_smoothness() {
1554 let m = 50;
1556 let n = 1;
1557 let (data, t) = make_test_data(n, m);
1558
1559 let fdpar_small = make_bspline_fdpar(&t, 15, 1e-8);
1560 let fdpar_large = make_bspline_fdpar(&t, 15, 1.0);
1561
1562 let res_small = smooth_basis(&data, &t, &fdpar_small).unwrap();
1563 let res_large = smooth_basis(&data, &t, &fdpar_large).unwrap();
1564
1565 let roughness = |fitted: &FdMatrix, row: usize, ncols: usize| -> f64 {
1567 (1..ncols - 1)
1568 .map(|j| {
1569 let d2 = fitted[(row, j + 1)] - 2.0 * fitted[(row, j)] + fitted[(row, j - 1)];
1570 d2 * d2
1571 })
1572 .sum::<f64>()
1573 };
1574
1575 let r_small = roughness(&res_small.fitted, 0, m);
1576 let r_large = roughness(&res_large.fitted, 0, m);
1577 assert!(
1578 r_large < r_small,
1579 "Larger lambda should produce smoother fit: roughness_large={}, roughness_small={}",
1580 r_large,
1581 r_small
1582 );
1583 }
1584
1585 #[test]
1586 fn test_smooth_basis_bspline_single_curve() {
1587 let m = 50;
1588 let (data, t) = make_test_data(1, m);
1589 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1590 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1591 assert_eq!(res.fitted.nrows(), 1);
1592 assert_eq!(res.fitted.ncols(), m);
1593 assert!(res.gcv.is_finite());
1594 }
1595
1596 #[test]
1597 fn test_smooth_basis_bspline_many_curves() {
1598 let m = 50;
1599 let n = 20;
1600 let (data, t) = make_test_data(n, m);
1601 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1602 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1603 assert_eq!(res.fitted.nrows(), n);
1604 assert_eq!(res.coefficients.nrows(), n);
1605 }
1606
1607 #[test]
1608 fn test_smooth_basis_bspline_minimal_nbasis() {
1609 let m = 50;
1611 let (data, t) = make_test_data(1, m);
1612 let fdpar = make_bspline_fdpar(&t, 2, 1e-4);
1613 let res = smooth_basis(&data, &t, &fdpar);
1614 assert!(res.is_ok());
1616 }
1617
1618 #[test]
1619 fn test_smooth_basis_bspline_different_orders() {
1620 let m = 50;
1621 let (data, t) = make_test_data(2, m);
1622 let penalty3 = bspline_penalty_matrix(&t, 10, 3, 2);
1624 let fdpar3 = FdPar {
1625 basis_type: BasisType::Bspline { order: 3 },
1626 nbasis: 10,
1627 lambda: 1e-4,
1628 lfd_order: 2,
1629 penalty_matrix: penalty3,
1630 };
1631 let res3 = smooth_basis(&data, &t, &fdpar3);
1632 assert!(res3.is_ok());
1633
1634 let penalty5 = bspline_penalty_matrix(&t, 10, 5, 2);
1636 let fdpar5 = FdPar {
1637 basis_type: BasisType::Bspline { order: 5 },
1638 nbasis: 10,
1639 lambda: 1e-4,
1640 lfd_order: 2,
1641 penalty_matrix: penalty5,
1642 };
1643 let res5 = smooth_basis(&data, &t, &fdpar5);
1644 assert!(res5.is_ok());
1645 }
1646
1647 #[test]
1650 fn test_smooth_basis_fourier_coefficient_shape() {
1651 let m = 50;
1652 let n = 3;
1653 let t = uniform_grid(m);
1654 let mut data = FdMatrix::zeros(n, m);
1655 for i in 0..n {
1656 for j in 0..m {
1657 data[(i, j)] = (2.0 * PI * t[j]).sin();
1658 }
1659 }
1660 let nbasis = 7;
1661 let fdpar = make_fourier_fdpar(nbasis, 1.0, 1e-6);
1662 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1663 assert_eq!(res.coefficients.nrows(), n);
1664 assert_eq!(res.coefficients.ncols(), nbasis);
1665 assert_eq!(res.nbasis, nbasis);
1666 }
1667
1668 #[test]
1669 fn test_smooth_basis_fourier_fits_pure_sine() {
1670 let m = 100;
1672 let t = uniform_grid(m);
1673 let mut data = FdMatrix::zeros(1, m);
1674 for j in 0..m {
1675 data[(0, j)] = (2.0 * PI * t[j]).sin();
1676 }
1677 let fdpar = make_fourier_fdpar(5, 1.0, 1e-8);
1678 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1679
1680 for j in 0..m {
1681 let expected = (2.0 * PI * t[j]).sin();
1682 assert!(
1683 (res.fitted[(0, j)] - expected).abs() < 0.05,
1684 "Fourier should fit pure sine; j={}, got={}, expected={}",
1685 j,
1686 res.fitted[(0, j)],
1687 expected
1688 );
1689 }
1690 }
1691
1692 #[test]
1693 fn test_smooth_basis_fourier_different_periods() {
1694 let m = 50;
1695 let t = uniform_grid(m);
1696 let mut data = FdMatrix::zeros(1, m);
1697 for j in 0..m {
1698 data[(0, j)] = (2.0 * PI * t[j]).sin();
1699 }
1700
1701 let fdpar1 = make_fourier_fdpar(7, 1.0, 1e-6);
1703 let res1 = smooth_basis(&data, &t, &fdpar1).unwrap();
1704
1705 let fdpar2 = make_fourier_fdpar(7, 2.0, 1e-6);
1707 let res2 = smooth_basis(&data, &t, &fdpar2).unwrap();
1708
1709 assert_eq!(res1.fitted.shape(), (1, m));
1711 assert_eq!(res2.fitted.shape(), (1, m));
1712 }
1713
1714 #[test]
1715 fn test_smooth_basis_fourier_zero_lambda() {
1716 let m = 50;
1717 let t = uniform_grid(m);
1718 let mut data = FdMatrix::zeros(1, m);
1719 for j in 0..m {
1720 data[(0, j)] = (2.0 * PI * t[j]).sin() + (4.0 * PI * t[j]).cos();
1721 }
1722 let fdpar = make_fourier_fdpar(9, 1.0, 0.0);
1723 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1724 assert_eq!(res.fitted.shape(), (1, m));
1725 assert!(res.edf > 1.0);
1727 }
1728
1729 #[test]
1730 fn test_smooth_basis_fourier_large_lambda() {
1731 let m = 50;
1732 let t = uniform_grid(m);
1733 let mut data = FdMatrix::zeros(1, m);
1734 for j in 0..m {
1735 data[(0, j)] = (2.0 * PI * t[j]).sin();
1736 }
1737 let fdpar = make_fourier_fdpar(9, 1.0, 1e6);
1738 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1739 assert!(
1741 res.edf < 5.0,
1742 "Large lambda should reduce EDF; edf={}",
1743 res.edf
1744 );
1745 }
1746
1747 #[test]
1750 fn test_smooth_basis_lambda_gradient_edf() {
1751 let m = 50;
1753 let (data, t) = make_test_data(3, m);
1754 let lambdas = [1e-8, 1e-4, 1e-2, 1.0, 1e2];
1755 let mut prev_edf = f64::INFINITY;
1756 for &lam in &lambdas {
1757 let fdpar = make_bspline_fdpar(&t, 12, lam);
1758 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1759 assert!(
1760 res.edf <= prev_edf + 0.01,
1761 "EDF should decrease: lambda={}, edf={}, prev_edf={}",
1762 lam,
1763 res.edf,
1764 prev_edf
1765 );
1766 prev_edf = res.edf;
1767 }
1768 }
1769
1770 #[test]
1771 fn test_smooth_basis_lambda_gradient_rss() {
1772 let m = 50;
1774 let n = 2;
1775 let (data, t) = make_test_data(n, m);
1776 let lambdas = [0.0, 1e-6, 1e-2, 1.0, 1e4];
1777 let mut prev_rss = -1.0;
1778 for &lam in &lambdas {
1779 let fdpar = make_bspline_fdpar(&t, 12, lam);
1780 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1781 let mut rss = 0.0;
1782 for i in 0..n {
1783 for j in 0..m {
1784 rss += (data[(i, j)] - res.fitted[(i, j)]).powi(2);
1785 }
1786 }
1787 assert!(
1788 rss >= prev_rss - 1e-8,
1789 "RSS should increase: lambda={}, rss={}, prev_rss={}",
1790 lam,
1791 rss,
1792 prev_rss
1793 );
1794 prev_rss = rss;
1795 }
1796 }
1797
1798 #[test]
1801 fn test_smooth_basis_empty_data_rows() {
1802 let t = uniform_grid(50);
1803 let data = FdMatrix::zeros(0, 50);
1804 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1805 let res = smooth_basis(&data, &t, &fdpar);
1806 assert!(res.is_err());
1807 }
1808
1809 #[test]
1810 fn test_smooth_basis_empty_data_cols() {
1811 let data = FdMatrix::zeros(5, 0);
1812 let fdpar = FdPar {
1813 basis_type: BasisType::Bspline { order: 4 },
1814 nbasis: 10,
1815 lambda: 1e-4,
1816 lfd_order: 2,
1817 penalty_matrix: vec![0.0; 100],
1818 };
1819 let res = smooth_basis(&data, &[], &fdpar);
1820 assert!(res.is_err());
1821 }
1822
1823 #[test]
1824 fn test_smooth_basis_mismatched_argvals() {
1825 let t = uniform_grid(50);
1826 let data = FdMatrix::zeros(3, 40); let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1828 let res = smooth_basis(&data, &t, &fdpar);
1829 assert!(res.is_err());
1830 }
1831
1832 #[test]
1833 fn test_smooth_basis_nbasis_too_small() {
1834 let t = uniform_grid(50);
1835 let data = FdMatrix::zeros(3, 50);
1836 let fdpar = FdPar {
1838 basis_type: BasisType::Bspline { order: 4 },
1839 nbasis: 1,
1840 lambda: 1e-4,
1841 lfd_order: 2,
1842 penalty_matrix: vec![0.0; 1],
1843 };
1844 let res = smooth_basis(&data, &t, &fdpar);
1845 assert!(res.is_err());
1846 }
1847
1848 #[test]
1849 fn test_smooth_basis_error_is_invalid_dimension() {
1850 let t = uniform_grid(50);
1851 let data = FdMatrix::zeros(0, 50);
1852 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1853 let err = smooth_basis(&data, &t, &fdpar).unwrap_err();
1854 match err {
1855 crate::FdarError::InvalidDimension { .. } => {} other => panic!("Expected InvalidDimension, got {:?}", other),
1857 }
1858 }
1859
1860 #[test]
1863 fn test_bspline_penalty_matrix_different_orders() {
1864 let t = uniform_grid(101);
1865 let p1 = bspline_penalty_matrix(&t, 10, 4, 1);
1867 let p2 = bspline_penalty_matrix(&t, 10, 4, 2);
1869 assert_eq!(p1.len(), p2.len());
1871 let diff: f64 = p1.iter().zip(p2.iter()).map(|(a, b)| (a - b).abs()).sum();
1873 assert!(
1874 diff > 1e-10,
1875 "Different lfd_orders should produce different penalties"
1876 );
1877 }
1878
1879 #[test]
1880 fn test_bspline_penalty_matrix_edge_cases() {
1881 let t = vec![0.0];
1883 let p = bspline_penalty_matrix(&t, 10, 4, 2);
1884 assert!(p.iter().all(|&v| v == 0.0));
1886
1887 let t2 = uniform_grid(50);
1889 let p2 = bspline_penalty_matrix(&t2, 1, 4, 2);
1890 assert!(p2.iter().all(|&v| v == 0.0));
1891
1892 let p3 = bspline_penalty_matrix(&t2, 10, 4, 4);
1894 assert!(p3.iter().all(|&v| v == 0.0));
1895 }
1896
1897 #[test]
1898 fn test_bspline_penalty_nonnegative_diagonal() {
1899 let t = uniform_grid(101);
1900 for nbasis in [5, 10, 20] {
1901 let p = bspline_penalty_matrix(&t, nbasis, 4, 2);
1902 let k = (p.len() as f64).sqrt() as usize;
1903 for i in 0..k {
1904 assert!(
1905 p[i + i * k] >= -1e-10,
1906 "Diagonal ({},{}) negative for nbasis={}: {}",
1907 i,
1908 i,
1909 nbasis,
1910 p[i + i * k]
1911 );
1912 }
1913 }
1914 }
1915
1916 #[test]
1917 fn test_fourier_penalty_increasing_with_frequency() {
1918 let penalty = fourier_penalty_matrix(11, 1.0, 2);
1919 let k = 11;
1920 assert!(penalty[0].abs() < 1e-15);
1922 let mut prev_eigenval = 0.0;
1924 for freq in 1..=5 {
1925 let idx_sin = 2 * freq - 1;
1926 let eigenval = penalty[idx_sin + idx_sin * k];
1927 assert!(
1928 eigenval > prev_eigenval,
1929 "Higher frequency should have larger penalty: freq={}, eigenval={}, prev={}",
1930 freq,
1931 eigenval,
1932 prev_eigenval
1933 );
1934 prev_eigenval = eigenval;
1935 let idx_cos = 2 * freq;
1937 if idx_cos < k {
1938 assert!(
1939 (penalty[idx_cos + idx_cos * k] - eigenval).abs() < 1e-10,
1940 "Sin and cos penalty should match at freq {}",
1941 freq
1942 );
1943 }
1944 }
1945 }
1946
1947 #[test]
1948 fn test_fourier_penalty_different_periods() {
1949 let p1 = fourier_penalty_matrix(7, 1.0, 2);
1950 let p2 = fourier_penalty_matrix(7, 2.0, 2);
1951 for i in 1..7 {
1953 assert!(
1954 p2[i + i * 7] < p1[i + i * 7] || (p1[i + i * 7] == 0.0 && p2[i + i * 7] == 0.0),
1955 "Longer period should have smaller penalties at i={}",
1956 i
1957 );
1958 }
1959 }
1960
1961 #[test]
1962 fn test_fourier_penalty_first_order() {
1963 let p = fourier_penalty_matrix(5, 1.0, 1);
1965 let omega1 = 2.0 * PI;
1967 let expected1 = omega1.powi(2);
1968 assert!(
1969 (p[1 + 5] - expected1).abs() < 1e-6,
1970 "First-order penalty eigenval: got {}, expected {}",
1971 p[1 + 5],
1972 expected1
1973 );
1974 }
1975
1976 #[test]
1977 fn test_fourier_penalty_zero_nbasis() {
1978 let p = fourier_penalty_matrix(0, 1.0, 2);
1979 assert!(p.is_empty());
1980 }
1981
1982 #[test]
1983 fn test_fourier_penalty_nbasis_one() {
1984 let p = fourier_penalty_matrix(1, 1.0, 2);
1985 assert_eq!(p.len(), 1);
1986 assert!(p[0].abs() < 1e-15); }
1988
1989 #[test]
1992 fn test_smooth_basis_gcv_returns_valid_result() {
1993 let (data, t) = make_test_data(5, 50);
1994 let bt = BasisType::Bspline { order: 4 };
1995 let result = smooth_basis_gcv(&data, &t, &bt, 12, 2, (-6.0, 2.0), 20);
1996 assert!(result.is_some());
1997 let res = result.unwrap();
1998 assert_eq!(res.fitted.shape(), (5, 50));
1999 assert!(res.gcv.is_finite());
2000 assert!(res.edf > 0.0);
2001 }
2002
2003 #[test]
2004 fn test_smooth_basis_gcv_fourier() {
2005 let m = 80;
2006 let t = uniform_grid(m);
2007 let mut data = FdMatrix::zeros(3, m);
2008 for i in 0..3 {
2009 for j in 0..m {
2010 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.5 * (4.0 * PI * t[j]).cos();
2011 }
2012 }
2013 let bt = BasisType::Fourier { period: 1.0 };
2014 let result = smooth_basis_gcv(&data, &t, &bt, 9, 2, (-8.0, 4.0), 25);
2015 assert!(result.is_some());
2016 let res = result.unwrap();
2017 assert_eq!(res.fitted.nrows(), 3);
2018 assert_eq!(res.nbasis, 9);
2019 }
2020
2021 #[test]
2022 fn test_smooth_basis_gcv_selects_finite_gcv() {
2023 let (data, t) = make_test_data(5, 60);
2024 let bt = BasisType::Bspline { order: 4 };
2025 let res = smooth_basis_gcv(&data, &t, &bt, 12, 2, (-6.0, 2.0), 15).unwrap();
2026 assert!(res.gcv.is_finite());
2027 assert!(res.gcv > 0.0);
2028 }
2029
2030 #[test]
2031 fn test_smooth_basis_gcv_empty_data() {
2032 let data = FdMatrix::zeros(0, 50);
2033 let t = uniform_grid(50);
2034 let bt = BasisType::Bspline { order: 4 };
2035 let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-6.0, 2.0), 10);
2036 assert!(result.is_none());
2038 }
2039
2040 #[test]
2041 fn test_smooth_basis_gcv_empty_argvals() {
2042 let data = FdMatrix::zeros(5, 0);
2043 let bt = BasisType::Bspline { order: 4 };
2044 let result = smooth_basis_gcv(&data, &[], &bt, 10, 2, (-6.0, 2.0), 10);
2045 assert!(result.is_none());
2046 }
2047
2048 #[test]
2049 fn test_smooth_basis_gcv_nbasis_too_small() {
2050 let (data, t) = make_test_data(5, 50);
2051 let bt = BasisType::Bspline { order: 4 };
2052 let result = smooth_basis_gcv(&data, &t, &bt, 1, 2, (-6.0, 2.0), 10);
2053 assert!(result.is_none());
2054 }
2055
2056 #[test]
2057 fn test_smooth_basis_gcv_ngrid_too_small() {
2058 let (data, t) = make_test_data(5, 50);
2059 let bt = BasisType::Bspline { order: 4 };
2060 let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-6.0, 2.0), 1);
2061 assert!(result.is_none());
2062 }
2063
2064 #[test]
2065 fn test_smooth_basis_gcv_narrow_range() {
2066 let (data, t) = make_test_data(3, 50);
2067 let bt = BasisType::Bspline { order: 4 };
2068 let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-3.0, -2.0), 5);
2070 assert!(result.is_some());
2071 }
2072
2073 #[test]
2074 fn test_smooth_basis_gcv_wide_range() {
2075 let (data, t) = make_test_data(3, 50);
2076 let bt = BasisType::Bspline { order: 4 };
2077 let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-12.0, 8.0), 30);
2079 assert!(result.is_some());
2080 }
2081
2082 #[test]
2085 fn test_basis_nbasis_cv_scores_length() {
2086 let (data, t) = make_test_data(5, 50);
2087 let nbasis_range: Vec<usize> = vec![4, 6, 8, 10, 12];
2088 let res = basis_nbasis_cv(
2089 &data,
2090 &t,
2091 &nbasis_range,
2092 &BasisType::Bspline { order: 4 },
2093 BasisCriterion::Gcv,
2094 5,
2095 1e-4,
2096 )
2097 .unwrap();
2098 assert_eq!(res.scores.len(), 5);
2099 assert_eq!(res.nbasis_range.len(), 5);
2100 assert_eq!(res.nbasis_range, nbasis_range);
2101 }
2102
2103 #[test]
2104 fn test_basis_nbasis_cv_optimal_within_range() {
2105 let (data, t) = make_test_data(8, 50);
2106 let nbasis_range: Vec<usize> = vec![5, 7, 9, 11, 13, 15];
2107 for criterion in [
2108 BasisCriterion::Gcv,
2109 BasisCriterion::Aic,
2110 BasisCriterion::Bic,
2111 ] {
2112 let res = basis_nbasis_cv(
2113 &data,
2114 &t,
2115 &nbasis_range,
2116 &BasisType::Bspline { order: 4 },
2117 criterion,
2118 5,
2119 1e-4,
2120 )
2121 .unwrap();
2122 assert!(
2123 nbasis_range.contains(&res.optimal_nbasis),
2124 "optimal_nbasis {} not in range for {:?}",
2125 res.optimal_nbasis,
2126 criterion
2127 );
2128 }
2129 }
2130
2131 #[test]
2132 fn test_basis_nbasis_cv_fourier_gcv() {
2133 let m = 80;
2134 let t = uniform_grid(m);
2135 let mut data = FdMatrix::zeros(5, m);
2136 for i in 0..5 {
2137 for j in 0..m {
2138 data[(i, j)] = (2.0 * PI * t[j]).sin()
2139 + 0.3 * (4.0 * PI * t[j]).cos()
2140 + 0.02 * ((i * 7 + j * 3) % 10) as f64;
2141 }
2142 }
2143 let nbasis_range: Vec<usize> = vec![5, 7, 9, 11];
2144 let res = basis_nbasis_cv(
2145 &data,
2146 &t,
2147 &nbasis_range,
2148 &BasisType::Fourier { period: 1.0 },
2149 BasisCriterion::Gcv,
2150 5,
2151 1e-4,
2152 )
2153 .unwrap();
2154 assert!(nbasis_range.contains(&res.optimal_nbasis));
2155 }
2156
2157 #[test]
2158 fn test_basis_nbasis_cv_fourier_cv() {
2159 let m = 60;
2160 let t = uniform_grid(m);
2161 let n = 10;
2162 let mut data = FdMatrix::zeros(n, m);
2163 for i in 0..n {
2164 for j in 0..m {
2165 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.02 * ((i * 11 + j) % 15) as f64;
2166 }
2167 }
2168 let nbasis_range: Vec<usize> = vec![5, 7, 9];
2169 let res = basis_nbasis_cv(
2170 &data,
2171 &t,
2172 &nbasis_range,
2173 &BasisType::Fourier { period: 1.0 },
2174 BasisCriterion::Cv,
2175 5,
2176 1e-4,
2177 )
2178 .unwrap();
2179 assert!(nbasis_range.contains(&res.optimal_nbasis));
2180 assert_eq!(res.criterion, BasisCriterion::Cv);
2181 }
2182
2183 #[test]
2184 fn test_basis_nbasis_cv_with_nbasis_below_minimum() {
2185 let (data, t) = make_test_data(5, 50);
2187 let nbasis_range: Vec<usize> = vec![1, 5, 10];
2188 let res = basis_nbasis_cv(
2189 &data,
2190 &t,
2191 &nbasis_range,
2192 &BasisType::Bspline { order: 4 },
2193 BasisCriterion::Gcv,
2194 5,
2195 1e-4,
2196 )
2197 .unwrap();
2198 assert!(
2200 res.optimal_nbasis >= 5,
2201 "Should skip invalid nbasis=1, got optimal={}",
2202 res.optimal_nbasis
2203 );
2204 assert!(res.scores[0].is_infinite());
2205 }
2206
2207 #[test]
2208 fn test_basis_nbasis_cv_empty_range() {
2209 let (data, t) = make_test_data(5, 50);
2210 let nbasis_range: Vec<usize> = vec![];
2211 let result = basis_nbasis_cv(
2212 &data,
2213 &t,
2214 &nbasis_range,
2215 &BasisType::Bspline { order: 4 },
2216 BasisCriterion::Gcv,
2217 5,
2218 1e-4,
2219 );
2220 assert!(result.is_none());
2221 }
2222
2223 #[test]
2224 fn test_basis_nbasis_cv_empty_data() {
2225 let data = FdMatrix::zeros(0, 50);
2226 let t = uniform_grid(50);
2227 let nbasis_range: Vec<usize> = vec![5, 10];
2228 let result = basis_nbasis_cv(
2229 &data,
2230 &t,
2231 &nbasis_range,
2232 &BasisType::Bspline { order: 4 },
2233 BasisCriterion::Gcv,
2234 5,
2235 1e-4,
2236 );
2237 assert!(result.is_none());
2238 }
2239
2240 #[test]
2241 fn test_basis_nbasis_cv_mismatched_argvals() {
2242 let data = FdMatrix::zeros(5, 50);
2243 let t = uniform_grid(40); let nbasis_range: Vec<usize> = vec![5, 10];
2245 let result = basis_nbasis_cv(
2246 &data,
2247 &t,
2248 &nbasis_range,
2249 &BasisType::Bspline { order: 4 },
2250 BasisCriterion::Gcv,
2251 5,
2252 1e-4,
2253 );
2254 assert!(result.is_none());
2255 }
2256
2257 #[test]
2258 fn test_basis_nbasis_cv_single_nbasis() {
2259 let (data, t) = make_test_data(5, 50);
2260 let nbasis_range: Vec<usize> = vec![10];
2261 let res = basis_nbasis_cv(
2262 &data,
2263 &t,
2264 &nbasis_range,
2265 &BasisType::Bspline { order: 4 },
2266 BasisCriterion::Gcv,
2267 5,
2268 1e-4,
2269 )
2270 .unwrap();
2271 assert_eq!(res.optimal_nbasis, 10);
2272 assert_eq!(res.scores.len(), 1);
2273 }
2274
2275 #[test]
2276 fn test_basis_nbasis_cv_bic_penalizes_more_than_aic() {
2277 let (data, t) = make_test_data(5, 80);
2280 let nbasis_range: Vec<usize> = (4..=20).step_by(2).collect();
2281
2282 let aic_res = basis_nbasis_cv(
2283 &data,
2284 &t,
2285 &nbasis_range,
2286 &BasisType::Bspline { order: 4 },
2287 BasisCriterion::Aic,
2288 5,
2289 1e-4,
2290 )
2291 .unwrap();
2292 let bic_res = basis_nbasis_cv(
2293 &data,
2294 &t,
2295 &nbasis_range,
2296 &BasisType::Bspline { order: 4 },
2297 BasisCriterion::Bic,
2298 5,
2299 1e-4,
2300 )
2301 .unwrap();
2302 assert!(
2305 bic_res.optimal_nbasis <= aic_res.optimal_nbasis + 4,
2306 "BIC selected {} vs AIC selected {} -- BIC should not select much more than AIC",
2307 bic_res.optimal_nbasis,
2308 aic_res.optimal_nbasis
2309 );
2310 }
2311
2312 #[test]
2315 fn test_smooth_basis_fitted_close_to_data() {
2316 let m = 50;
2318 let n = 3;
2319 let t = uniform_grid(m);
2320 let mut data = FdMatrix::zeros(n, m);
2321 for i in 0..n {
2322 for j in 0..m {
2323 data[(i, j)] = (2.0 * PI * t[j]).sin();
2324 }
2325 }
2326 let fdpar = make_bspline_fdpar(&t, 15, 1e-6);
2327 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2328
2329 let mut max_err = 0.0_f64;
2330 for i in 0..n {
2331 for j in 0..m {
2332 let err = (data[(i, j)] - res.fitted[(i, j)]).abs();
2333 max_err = max_err.max(err);
2334 }
2335 }
2336 assert!(
2337 max_err < 0.1,
2338 "Fitted should be close to smooth data; max_err={}",
2339 max_err
2340 );
2341 }
2342
2343 #[test]
2344 fn test_smooth_basis_constant_data() {
2345 let m = 50;
2347 let n = 2;
2348 let t = uniform_grid(m);
2349 let mut data = FdMatrix::zeros(n, m);
2350 for i in 0..n {
2351 for j in 0..m {
2352 data[(i, j)] = 3.15;
2353 }
2354 }
2355 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2356 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2357 for i in 0..n {
2358 for j in 0..m {
2359 assert!(
2360 (res.fitted[(i, j)] - 3.15).abs() < 0.01,
2361 "Constant data should be fit well at ({},{}): got {}",
2362 i,
2363 j,
2364 res.fitted[(i, j)]
2365 );
2366 }
2367 }
2368 }
2369
2370 #[test]
2371 fn test_smooth_basis_linear_data() {
2372 let m = 50;
2374 let t = uniform_grid(m);
2375 let mut data = FdMatrix::zeros(1, m);
2376 for j in 0..m {
2377 data[(0, j)] = 2.0 * t[j] + 1.0;
2378 }
2379 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2380 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2381 for j in 0..m {
2382 let expected = 2.0 * t[j] + 1.0;
2383 assert!(
2384 (res.fitted[(0, j)] - expected).abs() < 0.05,
2385 "Linear data should be fit well at j={}: got {}, expected {}",
2386 j,
2387 res.fitted[(0, j)],
2388 expected
2389 );
2390 }
2391 }
2392
2393 #[test]
2396 fn test_smooth_basis_edf_bounded() {
2397 let m = 50;
2398 let (data, t) = make_test_data(3, m);
2399 let fdpar = make_bspline_fdpar(&t, 12, 1e-4);
2400 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2401 assert!(
2403 res.edf > 0.0 && res.edf <= m as f64,
2404 "EDF should be in (0, {}]; got {}",
2405 m,
2406 res.edf
2407 );
2408 }
2409
2410 #[test]
2411 fn test_smooth_basis_gcv_aic_bic_all_finite() {
2412 let (data, t) = make_test_data(4, 60);
2413 let fdpar = make_bspline_fdpar(&t, 12, 1e-3);
2414 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2415 assert!(res.gcv.is_finite(), "GCV should be finite: {}", res.gcv);
2416 assert!(res.aic.is_finite(), "AIC should be finite: {}", res.aic);
2417 assert!(res.bic.is_finite(), "BIC should be finite: {}", res.bic);
2418 }
2419
2420 #[test]
2423 fn test_smooth_basis_penalty_matrix_in_result() {
2424 let (data, t) = make_test_data(3, 50);
2425 let nbasis = 10;
2426 let fdpar = make_bspline_fdpar(&t, nbasis, 1e-4);
2427 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2428 let k = res.nbasis;
2429 assert_eq!(
2430 res.penalty_matrix.len(),
2431 k * k,
2432 "Penalty matrix should be k*k = {}*{} = {}; got {}",
2433 k,
2434 k,
2435 k * k,
2436 res.penalty_matrix.len()
2437 );
2438 }
2439
2440 #[test]
2443 fn test_smooth_basis_identical_curves_same_coefficients() {
2444 let m = 50;
2445 let t = uniform_grid(m);
2446 let curve: Vec<f64> = (0..m).map(|j| (2.0 * PI * t[j]).sin()).collect();
2447 let n = 4;
2448 let mut data = FdMatrix::zeros(n, m);
2449 for i in 0..n {
2450 for j in 0..m {
2451 data[(i, j)] = curve[j];
2452 }
2453 }
2454 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2455 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2456
2457 let k = res.coefficients.ncols();
2459 for i in 1..n {
2460 for j in 0..k {
2461 assert!(
2462 (res.coefficients[(i, j)] - res.coefficients[(0, j)]).abs() < 1e-10,
2463 "Identical curves should have identical coefficients: curve {} col {} differs",
2464 i,
2465 j
2466 );
2467 }
2468 }
2469 }
2470
2471 #[test]
2474 fn test_basis_nbasis_cv_different_nfolds() {
2475 let (data, t) = make_test_data(12, 50);
2476 let nbasis_range: Vec<usize> = vec![5, 8, 11];
2477 for nfolds in [2, 3, 5, 10] {
2478 let res = basis_nbasis_cv(
2479 &data,
2480 &t,
2481 &nbasis_range,
2482 &BasisType::Bspline { order: 4 },
2483 BasisCriterion::Cv,
2484 nfolds,
2485 1e-4,
2486 );
2487 assert!(res.is_some(), "CV should succeed with nfolds={}", nfolds);
2488 let r = res.unwrap();
2489 assert!(nbasis_range.contains(&r.optimal_nbasis));
2490 }
2491 }
2492
2493 #[test]
2496 fn test_smooth_basis_many_basis_functions() {
2497 let m = 100;
2498 let (data, t) = make_test_data(2, m);
2499 let fdpar = make_bspline_fdpar(&t, 40, 1e-2);
2501 let res = smooth_basis(&data, &t, &fdpar);
2502 assert!(
2503 res.is_ok(),
2504 "Should handle many basis functions with penalty"
2505 );
2506 }
2507
2508 #[test]
2511 fn test_smooth_basis_bspline_vs_fourier_different_results() {
2512 let m = 50;
2513 let (data, t) = make_test_data(2, m);
2514 let fdpar_bs = make_bspline_fdpar(&t, 9, 1e-4);
2515 let fdpar_f = make_fourier_fdpar(9, 1.0, 1e-4);
2516 let res_bs = smooth_basis(&data, &t, &fdpar_bs).unwrap();
2517 let res_f = smooth_basis(&data, &t, &fdpar_f).unwrap();
2518 let diff: f64 = (0..m)
2520 .map(|j| (res_bs.fitted[(0, j)] - res_f.fitted[(0, j)]).abs())
2521 .sum();
2522 assert!(
2524 diff > 1e-10,
2525 "B-spline and Fourier fits should differ for the same data"
2526 );
2527 }
2528
2529 #[test]
2532 fn test_smooth_basis_gcv_positive_for_noisy_data() {
2533 let m = 50;
2534 let t = uniform_grid(m);
2535 let mut data = FdMatrix::zeros(1, m);
2536 for j in 0..m {
2537 data[(0, j)] = (2.0 * PI * t[j]).sin() + 0.5 * ((j * 37) % 20) as f64 / 20.0 - 0.25;
2539 }
2540 let fdpar = make_bspline_fdpar(&t, 10, 1e-3);
2541 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2542 assert!(res.gcv > 0.0, "GCV should be positive for noisy data");
2543 }
2544
2545 #[test]
2548 fn test_smooth_basis_different_lfd_orders() {
2549 let m = 50;
2550 let (data, t) = make_test_data(2, m);
2551
2552 let penalty1 = bspline_penalty_matrix(&t, 10, 4, 1);
2554 let fdpar1 = FdPar {
2555 basis_type: BasisType::Bspline { order: 4 },
2556 nbasis: 10,
2557 lambda: 1e-2,
2558 lfd_order: 1,
2559 penalty_matrix: penalty1,
2560 };
2561 let res1 = smooth_basis(&data, &t, &fdpar1);
2562 assert!(res1.is_ok());
2563
2564 let penalty2 = bspline_penalty_matrix(&t, 10, 4, 2);
2566 let fdpar2 = FdPar {
2567 basis_type: BasisType::Bspline { order: 4 },
2568 nbasis: 10,
2569 lambda: 1e-2,
2570 lfd_order: 2,
2571 penalty_matrix: penalty2,
2572 };
2573 let res2 = smooth_basis(&data, &t, &fdpar2);
2574 assert!(res2.is_ok());
2575
2576 let r1 = res1.unwrap();
2578 let r2 = res2.unwrap();
2579 let diff: f64 = (0..m)
2580 .map(|j| (r1.fitted[(0, j)] - r2.fitted[(0, j)]).abs())
2581 .sum();
2582 assert!(
2583 diff > 1e-10,
2584 "Different lfd_orders should produce different fits"
2585 );
2586 }
2587
2588 #[test]
2591 fn test_basis_nbasis_cv_result_fields() {
2592 let (data, t) = make_test_data(6, 50);
2593 let nbasis_range: Vec<usize> = vec![5, 7, 9, 11, 13];
2594 let res = basis_nbasis_cv(
2595 &data,
2596 &t,
2597 &nbasis_range,
2598 &BasisType::Bspline { order: 4 },
2599 BasisCriterion::Aic,
2600 5,
2601 1e-4,
2602 )
2603 .unwrap();
2604
2605 assert!(nbasis_range.contains(&res.optimal_nbasis));
2606 assert_eq!(res.scores.len(), nbasis_range.len());
2607 assert_eq!(res.nbasis_range, nbasis_range);
2608 assert_eq!(res.criterion, BasisCriterion::Aic);
2609 let min_score = res.scores.iter().copied().fold(f64::INFINITY, f64::min);
2611 let best_idx = res
2612 .scores
2613 .iter()
2614 .position(|&s| (s - min_score).abs() < 1e-15)
2615 .unwrap();
2616 assert_eq!(res.optimal_nbasis, nbasis_range[best_idx]);
2617 }
2618
2619 #[test]
2620 fn test_basis_nbasis_cv_result_clone() {
2621 let (data, t) = make_test_data(5, 50);
2622 let nbasis_range: Vec<usize> = vec![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 let cloned = res.clone();
2634 assert_eq!(res, cloned);
2635 }
2636
2637 #[test]
2640 fn test_smooth_basis_nonuniform_argvals() {
2641 let m = 50;
2642 let t: Vec<f64> = (0..m)
2644 .map(|i| {
2645 let x = i as f64 / (m - 1) as f64;
2646 0.5 * (1.0 - (PI * x).cos())
2647 })
2648 .collect();
2649 let mut data = FdMatrix::zeros(2, m);
2650 for i in 0..2 {
2651 for j in 0..m {
2652 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.1 * i as f64;
2653 }
2654 }
2655 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2656 let res = smooth_basis(&data, &t, &fdpar);
2657 assert!(res.is_ok(), "Should handle non-uniform argvals");
2658 let r = res.unwrap();
2659 assert_eq!(r.fitted.shape(), (2, m));
2660 }
2661
2662 #[test]
2665 fn test_smooth_basis_very_small_lambda() {
2666 let m = 50;
2667 let (data, t) = make_test_data(2, m);
2668 let fdpar = make_bspline_fdpar(&t, 10, 1e-15);
2669 let res = smooth_basis(&data, &t, &fdpar);
2670 assert!(res.is_ok(), "Should handle very small lambda");
2671 }
2672
2673 #[test]
2674 fn test_smooth_basis_very_large_lambda() {
2675 let m = 50;
2676 let (data, t) = make_test_data(2, m);
2677 let fdpar = make_bspline_fdpar(&t, 10, 1e10);
2678 let res = smooth_basis(&data, &t, &fdpar);
2679 assert!(res.is_ok(), "Should handle very large lambda");
2680 }
2681
2682 #[test]
2685 fn test_smooth_basis_multi_curve_vs_single_curve() {
2686 let m = 50;
2688 let n = 3;
2689 let (data, t) = make_test_data(n, m);
2690 let fdpar = make_bspline_fdpar(&t, 10, 1e-3);
2691
2692 let res_all = smooth_basis(&data, &t, &fdpar).unwrap();
2694
2695 for i in 0..n {
2697 let mut single = FdMatrix::zeros(1, m);
2698 for j in 0..m {
2699 single[(0, j)] = data[(i, j)];
2700 }
2701 let res_single = smooth_basis(&single, &t, &fdpar).unwrap();
2702 for j in 0..m {
2703 assert!(
2704 (res_all.fitted[(i, j)] - res_single.fitted[(0, j)]).abs() < 1e-10,
2705 "Multi-curve fit should match single-curve fit: curve {} point {}",
2706 i,
2707 j
2708 );
2709 }
2710 }
2711 }
2712
2713 #[test]
2716 fn test_basis_nbasis_cv_all_criteria_finite_scores() {
2717 let (data, t) = make_test_data(10, 60);
2718 let nbasis_range: Vec<usize> = vec![5, 7, 9, 11];
2719
2720 for criterion in [
2721 BasisCriterion::Gcv,
2722 BasisCriterion::Aic,
2723 BasisCriterion::Bic,
2724 BasisCriterion::Cv,
2725 ] {
2726 let res = basis_nbasis_cv(
2727 &data,
2728 &t,
2729 &nbasis_range,
2730 &BasisType::Bspline { order: 4 },
2731 criterion,
2732 5,
2733 1e-4,
2734 )
2735 .unwrap();
2736 let finite_count = res.scores.iter().filter(|s| s.is_finite()).count();
2738 assert!(
2739 finite_count > 0,
2740 "At least one score should be finite for {:?}",
2741 criterion
2742 );
2743 }
2744 }
2745
2746 #[test]
2749 fn test_smooth_basis_gcv_config_default() {
2750 let config = SmoothBasisGcvConfig::default();
2751 assert_eq!(config.basis_type, BasisType::Bspline { order: 4 });
2752 assert_eq!(config.nbasis, 15);
2753 assert_eq!(config.lfd_order, 2);
2754 assert_eq!(config.log_lambda_range, (-10.0, 2.0));
2755 assert_eq!(config.n_grid, 50);
2756 }
2757
2758 #[test]
2759 fn test_smooth_basis_gcv_config_clone_eq() {
2760 let config = SmoothBasisGcvConfig {
2761 nbasis: 20,
2762 ..SmoothBasisGcvConfig::default()
2763 };
2764 let cloned = config.clone();
2765 assert_eq!(config, cloned);
2766 }
2767
2768 #[test]
2769 fn test_smooth_basis_gcv_config_debug() {
2770 let config = SmoothBasisGcvConfig::default();
2771 let debug_str = format!("{:?}", config);
2772 assert!(debug_str.contains("SmoothBasisGcvConfig"));
2773 assert!(debug_str.contains("nbasis"));
2774 }
2775
2776 #[test]
2777 fn test_smooth_basis_gcv_config_partial_override() {
2778 let config = SmoothBasisGcvConfig {
2779 basis_type: BasisType::Fourier { period: 2.0 },
2780 n_grid: 100,
2781 ..SmoothBasisGcvConfig::default()
2782 };
2783 assert_eq!(config.basis_type, BasisType::Fourier { period: 2.0 });
2784 assert_eq!(config.n_grid, 100);
2785 assert_eq!(config.nbasis, 15);
2787 assert_eq!(config.lfd_order, 2);
2788 }
2789
2790 #[test]
2791 fn test_smooth_basis_gcv_with_config_default() {
2792 let (data, t) = make_test_data(5, 101);
2793 let config = SmoothBasisGcvConfig::default();
2794 let result = smooth_basis_gcv_with_config(&data, &t, &config);
2795 assert!(result.is_ok(), "GCV with default config should succeed");
2796 let res = result.unwrap();
2797 assert_eq!(res.fitted.shape(), (5, 101));
2798 assert!(res.edf > 0.0);
2799 assert!(res.gcv.is_finite());
2800 }
2801
2802 #[test]
2803 fn test_smooth_basis_gcv_with_config_custom() {
2804 let (data, t) = make_test_data(3, 50);
2805 let config = SmoothBasisGcvConfig {
2806 nbasis: 10,
2807 log_lambda_range: (-6.0, 0.0),
2808 n_grid: 15,
2809 ..SmoothBasisGcvConfig::default()
2810 };
2811 let result = smooth_basis_gcv_with_config(&data, &t, &config);
2812 assert!(result.is_ok());
2813 }
2814
2815 #[test]
2816 fn test_smooth_basis_gcv_with_config_matches_direct() {
2817 let (data, t) = make_test_data(3, 50);
2818 let config = SmoothBasisGcvConfig {
2819 nbasis: 10,
2820 log_lambda_range: (-6.0, 0.0),
2821 n_grid: 20,
2822 ..SmoothBasisGcvConfig::default()
2823 };
2824 let with_config = smooth_basis_gcv_with_config(&data, &t, &config).unwrap();
2825 let direct = smooth_basis_gcv(
2826 &data,
2827 &t,
2828 &config.basis_type,
2829 config.nbasis,
2830 config.lfd_order,
2831 config.log_lambda_range,
2832 config.n_grid,
2833 )
2834 .unwrap();
2835 assert_eq!(with_config.gcv, direct.gcv);
2836 assert_eq!(with_config.edf, direct.edf);
2837 assert_eq!(with_config.nbasis, direct.nbasis);
2838 }
2839
2840 #[test]
2841 fn test_smooth_basis_gcv_with_config_fourier() {
2842 let m = 100;
2843 let t = uniform_grid(m);
2844 let mut data = FdMatrix::zeros(2, m);
2845 for i in 0..2 {
2846 for j in 0..m {
2847 data[(i, j)] = (2.0 * PI * t[j]).sin() + (4.0 * PI * t[j]).cos();
2848 }
2849 }
2850 let config = SmoothBasisGcvConfig {
2851 basis_type: BasisType::Fourier { period: 1.0 },
2852 nbasis: 7,
2853 n_grid: 20,
2854 ..SmoothBasisGcvConfig::default()
2855 };
2856 let result = smooth_basis_gcv_with_config(&data, &t, &config);
2857 assert!(result.is_ok());
2858 }
2859
2860 #[test]
2863 fn test_basis_nbasis_cv_config_default() {
2864 let config = BasisNbasisCvConfig::default();
2865 assert_eq!(config.basis_type, BasisType::Bspline { order: 4 });
2866 assert_eq!(config.nbasis_range, (5, 30));
2867 assert!((config.lambda - 1e-4).abs() < 1e-15);
2868 assert_eq!(config.lfd_order, 2);
2869 assert_eq!(config.n_folds, 5);
2870 assert_eq!(config.criterion, BasisCriterion::Gcv);
2871 }
2872
2873 #[test]
2874 fn test_basis_nbasis_cv_config_clone_eq() {
2875 let config = BasisNbasisCvConfig {
2876 nbasis_range: (4, 15),
2877 ..BasisNbasisCvConfig::default()
2878 };
2879 let cloned = config.clone();
2880 assert_eq!(config, cloned);
2881 }
2882
2883 #[test]
2884 fn test_basis_nbasis_cv_config_debug() {
2885 let config = BasisNbasisCvConfig::default();
2886 let debug_str = format!("{:?}", config);
2887 assert!(debug_str.contains("BasisNbasisCvConfig"));
2888 assert!(debug_str.contains("nbasis_range"));
2889 }
2890
2891 #[test]
2892 fn test_basis_nbasis_cv_config_partial_override() {
2893 let config = BasisNbasisCvConfig {
2894 criterion: BasisCriterion::Aic,
2895 lambda: 1e-2,
2896 ..BasisNbasisCvConfig::default()
2897 };
2898 assert_eq!(config.criterion, BasisCriterion::Aic);
2899 assert!((config.lambda - 1e-2).abs() < 1e-15);
2900 assert_eq!(config.nbasis_range, (5, 30));
2902 assert_eq!(config.n_folds, 5);
2903 }
2904
2905 #[test]
2906 fn test_basis_nbasis_cv_with_config_default() {
2907 let (data, t) = make_test_data(5, 51);
2908 let config = BasisNbasisCvConfig {
2909 nbasis_range: (5, 12),
2910 ..BasisNbasisCvConfig::default()
2911 };
2912 let result = basis_nbasis_cv_with_config(&data, &t, &config);
2913 assert!(
2914 result.is_ok(),
2915 "nbasis CV with default config should succeed"
2916 );
2917 let res = result.unwrap();
2918 assert!(res.optimal_nbasis >= 5 && res.optimal_nbasis <= 12);
2919 assert_eq!(res.scores.len(), 8); assert_eq!(res.criterion, BasisCriterion::Gcv);
2921 }
2922
2923 #[test]
2924 fn test_basis_nbasis_cv_with_config_aic() {
2925 let (data, t) = make_test_data(5, 51);
2926 let config = BasisNbasisCvConfig {
2927 nbasis_range: (5, 10),
2928 criterion: BasisCriterion::Aic,
2929 ..BasisNbasisCvConfig::default()
2930 };
2931 let result = basis_nbasis_cv_with_config(&data, &t, &config);
2932 assert!(result.is_ok());
2933 assert_eq!(result.unwrap().criterion, BasisCriterion::Aic);
2934 }
2935
2936 #[test]
2937 fn test_basis_nbasis_cv_with_config_cv_folds() {
2938 let (data, t) = make_test_data(10, 51);
2939 let config = BasisNbasisCvConfig {
2940 nbasis_range: (5, 9),
2941 criterion: BasisCriterion::Cv,
2942 n_folds: 3,
2943 ..BasisNbasisCvConfig::default()
2944 };
2945 let result = basis_nbasis_cv_with_config(&data, &t, &config);
2946 assert!(result.is_ok());
2947 assert_eq!(result.unwrap().criterion, BasisCriterion::Cv);
2948 }
2949
2950 #[test]
2951 fn test_basis_nbasis_cv_with_config_matches_direct() {
2952 let (data, t) = make_test_data(5, 51);
2953 let config = BasisNbasisCvConfig {
2954 nbasis_range: (5, 10),
2955 criterion: BasisCriterion::Bic,
2956 lambda: 1e-3,
2957 ..BasisNbasisCvConfig::default()
2958 };
2959 let with_config = basis_nbasis_cv_with_config(&data, &t, &config).unwrap();
2960 let nbasis_range: Vec<usize> = (5..=10).collect();
2961 let direct = basis_nbasis_cv(
2962 &data,
2963 &t,
2964 &nbasis_range,
2965 &config.basis_type,
2966 config.criterion,
2967 config.n_folds,
2968 config.lambda,
2969 )
2970 .unwrap();
2971 assert_eq!(with_config.optimal_nbasis, direct.optimal_nbasis);
2972 assert_eq!(with_config.scores, direct.scores);
2973 assert_eq!(with_config.nbasis_range, direct.nbasis_range);
2974 }
2975
2976 #[test]
2977 fn test_basis_nbasis_cv_with_config_nbasis_range_expansion() {
2978 let (data, t) = make_test_data(5, 51);
2979 let config = BasisNbasisCvConfig {
2980 nbasis_range: (7, 7), ..BasisNbasisCvConfig::default()
2982 };
2983 let result = basis_nbasis_cv_with_config(&data, &t, &config);
2984 assert!(result.is_ok());
2985 let res = result.unwrap();
2986 assert_eq!(res.optimal_nbasis, 7);
2987 assert_eq!(res.scores.len(), 1);
2988 }
2989}