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
301#[derive(Debug, Clone, PartialEq)]
319pub struct SmoothBasisGcvConfig {
320 pub basis_type: BasisType,
322 pub nbasis: usize,
324 pub lfd_order: usize,
326 pub log_lambda_range: (f64, f64),
328 pub n_grid: usize,
330}
331
332impl Default for SmoothBasisGcvConfig {
333 fn default() -> Self {
334 Self {
335 basis_type: BasisType::Bspline { order: 4 },
336 nbasis: 15,
337 lfd_order: 2,
338 log_lambda_range: (-10.0, 2.0),
339 n_grid: 50,
340 }
341 }
342}
343
344#[must_use = "expensive computation whose result should not be discarded"]
359pub fn smooth_basis_gcv_with_config(
360 data: &FdMatrix,
361 argvals: &[f64],
362 config: &SmoothBasisGcvConfig,
363) -> Result<SmoothBasisResult, crate::FdarError> {
364 smooth_basis_gcv(
365 data,
366 argvals,
367 &config.basis_type,
368 config.nbasis,
369 config.lfd_order,
370 config.log_lambda_range,
371 config.n_grid,
372 )
373 .ok_or_else(|| crate::FdarError::ComputationFailed {
374 operation: "smooth_basis_gcv_with_config",
375 detail: "no valid smoothing result found in GCV lambda search".to_string(),
376 })
377}
378
379#[derive(Debug, Clone, PartialEq)]
395pub struct BasisNbasisCvConfig {
396 pub basis_type: BasisType,
398 pub nbasis_range: (usize, usize),
400 pub lambda: f64,
402 pub lfd_order: usize,
404 pub n_folds: usize,
406 pub criterion: BasisCriterion,
408}
409
410impl Default for BasisNbasisCvConfig {
411 fn default() -> Self {
412 Self {
413 basis_type: BasisType::Bspline { order: 4 },
414 nbasis_range: (5, 30),
415 lambda: 1e-4,
416 lfd_order: 2,
417 n_folds: 5,
418 criterion: BasisCriterion::Gcv,
419 }
420 }
421}
422
423#[must_use = "expensive computation whose result should not be discarded"]
441pub fn basis_nbasis_cv_with_config(
442 data: &FdMatrix,
443 argvals: &[f64],
444 config: &BasisNbasisCvConfig,
445) -> Result<BasisNbasisCvResult, crate::FdarError> {
446 let nbasis_range: Vec<usize> = (config.nbasis_range.0..=config.nbasis_range.1).collect();
447 basis_nbasis_cv(
448 data,
449 argvals,
450 &nbasis_range,
451 &config.basis_type,
452 config.criterion,
453 config.n_folds,
454 config.lambda,
455 )
456 .ok_or_else(|| crate::FdarError::ComputationFailed {
457 operation: "basis_nbasis_cv_with_config",
458 detail: "no valid result found in nbasis CV search".to_string(),
459 })
460}
461
462fn differentiate_basis_columns(
466 basis: &[f64],
467 n_quad: usize,
468 nbasis: usize,
469 h: f64,
470 lfd_order: usize,
471) -> Vec<f64> {
472 let mut deriv = basis.to_vec();
473 for _ in 0..lfd_order {
474 let mut new_deriv = vec![0.0; n_quad * nbasis];
475 for j in 0..nbasis {
476 let col: Vec<f64> = (0..n_quad).map(|i| deriv[i + j * n_quad]).collect();
477 let grad = crate::helpers::gradient_uniform(&col, h);
478 for i in 0..n_quad {
479 new_deriv[i + j * n_quad] = grad[i];
480 }
481 }
482 deriv = new_deriv;
483 }
484 deriv
485}
486
487fn integrate_symmetric_penalty(
489 deriv_basis: &[f64],
490 weights: &[f64],
491 k: usize,
492 n_quad: usize,
493) -> Vec<f64> {
494 let mut penalty = vec![0.0; k * k];
495 for j in 0..k {
496 for l in j..k {
497 let mut val = 0.0;
498 for i in 0..n_quad {
499 val += deriv_basis[i + j * n_quad] * deriv_basis[i + l * n_quad] * weights[i];
500 }
501 penalty[j + l * k] = val;
502 penalty[l + j * k] = val;
503 }
504 }
505 penalty
506}
507
508fn evaluate_basis(argvals: &[f64], basis_type: &BasisType, nbasis: usize) -> (Vec<f64>, usize) {
510 let m = argvals.len();
511 match basis_type {
512 BasisType::Bspline { order } => {
513 let nknots = nbasis.saturating_sub(*order).max(2);
514 let basis = bspline_basis(argvals, nknots, *order);
515 let actual = basis.len() / m;
516 (basis, actual)
517 }
518 BasisType::Fourier { period } => {
519 let basis = fourier_basis_with_period(argvals, nbasis, *period);
520 (basis, nbasis)
521 }
522 }
523}
524
525fn invert_penalized_system(system: &DMatrix<f64>, k: usize) -> Option<DMatrix<f64>> {
527 if let Some(chol) = system.clone().cholesky() {
528 return Some(chol.inverse());
529 }
530 let svd = nalgebra::SVD::new(system.clone(), true, true);
532 let u = svd.u.as_ref()?;
533 let v_t = svd.v_t.as_ref()?;
534 let max_sv: f64 = svd.singular_values.iter().copied().fold(0.0_f64, f64::max);
535 let eps = 1e-10 * max_sv;
536 let mut inv = DMatrix::<f64>::zeros(k, k);
537 for ii in 0..k {
538 for jj in 0..k {
539 let mut sum = 0.0;
540 for s in 0..k.min(svd.singular_values.len()) {
541 if svd.singular_values[s] > eps {
542 sum += v_t[(s, ii)] / svd.singular_values[s] * u[(jj, s)];
543 }
544 }
545 inv[(ii, jj)] = sum;
546 }
547 }
548 Some(inv)
549}
550
551fn project_all_curves(
553 data: &FdMatrix,
554 b_mat: &DMatrix<f64>,
555 proj: &DMatrix<f64>,
556 n: usize,
557 m: usize,
558 k: usize,
559) -> (FdMatrix, FdMatrix, f64) {
560 let mut all_coefs = FdMatrix::zeros(n, k);
561 let mut all_fitted = FdMatrix::zeros(n, m);
562 let mut total_rss = 0.0;
563
564 for i in 0..n {
565 let curve: Vec<f64> = (0..m).map(|j| data[(i, j)]).collect();
566 let y_vec = nalgebra::DVector::from_vec(curve.clone());
567 let coefs = proj * &y_vec;
568
569 for j in 0..k {
570 all_coefs[(i, j)] = coefs[j];
571 }
572 let fitted = b_mat * &coefs;
573 for j in 0..m {
574 all_fitted[(i, j)] = fitted[j];
575 let resid = curve[j] - fitted[j];
576 total_rss += resid * resid;
577 }
578 }
579
580 (all_coefs, all_fitted, total_rss)
581}
582
583fn compute_gcv(rss: f64, n_points: f64, edf: f64, m: usize) -> f64 {
585 let gcv_denom = 1.0 - edf / m as f64;
586 if gcv_denom.abs() > 1e-10 {
587 (rss / n_points) / (gcv_denom * gcv_denom)
588 } else {
589 f64::INFINITY
590 }
591}
592
593#[derive(Debug, Clone, Copy, PartialEq)]
597pub enum BasisCriterion {
598 Gcv,
600 Cv,
602 Aic,
604 Bic,
606}
607
608#[derive(Debug, Clone, PartialEq)]
610#[non_exhaustive]
611pub struct BasisNbasisCvResult {
612 pub optimal_nbasis: usize,
614 pub scores: Vec<f64>,
616 pub nbasis_range: Vec<usize>,
618 pub criterion: BasisCriterion,
620}
621
622fn evaluate_nbasis_info_criterion(
624 data: &FdMatrix,
625 argvals: &[f64],
626 nbasis_range: &[usize],
627 basis_type: &BasisType,
628 criterion: BasisCriterion,
629 lambda: f64,
630) -> Vec<f64> {
631 let mut scores = Vec::with_capacity(nbasis_range.len());
632 for &nb in nbasis_range {
633 if nb < 2 {
634 scores.push(f64::INFINITY);
635 continue;
636 }
637 let penalty = match basis_type {
638 BasisType::Bspline { order } => bspline_penalty_matrix(argvals, nb, *order, 2),
639 BasisType::Fourier { period } => fourier_penalty_matrix(nb, *period, 2),
640 };
641 let fdpar = FdPar {
642 basis_type: basis_type.clone(),
643 nbasis: nb,
644 lambda,
645 lfd_order: 2,
646 penalty_matrix: penalty,
647 };
648 match smooth_basis(data, argvals, &fdpar) {
649 Ok(result) => {
650 let score = match criterion {
651 BasisCriterion::Gcv => result.gcv,
652 BasisCriterion::Aic => result.aic,
653 BasisCriterion::Bic => result.bic,
654 BasisCriterion::Cv => unreachable!(),
655 };
656 scores.push(score);
657 }
658 Err(_) => scores.push(f64::INFINITY),
659 }
660 }
661 scores
662}
663
664fn evaluate_nbasis_cv(
666 data: &FdMatrix,
667 argvals: &[f64],
668 nbasis_range: &[usize],
669 basis_type: &BasisType,
670 lambda: f64,
671 n_folds: usize,
672) -> Vec<f64> {
673 let (n, m) = data.shape();
674 let n_folds = n_folds.max(2).min(m);
682 let point_folds = crate::cv::create_folds(m, n_folds, 42);
683 let mut scores = Vec::with_capacity(nbasis_range.len());
684
685 for &nb in nbasis_range {
686 if nb < 2 {
687 scores.push(f64::INFINITY);
688 continue;
689 }
690 let penalty = match basis_type {
691 BasisType::Bspline { order } => bspline_penalty_matrix(argvals, nb, *order, 2),
692 BasisType::Fourier { period } => fourier_penalty_matrix(nb, *period, 2),
693 };
694 let (basis_flat, actual_k) = evaluate_basis(argvals, basis_type, nb);
695 let b_full = DMatrix::from_column_slice(m, actual_k, &basis_flat);
696 let r_mat = DMatrix::from_column_slice(actual_k, actual_k, &penalty);
697
698 let mut total_se = 0.0;
699 let mut count = 0usize;
700
701 for fold in 0..n_folds {
702 let (train_pts, test_pts) = crate::cv::fold_indices(&point_folds, fold);
703 if train_pts.is_empty() || test_pts.is_empty() {
704 continue;
705 }
706 let b_train = b_full.select_rows(train_pts.iter());
709 let b_test = b_full.select_rows(test_pts.iter());
710 let btb = b_train.transpose() * &b_train;
711 let ridge_eps = 1e-10;
712 let system: DMatrix<f64> =
713 &btb + lambda * &r_mat + ridge_eps * DMatrix::<f64>::identity(actual_k, actual_k);
714 let Some(system_inv) = invert_penalized_system(&system, actual_k) else {
715 continue;
716 };
717 let proj = &system_inv * b_train.transpose(); for i in 0..n {
720 let y_train = nalgebra::DVector::from_iterator(
721 train_pts.len(),
722 train_pts.iter().map(|&j| data[(i, j)]),
723 );
724 let coefs = &proj * &y_train;
725 let pred = &b_test * &coefs; for (t_idx, &j) in test_pts.iter().enumerate() {
727 let err = data[(i, j)] - pred[t_idx];
728 total_se += err * err;
729 count += 1;
730 }
731 }
732 }
733
734 if count > 0 {
735 scores.push(total_se / count as f64);
736 } else {
737 scores.push(f64::INFINITY);
738 }
739 }
740 scores
741}
742
743pub fn basis_nbasis_cv(
746 data: &FdMatrix,
747 argvals: &[f64],
748 nbasis_range: &[usize],
749 basis_type: &BasisType,
750 criterion: BasisCriterion,
751 n_folds: usize,
752 lambda: f64,
753) -> Option<BasisNbasisCvResult> {
754 let (n, m) = data.shape();
755 if n == 0 || m == 0 || argvals.len() != m || nbasis_range.is_empty() {
756 return None;
757 }
758
759 let scores = match criterion {
760 BasisCriterion::Gcv | BasisCriterion::Aic | BasisCriterion::Bic => {
761 evaluate_nbasis_info_criterion(
762 data,
763 argvals,
764 nbasis_range,
765 basis_type,
766 criterion,
767 lambda,
768 )
769 }
770 BasisCriterion::Cv => {
771 evaluate_nbasis_cv(data, argvals, nbasis_range, basis_type, lambda, n_folds)
772 }
773 };
774
775 let (best_idx, _) = scores
776 .iter()
777 .enumerate()
778 .min_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))?;
779
780 Some(BasisNbasisCvResult {
781 optimal_nbasis: nbasis_range[best_idx],
782 scores,
783 nbasis_range: nbasis_range.to_vec(),
784 criterion,
785 })
786}
787
788#[cfg(test)]
789mod tests {
790 use super::*;
791 use crate::test_helpers::uniform_grid;
792 use std::f64::consts::PI;
793
794 #[test]
795 fn test_bspline_penalty_matrix_symmetric() {
796 let t = uniform_grid(101);
797 let penalty = bspline_penalty_matrix(&t, 15, 4, 2);
798 let _k = 15; let actual_k = (penalty.len() as f64).sqrt() as usize;
800 for i in 0..actual_k {
801 for j in 0..actual_k {
802 assert!(
803 (penalty[i + j * actual_k] - penalty[j + i * actual_k]).abs() < 1e-10,
804 "Penalty matrix not symmetric at ({}, {})",
805 i,
806 j
807 );
808 }
809 }
810 }
811
812 #[test]
813 fn test_bspline_penalty_matrix_positive_semidefinite() {
814 let t = uniform_grid(101);
815 let penalty = bspline_penalty_matrix(&t, 10, 4, 2);
816 let k = (penalty.len() as f64).sqrt() as usize;
817 for i in 0..k {
819 assert!(
820 penalty[i + i * k] >= -1e-10,
821 "Diagonal element {} is negative: {}",
822 i,
823 penalty[i + i * k]
824 );
825 }
826 }
827
828 #[test]
829 fn test_fourier_penalty_diagonal() {
830 let penalty = fourier_penalty_matrix(7, 1.0, 2);
831 for i in 0..7 {
833 for j in 0..7 {
834 if i != j {
835 assert!(
836 penalty[i + j * 7].abs() < 1e-10,
837 "Off-diagonal ({},{}) = {}",
838 i,
839 j,
840 penalty[i + j * 7]
841 );
842 }
843 }
844 }
845 assert!(penalty[0].abs() < 1e-10);
847 assert!(penalty[1 + 7] > 0.0);
849 assert!(penalty[3 + 3 * 7] > penalty[1 + 7]);
850 }
851
852 #[test]
853 fn test_smooth_basis_bspline() {
854 let m = 101;
855 let n = 5;
856 let t = uniform_grid(m);
857
858 let mut data = FdMatrix::zeros(n, m);
860 for i in 0..n {
861 for j in 0..m {
862 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.1 * (i as f64 * 0.3 + j as f64 * 0.01);
863 }
864 }
865
866 let nbasis = 15;
867 let penalty = bspline_penalty_matrix(&t, nbasis, 4, 2);
868 let _actual_k = (penalty.len() as f64).sqrt() as usize;
869
870 let fdpar = FdPar {
871 basis_type: BasisType::Bspline { order: 4 },
872 nbasis,
873 lambda: 1e-4,
874 lfd_order: 2,
875 penalty_matrix: penalty,
876 };
877
878 let result = smooth_basis(&data, &t, &fdpar);
879 assert!(result.is_ok(), "smooth_basis should succeed");
880
881 let res = result.unwrap();
882 assert_eq!(res.fitted.shape(), (n, m));
883 assert_eq!(res.coefficients.nrows(), n);
884 assert!(res.edf > 0.0, "EDF should be positive");
885 assert!(res.gcv > 0.0, "GCV should be positive");
886 }
887
888 #[test]
889 fn test_smooth_basis_fourier() {
890 let m = 101;
891 let n = 3;
892 let t = uniform_grid(m);
893
894 let mut data = FdMatrix::zeros(n, m);
895 for i in 0..n {
896 for j in 0..m {
897 data[(i, j)] = (2.0 * PI * t[j]).sin() + (4.0 * PI * t[j]).cos();
898 }
899 }
900
901 let nbasis = 7;
902 let period = 1.0;
903 let penalty = fourier_penalty_matrix(nbasis, period, 2);
904
905 let fdpar = FdPar {
906 basis_type: BasisType::Fourier { period },
907 nbasis,
908 lambda: 1e-6,
909 lfd_order: 2,
910 penalty_matrix: penalty,
911 };
912
913 let result = smooth_basis(&data, &t, &fdpar);
914 assert!(result.is_ok());
915
916 let res = result.unwrap();
917 for j in 0..m {
919 let expected = (2.0 * PI * t[j]).sin() + (4.0 * PI * t[j]).cos();
920 assert!(
921 (res.fitted[(0, j)] - expected).abs() < 0.1,
922 "Fourier fit poor at j={}: got {}, expected {}",
923 j,
924 res.fitted[(0, j)],
925 expected
926 );
927 }
928 }
929
930 #[test]
931 fn test_smooth_basis_gcv_selects_reasonable_lambda() {
932 let m = 101;
933 let n = 5;
934 let t = uniform_grid(m);
935
936 let mut data = FdMatrix::zeros(n, m);
937 for i in 0..n {
938 for j in 0..m {
939 data[(i, j)] =
940 (2.0 * PI * t[j]).sin() + 0.1 * ((i * 37 + j * 13) % 20) as f64 / 20.0;
941 }
942 }
943
944 let basis_type = BasisType::Bspline { order: 4 };
945 let result = smooth_basis_gcv(&data, &t, &basis_type, 15, 2, (-8.0, 4.0), 25);
946 assert!(result.is_some(), "GCV search should succeed");
947 }
948
949 #[test]
950 fn test_smooth_basis_large_lambda_reduces_edf() {
951 let m = 101;
952 let n = 3;
953 let t = uniform_grid(m);
954
955 let mut data = FdMatrix::zeros(n, m);
956 for i in 0..n {
957 for j in 0..m {
958 data[(i, j)] = (2.0 * PI * t[j]).sin();
959 }
960 }
961
962 let nbasis = 15;
963 let penalty = bspline_penalty_matrix(&t, nbasis, 4, 2);
964 let _actual_k = (penalty.len() as f64).sqrt() as usize;
965
966 let fdpar_small = FdPar {
967 basis_type: BasisType::Bspline { order: 4 },
968 nbasis,
969 lambda: 1e-8,
970 lfd_order: 2,
971 penalty_matrix: penalty.clone(),
972 };
973 let fdpar_large = FdPar {
974 basis_type: BasisType::Bspline { order: 4 },
975 nbasis,
976 lambda: 1e2,
977 lfd_order: 2,
978 penalty_matrix: penalty,
979 };
980
981 let res_small = smooth_basis(&data, &t, &fdpar_small).unwrap();
982 let res_large = smooth_basis(&data, &t, &fdpar_large).unwrap();
983
984 assert!(
985 res_large.edf < res_small.edf,
986 "Larger lambda should reduce EDF: {} vs {}",
987 res_large.edf,
988 res_small.edf
989 );
990 }
991
992 #[test]
995 fn test_basis_nbasis_cv_gcv() {
996 let m = 101;
997 let n = 5;
998 let t = uniform_grid(m);
999 let mut data = FdMatrix::zeros(n, m);
1000 for i in 0..n {
1001 for j in 0..m {
1002 data[(i, j)] =
1003 (2.0 * PI * t[j]).sin() + 0.1 * ((i * 37 + j * 13) % 20) as f64 / 20.0;
1004 }
1005 }
1006
1007 let nbasis_range: Vec<usize> = (4..=20).step_by(2).collect();
1008 let result = basis_nbasis_cv(
1009 &data,
1010 &t,
1011 &nbasis_range,
1012 &BasisType::Bspline { order: 4 },
1013 BasisCriterion::Gcv,
1014 5,
1015 1e-4,
1016 );
1017 assert!(result.is_some());
1018 let res = result.unwrap();
1019 assert!(nbasis_range.contains(&res.optimal_nbasis));
1020 assert_eq!(res.scores.len(), nbasis_range.len());
1021 assert_eq!(res.criterion, BasisCriterion::Gcv);
1022 }
1023
1024 #[test]
1025 fn test_basis_nbasis_cv_aic_bic() {
1026 let m = 51;
1027 let n = 5;
1028 let t = uniform_grid(m);
1029 let mut data = FdMatrix::zeros(n, m);
1030 for i in 0..n {
1031 for j in 0..m {
1032 data[(i, j)] = (2.0 * PI * t[j]).sin();
1033 }
1034 }
1035
1036 let nbasis_range: Vec<usize> = vec![5, 7, 9, 11];
1037 let aic_result = basis_nbasis_cv(
1038 &data,
1039 &t,
1040 &nbasis_range,
1041 &BasisType::Bspline { order: 4 },
1042 BasisCriterion::Aic,
1043 5,
1044 0.0,
1045 );
1046 let bic_result = basis_nbasis_cv(
1047 &data,
1048 &t,
1049 &nbasis_range,
1050 &BasisType::Bspline { order: 4 },
1051 BasisCriterion::Bic,
1052 5,
1053 0.0,
1054 );
1055 assert!(aic_result.is_some());
1056 assert!(bic_result.is_some());
1057 }
1058
1059 #[test]
1060 fn test_basis_nbasis_cv_kfold() {
1061 let m = 51;
1062 let n = 10;
1063 let t = uniform_grid(m);
1064 let mut data = FdMatrix::zeros(n, m);
1065 for i in 0..n {
1066 for j in 0..m {
1067 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.05 * ((i * 7 + j * 3) % 10) as f64;
1068 }
1069 }
1070
1071 let nbasis_range: Vec<usize> = vec![5, 7, 9];
1072 let result = basis_nbasis_cv(
1073 &data,
1074 &t,
1075 &nbasis_range,
1076 &BasisType::Bspline { order: 4 },
1077 BasisCriterion::Cv,
1078 5,
1079 1e-4,
1080 );
1081 assert!(result.is_some());
1082 let res = result.unwrap();
1083 assert!(nbasis_range.contains(&res.optimal_nbasis));
1084 assert_eq!(res.criterion, BasisCriterion::Cv);
1085 }
1086
1087 #[test]
1092 fn test_basis_nbasis_cv_penalizes_overfitting() {
1093 let m = 120;
1094 let n = 6;
1095 let t = uniform_grid(m);
1096 let mut data = FdMatrix::zeros(n, m);
1097 for i in 0..n {
1098 for j in 0..m {
1099 let noise = 0.2 * (((i * 31 + j * 17) % 13) as f64 / 13.0 - 0.5);
1101 data[(i, j)] = (2.0 * PI * t[j]).sin() + noise;
1102 }
1103 }
1104
1105 let nbasis_range: Vec<usize> = vec![5, 8, 12, 20, 30];
1106 let res = basis_nbasis_cv(
1107 &data,
1108 &t,
1109 &nbasis_range,
1110 &BasisType::Bspline { order: 4 },
1111 BasisCriterion::Cv,
1112 5,
1113 1e-6,
1114 )
1115 .unwrap();
1116
1117 assert_ne!(
1118 res.optimal_nbasis, 30,
1119 "CV must not always select the maximum n_basis (GH #33); scores={:?}",
1120 res.scores
1121 );
1122 let monotone_decreasing = res.scores.windows(2).all(|w| w[1] <= w[0] + 1e-12);
1123 assert!(
1124 !monotone_decreasing,
1125 "CV scores must not be monotone-decreasing in n_basis; scores={:?}",
1126 res.scores
1127 );
1128 }
1129
1130 fn make_test_data(n: usize, m: usize) -> (FdMatrix, Vec<f64>) {
1134 let t: Vec<f64> = (0..m).map(|i| i as f64 / (m - 1) as f64).collect();
1135 let mut data = FdMatrix::zeros(n, m);
1136 for i in 0..n {
1137 for j in 0..m {
1138 data[(i, j)] = (2.0 * PI * t[j]).sin()
1139 + 0.1 * (10.0 * t[j]).sin()
1140 + 0.05 * ((i * 37 + j * 13) % 20) as f64 / 20.0;
1141 }
1142 }
1143 (data, t)
1144 }
1145
1146 fn make_bspline_fdpar(argvals: &[f64], nbasis: usize, lambda: f64) -> FdPar {
1148 let penalty = bspline_penalty_matrix(argvals, nbasis, 4, 2);
1149 FdPar {
1150 basis_type: BasisType::Bspline { order: 4 },
1151 nbasis,
1152 lambda,
1153 lfd_order: 2,
1154 penalty_matrix: penalty,
1155 }
1156 }
1157
1158 fn make_fourier_fdpar(nbasis: usize, period: f64, lambda: f64) -> FdPar {
1160 let penalty = fourier_penalty_matrix(nbasis, period, 2);
1161 FdPar {
1162 basis_type: BasisType::Fourier { period },
1163 nbasis,
1164 lambda,
1165 lfd_order: 2,
1166 penalty_matrix: penalty,
1167 }
1168 }
1169
1170 #[test]
1173 fn test_basis_type_bspline_variant() {
1174 let bt = BasisType::Bspline { order: 4 };
1175 assert_eq!(bt, BasisType::Bspline { order: 4 });
1176 assert_ne!(bt, BasisType::Bspline { order: 3 });
1178 }
1179
1180 #[test]
1181 fn test_basis_type_fourier_variant() {
1182 let bt = BasisType::Fourier { period: 1.0 };
1183 assert_eq!(bt, BasisType::Fourier { period: 1.0 });
1184 assert_ne!(bt, BasisType::Fourier { period: 2.0 });
1185 }
1186
1187 #[test]
1188 fn test_basis_type_cross_variant_inequality() {
1189 let bspline = BasisType::Bspline { order: 4 };
1190 let fourier = BasisType::Fourier { period: 1.0 };
1191 assert_ne!(bspline, fourier);
1192 }
1193
1194 #[test]
1195 fn test_basis_type_clone_and_debug() {
1196 let bt = BasisType::Bspline { order: 4 };
1197 let cloned = bt.clone();
1198 assert_eq!(bt, cloned);
1199 let debug_str = format!("{:?}", bt);
1200 assert!(debug_str.contains("Bspline"));
1201 assert!(debug_str.contains("4"));
1202 }
1203
1204 #[test]
1207 fn test_fdpar_construction_and_fields() {
1208 let penalty = vec![1.0, 0.0, 0.0, 1.0];
1209 let fdpar = FdPar {
1210 basis_type: BasisType::Bspline { order: 4 },
1211 nbasis: 2,
1212 lambda: 0.01,
1213 lfd_order: 2,
1214 penalty_matrix: penalty.clone(),
1215 };
1216 assert_eq!(fdpar.nbasis, 2);
1217 assert!((fdpar.lambda - 0.01).abs() < 1e-15);
1218 assert_eq!(fdpar.lfd_order, 2);
1219 assert_eq!(fdpar.penalty_matrix.len(), 4);
1220 }
1221
1222 #[test]
1223 fn test_fdpar_clone_and_debug() {
1224 let t = uniform_grid(50);
1225 let fdpar = make_bspline_fdpar(&t, 8, 1e-3);
1226 let cloned = fdpar.clone();
1227 assert_eq!(fdpar, cloned);
1228 let debug_str = format!("{:?}", fdpar);
1229 assert!(debug_str.contains("FdPar"));
1230 }
1231
1232 #[test]
1235 fn test_basis_criterion_variants() {
1236 assert_eq!(BasisCriterion::Gcv, BasisCriterion::Gcv);
1237 assert_eq!(BasisCriterion::Cv, BasisCriterion::Cv);
1238 assert_eq!(BasisCriterion::Aic, BasisCriterion::Aic);
1239 assert_eq!(BasisCriterion::Bic, BasisCriterion::Bic);
1240 assert_ne!(BasisCriterion::Gcv, BasisCriterion::Aic);
1241 assert_ne!(BasisCriterion::Cv, BasisCriterion::Bic);
1242 }
1243
1244 #[test]
1245 fn test_basis_criterion_copy() {
1246 let c = BasisCriterion::Gcv;
1247 let copied = c; assert_eq!(c, copied);
1249 }
1250
1251 #[test]
1252 fn test_basis_criterion_debug() {
1253 let debug_str = format!("{:?}", BasisCriterion::Bic);
1254 assert!(debug_str.contains("Bic"));
1255 }
1256
1257 #[test]
1260 fn test_smooth_basis_result_all_fields() {
1261 let (data, t) = make_test_data(3, 50);
1262 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1263 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1264
1265 assert_eq!(res.coefficients.nrows(), 3);
1267 assert!(res.coefficients.ncols() > 0);
1268 assert_eq!(res.nbasis, res.coefficients.ncols());
1269 assert_eq!(res.fitted.shape(), (3, 50));
1271 assert!(res.edf > 0.0 && res.edf <= res.nbasis as f64);
1273 assert!(res.gcv.is_finite());
1275 assert!(res.aic.is_finite());
1276 assert!(res.bic.is_finite());
1277 let k = res.nbasis;
1279 assert_eq!(res.penalty_matrix.len(), k * k);
1280 }
1281
1282 #[test]
1283 fn test_smooth_basis_result_clone() {
1284 let (data, t) = make_test_data(2, 50);
1285 let fdpar = make_bspline_fdpar(&t, 8, 1e-3);
1286 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1287 let cloned = res.clone();
1288 assert_eq!(res, cloned);
1289 }
1290
1291 #[test]
1294 fn test_smooth_basis_bspline_coefficient_shape() {
1295 let (data, t) = make_test_data(4, 50);
1296 let nbasis = 12;
1297 let fdpar = make_bspline_fdpar(&t, nbasis, 1e-4);
1298 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1299 assert_eq!(res.coefficients.nrows(), 4);
1300 assert!(res.coefficients.ncols() >= 2);
1302 assert_eq!(res.nbasis, res.coefficients.ncols());
1303 }
1304
1305 #[test]
1306 fn test_smooth_basis_bspline_fitted_values_shape() {
1307 let m = 80;
1308 let n = 6;
1309 let (data, t) = make_test_data(n, m);
1310 let fdpar = make_bspline_fdpar(&t, 15, 1e-4);
1311 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1312 assert_eq!(res.fitted.shape(), (n, m));
1313 }
1314
1315 #[test]
1316 fn test_smooth_basis_bspline_zero_lambda_interpolates() {
1317 let m = 30;
1319 let n = 2;
1320 let (data, t) = make_test_data(n, m);
1321 let fdpar = make_bspline_fdpar(&t, 15, 0.0);
1322 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1323
1324 let mut max_resid = 0.0_f64;
1326 for i in 0..n {
1327 for j in 0..m {
1328 let resid = (data[(i, j)] - res.fitted[(i, j)]).abs();
1329 max_resid = max_resid.max(resid);
1330 }
1331 }
1332 assert!(
1333 max_resid < 0.5,
1334 "Zero-lambda B-spline should closely interpolate; max_resid = {}",
1335 max_resid
1336 );
1337 }
1338
1339 #[test]
1340 fn test_smooth_basis_bspline_large_lambda_oversmooths() {
1341 let m = 50;
1344 let n = 1;
1345 let (data, t) = make_test_data(n, m);
1346
1347 let fdpar_small = make_bspline_fdpar(&t, 15, 1e-6);
1348 let res_small = smooth_basis(&data, &t, &fdpar_small).unwrap();
1349
1350 let fdpar_large = make_bspline_fdpar(&t, 15, 1e6);
1351 let res_large = smooth_basis(&data, &t, &fdpar_large).unwrap();
1352
1353 let compute_variance = |fitted: &FdMatrix, row: usize, ncols: usize| -> f64 {
1354 let vals: Vec<f64> = (0..ncols).map(|j| fitted[(row, j)]).collect();
1355 let mean = vals.iter().sum::<f64>() / ncols as f64;
1356 vals.iter().map(|&v| (v - mean).powi(2)).sum::<f64>() / ncols as f64
1357 };
1358
1359 let var_small = compute_variance(&res_small.fitted, 0, m);
1360 let var_large = compute_variance(&res_large.fitted, 0, m);
1361 assert!(
1362 var_large < var_small,
1363 "Large lambda should yield lower variance fit: var_large={}, var_small={}",
1364 var_large,
1365 var_small
1366 );
1367 }
1368
1369 #[test]
1370 fn test_smooth_basis_bspline_penalty_effect_on_smoothness() {
1371 let m = 50;
1373 let n = 1;
1374 let (data, t) = make_test_data(n, m);
1375
1376 let fdpar_small = make_bspline_fdpar(&t, 15, 1e-8);
1377 let fdpar_large = make_bspline_fdpar(&t, 15, 1.0);
1378
1379 let res_small = smooth_basis(&data, &t, &fdpar_small).unwrap();
1380 let res_large = smooth_basis(&data, &t, &fdpar_large).unwrap();
1381
1382 let roughness = |fitted: &FdMatrix, row: usize, ncols: usize| -> f64 {
1384 (1..ncols - 1)
1385 .map(|j| {
1386 let d2 = fitted[(row, j + 1)] - 2.0 * fitted[(row, j)] + fitted[(row, j - 1)];
1387 d2 * d2
1388 })
1389 .sum::<f64>()
1390 };
1391
1392 let r_small = roughness(&res_small.fitted, 0, m);
1393 let r_large = roughness(&res_large.fitted, 0, m);
1394 assert!(
1395 r_large < r_small,
1396 "Larger lambda should produce smoother fit: roughness_large={}, roughness_small={}",
1397 r_large,
1398 r_small
1399 );
1400 }
1401
1402 #[test]
1403 fn test_smooth_basis_bspline_single_curve() {
1404 let m = 50;
1405 let (data, t) = make_test_data(1, m);
1406 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1407 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1408 assert_eq!(res.fitted.nrows(), 1);
1409 assert_eq!(res.fitted.ncols(), m);
1410 assert!(res.gcv.is_finite());
1411 }
1412
1413 #[test]
1414 fn test_smooth_basis_bspline_many_curves() {
1415 let m = 50;
1416 let n = 20;
1417 let (data, t) = make_test_data(n, m);
1418 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1419 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1420 assert_eq!(res.fitted.nrows(), n);
1421 assert_eq!(res.coefficients.nrows(), n);
1422 }
1423
1424 #[test]
1425 fn test_smooth_basis_bspline_minimal_nbasis() {
1426 let m = 50;
1428 let (data, t) = make_test_data(1, m);
1429 let fdpar = make_bspline_fdpar(&t, 2, 1e-4);
1430 let res = smooth_basis(&data, &t, &fdpar);
1431 assert!(res.is_ok());
1433 }
1434
1435 #[test]
1436 fn test_smooth_basis_bspline_different_orders() {
1437 let m = 50;
1438 let (data, t) = make_test_data(2, m);
1439 let penalty3 = bspline_penalty_matrix(&t, 10, 3, 2);
1441 let fdpar3 = FdPar {
1442 basis_type: BasisType::Bspline { order: 3 },
1443 nbasis: 10,
1444 lambda: 1e-4,
1445 lfd_order: 2,
1446 penalty_matrix: penalty3,
1447 };
1448 let res3 = smooth_basis(&data, &t, &fdpar3);
1449 assert!(res3.is_ok());
1450
1451 let penalty5 = bspline_penalty_matrix(&t, 10, 5, 2);
1453 let fdpar5 = FdPar {
1454 basis_type: BasisType::Bspline { order: 5 },
1455 nbasis: 10,
1456 lambda: 1e-4,
1457 lfd_order: 2,
1458 penalty_matrix: penalty5,
1459 };
1460 let res5 = smooth_basis(&data, &t, &fdpar5);
1461 assert!(res5.is_ok());
1462 }
1463
1464 #[test]
1467 fn test_smooth_basis_fourier_coefficient_shape() {
1468 let m = 50;
1469 let n = 3;
1470 let t = uniform_grid(m);
1471 let mut data = FdMatrix::zeros(n, m);
1472 for i in 0..n {
1473 for j in 0..m {
1474 data[(i, j)] = (2.0 * PI * t[j]).sin();
1475 }
1476 }
1477 let nbasis = 7;
1478 let fdpar = make_fourier_fdpar(nbasis, 1.0, 1e-6);
1479 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1480 assert_eq!(res.coefficients.nrows(), n);
1481 assert_eq!(res.coefficients.ncols(), nbasis);
1482 assert_eq!(res.nbasis, nbasis);
1483 }
1484
1485 #[test]
1486 fn test_smooth_basis_fourier_fits_pure_sine() {
1487 let m = 100;
1489 let t = uniform_grid(m);
1490 let mut data = FdMatrix::zeros(1, m);
1491 for j in 0..m {
1492 data[(0, j)] = (2.0 * PI * t[j]).sin();
1493 }
1494 let fdpar = make_fourier_fdpar(5, 1.0, 1e-8);
1495 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1496
1497 for j in 0..m {
1498 let expected = (2.0 * PI * t[j]).sin();
1499 assert!(
1500 (res.fitted[(0, j)] - expected).abs() < 0.05,
1501 "Fourier should fit pure sine; j={}, got={}, expected={}",
1502 j,
1503 res.fitted[(0, j)],
1504 expected
1505 );
1506 }
1507 }
1508
1509 #[test]
1510 fn test_smooth_basis_fourier_different_periods() {
1511 let m = 50;
1512 let t = uniform_grid(m);
1513 let mut data = FdMatrix::zeros(1, m);
1514 for j in 0..m {
1515 data[(0, j)] = (2.0 * PI * t[j]).sin();
1516 }
1517
1518 let fdpar1 = make_fourier_fdpar(7, 1.0, 1e-6);
1520 let res1 = smooth_basis(&data, &t, &fdpar1).unwrap();
1521
1522 let fdpar2 = make_fourier_fdpar(7, 2.0, 1e-6);
1524 let res2 = smooth_basis(&data, &t, &fdpar2).unwrap();
1525
1526 assert_eq!(res1.fitted.shape(), (1, m));
1528 assert_eq!(res2.fitted.shape(), (1, m));
1529 }
1530
1531 #[test]
1532 fn test_smooth_basis_fourier_zero_lambda() {
1533 let m = 50;
1534 let t = uniform_grid(m);
1535 let mut data = FdMatrix::zeros(1, m);
1536 for j in 0..m {
1537 data[(0, j)] = (2.0 * PI * t[j]).sin() + (4.0 * PI * t[j]).cos();
1538 }
1539 let fdpar = make_fourier_fdpar(9, 1.0, 0.0);
1540 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1541 assert_eq!(res.fitted.shape(), (1, m));
1542 assert!(res.edf > 1.0);
1544 }
1545
1546 #[test]
1547 fn test_smooth_basis_fourier_large_lambda() {
1548 let m = 50;
1549 let t = uniform_grid(m);
1550 let mut data = FdMatrix::zeros(1, m);
1551 for j in 0..m {
1552 data[(0, j)] = (2.0 * PI * t[j]).sin();
1553 }
1554 let fdpar = make_fourier_fdpar(9, 1.0, 1e6);
1555 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1556 assert!(
1558 res.edf < 5.0,
1559 "Large lambda should reduce EDF; edf={}",
1560 res.edf
1561 );
1562 }
1563
1564 #[test]
1567 fn test_smooth_basis_lambda_gradient_edf() {
1568 let m = 50;
1570 let (data, t) = make_test_data(3, m);
1571 let lambdas = [1e-8, 1e-4, 1e-2, 1.0, 1e2];
1572 let mut prev_edf = f64::INFINITY;
1573 for &lam in &lambdas {
1574 let fdpar = make_bspline_fdpar(&t, 12, lam);
1575 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1576 assert!(
1577 res.edf <= prev_edf + 0.01,
1578 "EDF should decrease: lambda={}, edf={}, prev_edf={}",
1579 lam,
1580 res.edf,
1581 prev_edf
1582 );
1583 prev_edf = res.edf;
1584 }
1585 }
1586
1587 #[test]
1588 fn test_smooth_basis_lambda_gradient_rss() {
1589 let m = 50;
1591 let n = 2;
1592 let (data, t) = make_test_data(n, m);
1593 let lambdas = [0.0, 1e-6, 1e-2, 1.0, 1e4];
1594 let mut prev_rss = -1.0;
1595 for &lam in &lambdas {
1596 let fdpar = make_bspline_fdpar(&t, 12, lam);
1597 let res = smooth_basis(&data, &t, &fdpar).unwrap();
1598 let mut rss = 0.0;
1599 for i in 0..n {
1600 for j in 0..m {
1601 rss += (data[(i, j)] - res.fitted[(i, j)]).powi(2);
1602 }
1603 }
1604 assert!(
1605 rss >= prev_rss - 1e-8,
1606 "RSS should increase: lambda={}, rss={}, prev_rss={}",
1607 lam,
1608 rss,
1609 prev_rss
1610 );
1611 prev_rss = rss;
1612 }
1613 }
1614
1615 #[test]
1618 fn test_smooth_basis_empty_data_rows() {
1619 let t = uniform_grid(50);
1620 let data = FdMatrix::zeros(0, 50);
1621 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1622 let res = smooth_basis(&data, &t, &fdpar);
1623 assert!(res.is_err());
1624 }
1625
1626 #[test]
1627 fn test_smooth_basis_empty_data_cols() {
1628 let data = FdMatrix::zeros(5, 0);
1629 let fdpar = FdPar {
1630 basis_type: BasisType::Bspline { order: 4 },
1631 nbasis: 10,
1632 lambda: 1e-4,
1633 lfd_order: 2,
1634 penalty_matrix: vec![0.0; 100],
1635 };
1636 let res = smooth_basis(&data, &[], &fdpar);
1637 assert!(res.is_err());
1638 }
1639
1640 #[test]
1641 fn test_smooth_basis_mismatched_argvals() {
1642 let t = uniform_grid(50);
1643 let data = FdMatrix::zeros(3, 40); let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1645 let res = smooth_basis(&data, &t, &fdpar);
1646 assert!(res.is_err());
1647 }
1648
1649 #[test]
1650 fn test_smooth_basis_nbasis_too_small() {
1651 let t = uniform_grid(50);
1652 let data = FdMatrix::zeros(3, 50);
1653 let fdpar = FdPar {
1655 basis_type: BasisType::Bspline { order: 4 },
1656 nbasis: 1,
1657 lambda: 1e-4,
1658 lfd_order: 2,
1659 penalty_matrix: vec![0.0; 1],
1660 };
1661 let res = smooth_basis(&data, &t, &fdpar);
1662 assert!(res.is_err());
1663 }
1664
1665 #[test]
1666 fn test_smooth_basis_error_is_invalid_dimension() {
1667 let t = uniform_grid(50);
1668 let data = FdMatrix::zeros(0, 50);
1669 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1670 let err = smooth_basis(&data, &t, &fdpar).unwrap_err();
1671 match err {
1672 crate::FdarError::InvalidDimension { .. } => {} other => panic!("Expected InvalidDimension, got {:?}", other),
1674 }
1675 }
1676
1677 #[test]
1680 fn test_bspline_penalty_matrix_different_orders() {
1681 let t = uniform_grid(101);
1682 let p1 = bspline_penalty_matrix(&t, 10, 4, 1);
1684 let p2 = bspline_penalty_matrix(&t, 10, 4, 2);
1686 assert_eq!(p1.len(), p2.len());
1688 let diff: f64 = p1.iter().zip(p2.iter()).map(|(a, b)| (a - b).abs()).sum();
1690 assert!(
1691 diff > 1e-10,
1692 "Different lfd_orders should produce different penalties"
1693 );
1694 }
1695
1696 #[test]
1697 fn test_bspline_penalty_matrix_edge_cases() {
1698 let t = vec![0.0];
1700 let p = bspline_penalty_matrix(&t, 10, 4, 2);
1701 assert!(p.iter().all(|&v| v == 0.0));
1703
1704 let t2 = uniform_grid(50);
1706 let p2 = bspline_penalty_matrix(&t2, 1, 4, 2);
1707 assert!(p2.iter().all(|&v| v == 0.0));
1708
1709 let p3 = bspline_penalty_matrix(&t2, 10, 4, 4);
1711 assert!(p3.iter().all(|&v| v == 0.0));
1712 }
1713
1714 #[test]
1715 fn test_bspline_penalty_nonnegative_diagonal() {
1716 let t = uniform_grid(101);
1717 for nbasis in [5, 10, 20] {
1718 let p = bspline_penalty_matrix(&t, nbasis, 4, 2);
1719 let k = (p.len() as f64).sqrt() as usize;
1720 for i in 0..k {
1721 assert!(
1722 p[i + i * k] >= -1e-10,
1723 "Diagonal ({},{}) negative for nbasis={}: {}",
1724 i,
1725 i,
1726 nbasis,
1727 p[i + i * k]
1728 );
1729 }
1730 }
1731 }
1732
1733 #[test]
1734 fn test_fourier_penalty_increasing_with_frequency() {
1735 let penalty = fourier_penalty_matrix(11, 1.0, 2);
1736 let k = 11;
1737 assert!(penalty[0].abs() < 1e-15);
1739 let mut prev_eigenval = 0.0;
1741 for freq in 1..=5 {
1742 let idx_sin = 2 * freq - 1;
1743 let eigenval = penalty[idx_sin + idx_sin * k];
1744 assert!(
1745 eigenval > prev_eigenval,
1746 "Higher frequency should have larger penalty: freq={}, eigenval={}, prev={}",
1747 freq,
1748 eigenval,
1749 prev_eigenval
1750 );
1751 prev_eigenval = eigenval;
1752 let idx_cos = 2 * freq;
1754 if idx_cos < k {
1755 assert!(
1756 (penalty[idx_cos + idx_cos * k] - eigenval).abs() < 1e-10,
1757 "Sin and cos penalty should match at freq {}",
1758 freq
1759 );
1760 }
1761 }
1762 }
1763
1764 #[test]
1765 fn test_fourier_penalty_different_periods() {
1766 let p1 = fourier_penalty_matrix(7, 1.0, 2);
1767 let p2 = fourier_penalty_matrix(7, 2.0, 2);
1768 for i in 1..7 {
1770 assert!(
1771 p2[i + i * 7] < p1[i + i * 7] || (p1[i + i * 7] == 0.0 && p2[i + i * 7] == 0.0),
1772 "Longer period should have smaller penalties at i={}",
1773 i
1774 );
1775 }
1776 }
1777
1778 #[test]
1779 fn test_fourier_penalty_first_order() {
1780 let p = fourier_penalty_matrix(5, 1.0, 1);
1782 let omega1 = 2.0 * PI;
1784 let expected1 = omega1.powi(2);
1785 assert!(
1786 (p[1 + 5] - expected1).abs() < 1e-6,
1787 "First-order penalty eigenval: got {}, expected {}",
1788 p[1 + 5],
1789 expected1
1790 );
1791 }
1792
1793 #[test]
1794 fn test_fourier_penalty_zero_nbasis() {
1795 let p = fourier_penalty_matrix(0, 1.0, 2);
1796 assert!(p.is_empty());
1797 }
1798
1799 #[test]
1800 fn test_fourier_penalty_nbasis_one() {
1801 let p = fourier_penalty_matrix(1, 1.0, 2);
1802 assert_eq!(p.len(), 1);
1803 assert!(p[0].abs() < 1e-15); }
1805
1806 #[test]
1809 fn test_smooth_basis_gcv_returns_valid_result() {
1810 let (data, t) = make_test_data(5, 50);
1811 let bt = BasisType::Bspline { order: 4 };
1812 let result = smooth_basis_gcv(&data, &t, &bt, 12, 2, (-6.0, 2.0), 20);
1813 assert!(result.is_some());
1814 let res = result.unwrap();
1815 assert_eq!(res.fitted.shape(), (5, 50));
1816 assert!(res.gcv.is_finite());
1817 assert!(res.edf > 0.0);
1818 }
1819
1820 #[test]
1821 fn test_smooth_basis_gcv_fourier() {
1822 let m = 80;
1823 let t = uniform_grid(m);
1824 let mut data = FdMatrix::zeros(3, m);
1825 for i in 0..3 {
1826 for j in 0..m {
1827 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.5 * (4.0 * PI * t[j]).cos();
1828 }
1829 }
1830 let bt = BasisType::Fourier { period: 1.0 };
1831 let result = smooth_basis_gcv(&data, &t, &bt, 9, 2, (-8.0, 4.0), 25);
1832 assert!(result.is_some());
1833 let res = result.unwrap();
1834 assert_eq!(res.fitted.nrows(), 3);
1835 assert_eq!(res.nbasis, 9);
1836 }
1837
1838 #[test]
1839 fn test_smooth_basis_gcv_selects_finite_gcv() {
1840 let (data, t) = make_test_data(5, 60);
1841 let bt = BasisType::Bspline { order: 4 };
1842 let res = smooth_basis_gcv(&data, &t, &bt, 12, 2, (-6.0, 2.0), 15).unwrap();
1843 assert!(res.gcv.is_finite());
1844 assert!(res.gcv > 0.0);
1845 }
1846
1847 #[test]
1848 fn test_smooth_basis_gcv_empty_data() {
1849 let data = FdMatrix::zeros(0, 50);
1850 let t = uniform_grid(50);
1851 let bt = BasisType::Bspline { order: 4 };
1852 let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-6.0, 2.0), 10);
1853 assert!(result.is_none());
1855 }
1856
1857 #[test]
1858 fn test_smooth_basis_gcv_empty_argvals() {
1859 let data = FdMatrix::zeros(5, 0);
1860 let bt = BasisType::Bspline { order: 4 };
1861 let result = smooth_basis_gcv(&data, &[], &bt, 10, 2, (-6.0, 2.0), 10);
1862 assert!(result.is_none());
1863 }
1864
1865 #[test]
1866 fn test_smooth_basis_gcv_nbasis_too_small() {
1867 let (data, t) = make_test_data(5, 50);
1868 let bt = BasisType::Bspline { order: 4 };
1869 let result = smooth_basis_gcv(&data, &t, &bt, 1, 2, (-6.0, 2.0), 10);
1870 assert!(result.is_none());
1871 }
1872
1873 #[test]
1874 fn test_smooth_basis_gcv_ngrid_too_small() {
1875 let (data, t) = make_test_data(5, 50);
1876 let bt = BasisType::Bspline { order: 4 };
1877 let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-6.0, 2.0), 1);
1878 assert!(result.is_none());
1879 }
1880
1881 #[test]
1882 fn test_smooth_basis_gcv_narrow_range() {
1883 let (data, t) = make_test_data(3, 50);
1884 let bt = BasisType::Bspline { order: 4 };
1885 let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-3.0, -2.0), 5);
1887 assert!(result.is_some());
1888 }
1889
1890 #[test]
1891 fn test_smooth_basis_gcv_wide_range() {
1892 let (data, t) = make_test_data(3, 50);
1893 let bt = BasisType::Bspline { order: 4 };
1894 let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-12.0, 8.0), 30);
1896 assert!(result.is_some());
1897 }
1898
1899 #[test]
1902 fn test_basis_nbasis_cv_scores_length() {
1903 let (data, t) = make_test_data(5, 50);
1904 let nbasis_range: Vec<usize> = vec![4, 6, 8, 10, 12];
1905 let res = basis_nbasis_cv(
1906 &data,
1907 &t,
1908 &nbasis_range,
1909 &BasisType::Bspline { order: 4 },
1910 BasisCriterion::Gcv,
1911 5,
1912 1e-4,
1913 )
1914 .unwrap();
1915 assert_eq!(res.scores.len(), 5);
1916 assert_eq!(res.nbasis_range.len(), 5);
1917 assert_eq!(res.nbasis_range, nbasis_range);
1918 }
1919
1920 #[test]
1921 fn test_basis_nbasis_cv_optimal_within_range() {
1922 let (data, t) = make_test_data(8, 50);
1923 let nbasis_range: Vec<usize> = vec![5, 7, 9, 11, 13, 15];
1924 for criterion in [
1925 BasisCriterion::Gcv,
1926 BasisCriterion::Aic,
1927 BasisCriterion::Bic,
1928 ] {
1929 let res = basis_nbasis_cv(
1930 &data,
1931 &t,
1932 &nbasis_range,
1933 &BasisType::Bspline { order: 4 },
1934 criterion,
1935 5,
1936 1e-4,
1937 )
1938 .unwrap();
1939 assert!(
1940 nbasis_range.contains(&res.optimal_nbasis),
1941 "optimal_nbasis {} not in range for {:?}",
1942 res.optimal_nbasis,
1943 criterion
1944 );
1945 }
1946 }
1947
1948 #[test]
1949 fn test_basis_nbasis_cv_fourier_gcv() {
1950 let m = 80;
1951 let t = uniform_grid(m);
1952 let mut data = FdMatrix::zeros(5, m);
1953 for i in 0..5 {
1954 for j in 0..m {
1955 data[(i, j)] = (2.0 * PI * t[j]).sin()
1956 + 0.3 * (4.0 * PI * t[j]).cos()
1957 + 0.02 * ((i * 7 + j * 3) % 10) as f64;
1958 }
1959 }
1960 let nbasis_range: Vec<usize> = vec![5, 7, 9, 11];
1961 let res = basis_nbasis_cv(
1962 &data,
1963 &t,
1964 &nbasis_range,
1965 &BasisType::Fourier { period: 1.0 },
1966 BasisCriterion::Gcv,
1967 5,
1968 1e-4,
1969 )
1970 .unwrap();
1971 assert!(nbasis_range.contains(&res.optimal_nbasis));
1972 }
1973
1974 #[test]
1975 fn test_basis_nbasis_cv_fourier_cv() {
1976 let m = 60;
1977 let t = uniform_grid(m);
1978 let n = 10;
1979 let mut data = FdMatrix::zeros(n, m);
1980 for i in 0..n {
1981 for j in 0..m {
1982 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.02 * ((i * 11 + j) % 15) as f64;
1983 }
1984 }
1985 let nbasis_range: Vec<usize> = vec![5, 7, 9];
1986 let res = basis_nbasis_cv(
1987 &data,
1988 &t,
1989 &nbasis_range,
1990 &BasisType::Fourier { period: 1.0 },
1991 BasisCriterion::Cv,
1992 5,
1993 1e-4,
1994 )
1995 .unwrap();
1996 assert!(nbasis_range.contains(&res.optimal_nbasis));
1997 assert_eq!(res.criterion, BasisCriterion::Cv);
1998 }
1999
2000 #[test]
2001 fn test_basis_nbasis_cv_with_nbasis_below_minimum() {
2002 let (data, t) = make_test_data(5, 50);
2004 let nbasis_range: Vec<usize> = vec![1, 5, 10];
2005 let res = basis_nbasis_cv(
2006 &data,
2007 &t,
2008 &nbasis_range,
2009 &BasisType::Bspline { order: 4 },
2010 BasisCriterion::Gcv,
2011 5,
2012 1e-4,
2013 )
2014 .unwrap();
2015 assert!(
2017 res.optimal_nbasis >= 5,
2018 "Should skip invalid nbasis=1, got optimal={}",
2019 res.optimal_nbasis
2020 );
2021 assert!(res.scores[0].is_infinite());
2022 }
2023
2024 #[test]
2025 fn test_basis_nbasis_cv_empty_range() {
2026 let (data, t) = make_test_data(5, 50);
2027 let nbasis_range: Vec<usize> = vec![];
2028 let result = basis_nbasis_cv(
2029 &data,
2030 &t,
2031 &nbasis_range,
2032 &BasisType::Bspline { order: 4 },
2033 BasisCriterion::Gcv,
2034 5,
2035 1e-4,
2036 );
2037 assert!(result.is_none());
2038 }
2039
2040 #[test]
2041 fn test_basis_nbasis_cv_empty_data() {
2042 let data = FdMatrix::zeros(0, 50);
2043 let t = uniform_grid(50);
2044 let nbasis_range: Vec<usize> = vec![5, 10];
2045 let result = basis_nbasis_cv(
2046 &data,
2047 &t,
2048 &nbasis_range,
2049 &BasisType::Bspline { order: 4 },
2050 BasisCriterion::Gcv,
2051 5,
2052 1e-4,
2053 );
2054 assert!(result.is_none());
2055 }
2056
2057 #[test]
2058 fn test_basis_nbasis_cv_mismatched_argvals() {
2059 let data = FdMatrix::zeros(5, 50);
2060 let t = uniform_grid(40); let nbasis_range: Vec<usize> = vec![5, 10];
2062 let result = basis_nbasis_cv(
2063 &data,
2064 &t,
2065 &nbasis_range,
2066 &BasisType::Bspline { order: 4 },
2067 BasisCriterion::Gcv,
2068 5,
2069 1e-4,
2070 );
2071 assert!(result.is_none());
2072 }
2073
2074 #[test]
2075 fn test_basis_nbasis_cv_single_nbasis() {
2076 let (data, t) = make_test_data(5, 50);
2077 let nbasis_range: Vec<usize> = vec![10];
2078 let res = basis_nbasis_cv(
2079 &data,
2080 &t,
2081 &nbasis_range,
2082 &BasisType::Bspline { order: 4 },
2083 BasisCriterion::Gcv,
2084 5,
2085 1e-4,
2086 )
2087 .unwrap();
2088 assert_eq!(res.optimal_nbasis, 10);
2089 assert_eq!(res.scores.len(), 1);
2090 }
2091
2092 #[test]
2093 fn test_basis_nbasis_cv_bic_penalizes_more_than_aic() {
2094 let (data, t) = make_test_data(5, 80);
2097 let nbasis_range: Vec<usize> = (4..=20).step_by(2).collect();
2098
2099 let aic_res = basis_nbasis_cv(
2100 &data,
2101 &t,
2102 &nbasis_range,
2103 &BasisType::Bspline { order: 4 },
2104 BasisCriterion::Aic,
2105 5,
2106 1e-4,
2107 )
2108 .unwrap();
2109 let bic_res = basis_nbasis_cv(
2110 &data,
2111 &t,
2112 &nbasis_range,
2113 &BasisType::Bspline { order: 4 },
2114 BasisCriterion::Bic,
2115 5,
2116 1e-4,
2117 )
2118 .unwrap();
2119 assert!(
2122 bic_res.optimal_nbasis <= aic_res.optimal_nbasis + 4,
2123 "BIC selected {} vs AIC selected {} -- BIC should not select much more than AIC",
2124 bic_res.optimal_nbasis,
2125 aic_res.optimal_nbasis
2126 );
2127 }
2128
2129 #[test]
2132 fn test_smooth_basis_fitted_close_to_data() {
2133 let m = 50;
2135 let n = 3;
2136 let t = uniform_grid(m);
2137 let mut data = FdMatrix::zeros(n, m);
2138 for i in 0..n {
2139 for j in 0..m {
2140 data[(i, j)] = (2.0 * PI * t[j]).sin();
2141 }
2142 }
2143 let fdpar = make_bspline_fdpar(&t, 15, 1e-6);
2144 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2145
2146 let mut max_err = 0.0_f64;
2147 for i in 0..n {
2148 for j in 0..m {
2149 let err = (data[(i, j)] - res.fitted[(i, j)]).abs();
2150 max_err = max_err.max(err);
2151 }
2152 }
2153 assert!(
2154 max_err < 0.1,
2155 "Fitted should be close to smooth data; max_err={}",
2156 max_err
2157 );
2158 }
2159
2160 #[test]
2161 fn test_smooth_basis_constant_data() {
2162 let m = 50;
2164 let n = 2;
2165 let t = uniform_grid(m);
2166 let mut data = FdMatrix::zeros(n, m);
2167 for i in 0..n {
2168 for j in 0..m {
2169 data[(i, j)] = 3.15;
2170 }
2171 }
2172 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2173 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2174 for i in 0..n {
2175 for j in 0..m {
2176 assert!(
2177 (res.fitted[(i, j)] - 3.15).abs() < 0.01,
2178 "Constant data should be fit well at ({},{}): got {}",
2179 i,
2180 j,
2181 res.fitted[(i, j)]
2182 );
2183 }
2184 }
2185 }
2186
2187 #[test]
2188 fn test_smooth_basis_linear_data() {
2189 let m = 50;
2191 let t = uniform_grid(m);
2192 let mut data = FdMatrix::zeros(1, m);
2193 for j in 0..m {
2194 data[(0, j)] = 2.0 * t[j] + 1.0;
2195 }
2196 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2197 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2198 for j in 0..m {
2199 let expected = 2.0 * t[j] + 1.0;
2200 assert!(
2201 (res.fitted[(0, j)] - expected).abs() < 0.05,
2202 "Linear data should be fit well at j={}: got {}, expected {}",
2203 j,
2204 res.fitted[(0, j)],
2205 expected
2206 );
2207 }
2208 }
2209
2210 #[test]
2213 fn test_smooth_basis_edf_bounded() {
2214 let m = 50;
2215 let (data, t) = make_test_data(3, m);
2216 let fdpar = make_bspline_fdpar(&t, 12, 1e-4);
2217 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2218 assert!(
2220 res.edf > 0.0 && res.edf <= m as f64,
2221 "EDF should be in (0, {}]; got {}",
2222 m,
2223 res.edf
2224 );
2225 }
2226
2227 #[test]
2228 fn test_smooth_basis_gcv_aic_bic_all_finite() {
2229 let (data, t) = make_test_data(4, 60);
2230 let fdpar = make_bspline_fdpar(&t, 12, 1e-3);
2231 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2232 assert!(res.gcv.is_finite(), "GCV should be finite: {}", res.gcv);
2233 assert!(res.aic.is_finite(), "AIC should be finite: {}", res.aic);
2234 assert!(res.bic.is_finite(), "BIC should be finite: {}", res.bic);
2235 }
2236
2237 #[test]
2240 fn test_smooth_basis_penalty_matrix_in_result() {
2241 let (data, t) = make_test_data(3, 50);
2242 let nbasis = 10;
2243 let fdpar = make_bspline_fdpar(&t, nbasis, 1e-4);
2244 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2245 let k = res.nbasis;
2246 assert_eq!(
2247 res.penalty_matrix.len(),
2248 k * k,
2249 "Penalty matrix should be k*k = {}*{} = {}; got {}",
2250 k,
2251 k,
2252 k * k,
2253 res.penalty_matrix.len()
2254 );
2255 }
2256
2257 #[test]
2260 fn test_smooth_basis_identical_curves_same_coefficients() {
2261 let m = 50;
2262 let t = uniform_grid(m);
2263 let curve: Vec<f64> = (0..m).map(|j| (2.0 * PI * t[j]).sin()).collect();
2264 let n = 4;
2265 let mut data = FdMatrix::zeros(n, m);
2266 for i in 0..n {
2267 for j in 0..m {
2268 data[(i, j)] = curve[j];
2269 }
2270 }
2271 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2272 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2273
2274 let k = res.coefficients.ncols();
2276 for i in 1..n {
2277 for j in 0..k {
2278 assert!(
2279 (res.coefficients[(i, j)] - res.coefficients[(0, j)]).abs() < 1e-10,
2280 "Identical curves should have identical coefficients: curve {} col {} differs",
2281 i,
2282 j
2283 );
2284 }
2285 }
2286 }
2287
2288 #[test]
2291 fn test_basis_nbasis_cv_different_nfolds() {
2292 let (data, t) = make_test_data(12, 50);
2293 let nbasis_range: Vec<usize> = vec![5, 8, 11];
2294 for nfolds in [2, 3, 5, 10] {
2295 let res = basis_nbasis_cv(
2296 &data,
2297 &t,
2298 &nbasis_range,
2299 &BasisType::Bspline { order: 4 },
2300 BasisCriterion::Cv,
2301 nfolds,
2302 1e-4,
2303 );
2304 assert!(res.is_some(), "CV should succeed with nfolds={}", nfolds);
2305 let r = res.unwrap();
2306 assert!(nbasis_range.contains(&r.optimal_nbasis));
2307 }
2308 }
2309
2310 #[test]
2313 fn test_smooth_basis_many_basis_functions() {
2314 let m = 100;
2315 let (data, t) = make_test_data(2, m);
2316 let fdpar = make_bspline_fdpar(&t, 40, 1e-2);
2318 let res = smooth_basis(&data, &t, &fdpar);
2319 assert!(
2320 res.is_ok(),
2321 "Should handle many basis functions with penalty"
2322 );
2323 }
2324
2325 #[test]
2328 fn test_smooth_basis_bspline_vs_fourier_different_results() {
2329 let m = 50;
2330 let (data, t) = make_test_data(2, m);
2331 let fdpar_bs = make_bspline_fdpar(&t, 9, 1e-4);
2332 let fdpar_f = make_fourier_fdpar(9, 1.0, 1e-4);
2333 let res_bs = smooth_basis(&data, &t, &fdpar_bs).unwrap();
2334 let res_f = smooth_basis(&data, &t, &fdpar_f).unwrap();
2335 let diff: f64 = (0..m)
2337 .map(|j| (res_bs.fitted[(0, j)] - res_f.fitted[(0, j)]).abs())
2338 .sum();
2339 assert!(
2341 diff > 1e-10,
2342 "B-spline and Fourier fits should differ for the same data"
2343 );
2344 }
2345
2346 #[test]
2349 fn test_smooth_basis_gcv_positive_for_noisy_data() {
2350 let m = 50;
2351 let t = uniform_grid(m);
2352 let mut data = FdMatrix::zeros(1, m);
2353 for j in 0..m {
2354 data[(0, j)] = (2.0 * PI * t[j]).sin() + 0.5 * ((j * 37) % 20) as f64 / 20.0 - 0.25;
2356 }
2357 let fdpar = make_bspline_fdpar(&t, 10, 1e-3);
2358 let res = smooth_basis(&data, &t, &fdpar).unwrap();
2359 assert!(res.gcv > 0.0, "GCV should be positive for noisy data");
2360 }
2361
2362 #[test]
2365 fn test_smooth_basis_different_lfd_orders() {
2366 let m = 50;
2367 let (data, t) = make_test_data(2, m);
2368
2369 let penalty1 = bspline_penalty_matrix(&t, 10, 4, 1);
2371 let fdpar1 = FdPar {
2372 basis_type: BasisType::Bspline { order: 4 },
2373 nbasis: 10,
2374 lambda: 1e-2,
2375 lfd_order: 1,
2376 penalty_matrix: penalty1,
2377 };
2378 let res1 = smooth_basis(&data, &t, &fdpar1);
2379 assert!(res1.is_ok());
2380
2381 let penalty2 = bspline_penalty_matrix(&t, 10, 4, 2);
2383 let fdpar2 = FdPar {
2384 basis_type: BasisType::Bspline { order: 4 },
2385 nbasis: 10,
2386 lambda: 1e-2,
2387 lfd_order: 2,
2388 penalty_matrix: penalty2,
2389 };
2390 let res2 = smooth_basis(&data, &t, &fdpar2);
2391 assert!(res2.is_ok());
2392
2393 let r1 = res1.unwrap();
2395 let r2 = res2.unwrap();
2396 let diff: f64 = (0..m)
2397 .map(|j| (r1.fitted[(0, j)] - r2.fitted[(0, j)]).abs())
2398 .sum();
2399 assert!(
2400 diff > 1e-10,
2401 "Different lfd_orders should produce different fits"
2402 );
2403 }
2404
2405 #[test]
2408 fn test_basis_nbasis_cv_result_fields() {
2409 let (data, t) = make_test_data(6, 50);
2410 let nbasis_range: Vec<usize> = vec![5, 7, 9, 11, 13];
2411 let res = basis_nbasis_cv(
2412 &data,
2413 &t,
2414 &nbasis_range,
2415 &BasisType::Bspline { order: 4 },
2416 BasisCriterion::Aic,
2417 5,
2418 1e-4,
2419 )
2420 .unwrap();
2421
2422 assert!(nbasis_range.contains(&res.optimal_nbasis));
2423 assert_eq!(res.scores.len(), nbasis_range.len());
2424 assert_eq!(res.nbasis_range, nbasis_range);
2425 assert_eq!(res.criterion, BasisCriterion::Aic);
2426 let min_score = res.scores.iter().copied().fold(f64::INFINITY, f64::min);
2428 let best_idx = res
2429 .scores
2430 .iter()
2431 .position(|&s| (s - min_score).abs() < 1e-15)
2432 .unwrap();
2433 assert_eq!(res.optimal_nbasis, nbasis_range[best_idx]);
2434 }
2435
2436 #[test]
2437 fn test_basis_nbasis_cv_result_clone() {
2438 let (data, t) = make_test_data(5, 50);
2439 let nbasis_range: Vec<usize> = vec![5, 10];
2440 let res = basis_nbasis_cv(
2441 &data,
2442 &t,
2443 &nbasis_range,
2444 &BasisType::Bspline { order: 4 },
2445 BasisCriterion::Gcv,
2446 5,
2447 1e-4,
2448 )
2449 .unwrap();
2450 let cloned = res.clone();
2451 assert_eq!(res, cloned);
2452 }
2453
2454 #[test]
2457 fn test_smooth_basis_nonuniform_argvals() {
2458 let m = 50;
2459 let t: Vec<f64> = (0..m)
2461 .map(|i| {
2462 let x = i as f64 / (m - 1) as f64;
2463 0.5 * (1.0 - (PI * x).cos())
2464 })
2465 .collect();
2466 let mut data = FdMatrix::zeros(2, m);
2467 for i in 0..2 {
2468 for j in 0..m {
2469 data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.1 * i as f64;
2470 }
2471 }
2472 let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2473 let res = smooth_basis(&data, &t, &fdpar);
2474 assert!(res.is_ok(), "Should handle non-uniform argvals");
2475 let r = res.unwrap();
2476 assert_eq!(r.fitted.shape(), (2, m));
2477 }
2478
2479 #[test]
2482 fn test_smooth_basis_very_small_lambda() {
2483 let m = 50;
2484 let (data, t) = make_test_data(2, m);
2485 let fdpar = make_bspline_fdpar(&t, 10, 1e-15);
2486 let res = smooth_basis(&data, &t, &fdpar);
2487 assert!(res.is_ok(), "Should handle very small lambda");
2488 }
2489
2490 #[test]
2491 fn test_smooth_basis_very_large_lambda() {
2492 let m = 50;
2493 let (data, t) = make_test_data(2, m);
2494 let fdpar = make_bspline_fdpar(&t, 10, 1e10);
2495 let res = smooth_basis(&data, &t, &fdpar);
2496 assert!(res.is_ok(), "Should handle very large lambda");
2497 }
2498
2499 #[test]
2502 fn test_smooth_basis_multi_curve_vs_single_curve() {
2503 let m = 50;
2505 let n = 3;
2506 let (data, t) = make_test_data(n, m);
2507 let fdpar = make_bspline_fdpar(&t, 10, 1e-3);
2508
2509 let res_all = smooth_basis(&data, &t, &fdpar).unwrap();
2511
2512 for i in 0..n {
2514 let mut single = FdMatrix::zeros(1, m);
2515 for j in 0..m {
2516 single[(0, j)] = data[(i, j)];
2517 }
2518 let res_single = smooth_basis(&single, &t, &fdpar).unwrap();
2519 for j in 0..m {
2520 assert!(
2521 (res_all.fitted[(i, j)] - res_single.fitted[(0, j)]).abs() < 1e-10,
2522 "Multi-curve fit should match single-curve fit: curve {} point {}",
2523 i,
2524 j
2525 );
2526 }
2527 }
2528 }
2529
2530 #[test]
2533 fn test_basis_nbasis_cv_all_criteria_finite_scores() {
2534 let (data, t) = make_test_data(10, 60);
2535 let nbasis_range: Vec<usize> = vec![5, 7, 9, 11];
2536
2537 for criterion in [
2538 BasisCriterion::Gcv,
2539 BasisCriterion::Aic,
2540 BasisCriterion::Bic,
2541 BasisCriterion::Cv,
2542 ] {
2543 let res = basis_nbasis_cv(
2544 &data,
2545 &t,
2546 &nbasis_range,
2547 &BasisType::Bspline { order: 4 },
2548 criterion,
2549 5,
2550 1e-4,
2551 )
2552 .unwrap();
2553 let finite_count = res.scores.iter().filter(|s| s.is_finite()).count();
2555 assert!(
2556 finite_count > 0,
2557 "At least one score should be finite for {:?}",
2558 criterion
2559 );
2560 }
2561 }
2562
2563 #[test]
2566 fn test_smooth_basis_gcv_config_default() {
2567 let config = SmoothBasisGcvConfig::default();
2568 assert_eq!(config.basis_type, BasisType::Bspline { order: 4 });
2569 assert_eq!(config.nbasis, 15);
2570 assert_eq!(config.lfd_order, 2);
2571 assert_eq!(config.log_lambda_range, (-10.0, 2.0));
2572 assert_eq!(config.n_grid, 50);
2573 }
2574
2575 #[test]
2576 fn test_smooth_basis_gcv_config_clone_eq() {
2577 let config = SmoothBasisGcvConfig {
2578 nbasis: 20,
2579 ..SmoothBasisGcvConfig::default()
2580 };
2581 let cloned = config.clone();
2582 assert_eq!(config, cloned);
2583 }
2584
2585 #[test]
2586 fn test_smooth_basis_gcv_config_debug() {
2587 let config = SmoothBasisGcvConfig::default();
2588 let debug_str = format!("{:?}", config);
2589 assert!(debug_str.contains("SmoothBasisGcvConfig"));
2590 assert!(debug_str.contains("nbasis"));
2591 }
2592
2593 #[test]
2594 fn test_smooth_basis_gcv_config_partial_override() {
2595 let config = SmoothBasisGcvConfig {
2596 basis_type: BasisType::Fourier { period: 2.0 },
2597 n_grid: 100,
2598 ..SmoothBasisGcvConfig::default()
2599 };
2600 assert_eq!(config.basis_type, BasisType::Fourier { period: 2.0 });
2601 assert_eq!(config.n_grid, 100);
2602 assert_eq!(config.nbasis, 15);
2604 assert_eq!(config.lfd_order, 2);
2605 }
2606
2607 #[test]
2608 fn test_smooth_basis_gcv_with_config_default() {
2609 let (data, t) = make_test_data(5, 101);
2610 let config = SmoothBasisGcvConfig::default();
2611 let result = smooth_basis_gcv_with_config(&data, &t, &config);
2612 assert!(result.is_ok(), "GCV with default config should succeed");
2613 let res = result.unwrap();
2614 assert_eq!(res.fitted.shape(), (5, 101));
2615 assert!(res.edf > 0.0);
2616 assert!(res.gcv.is_finite());
2617 }
2618
2619 #[test]
2620 fn test_smooth_basis_gcv_with_config_custom() {
2621 let (data, t) = make_test_data(3, 50);
2622 let config = SmoothBasisGcvConfig {
2623 nbasis: 10,
2624 log_lambda_range: (-6.0, 0.0),
2625 n_grid: 15,
2626 ..SmoothBasisGcvConfig::default()
2627 };
2628 let result = smooth_basis_gcv_with_config(&data, &t, &config);
2629 assert!(result.is_ok());
2630 }
2631
2632 #[test]
2633 fn test_smooth_basis_gcv_with_config_matches_direct() {
2634 let (data, t) = make_test_data(3, 50);
2635 let config = SmoothBasisGcvConfig {
2636 nbasis: 10,
2637 log_lambda_range: (-6.0, 0.0),
2638 n_grid: 20,
2639 ..SmoothBasisGcvConfig::default()
2640 };
2641 let with_config = smooth_basis_gcv_with_config(&data, &t, &config).unwrap();
2642 let direct = smooth_basis_gcv(
2643 &data,
2644 &t,
2645 &config.basis_type,
2646 config.nbasis,
2647 config.lfd_order,
2648 config.log_lambda_range,
2649 config.n_grid,
2650 )
2651 .unwrap();
2652 assert_eq!(with_config.gcv, direct.gcv);
2653 assert_eq!(with_config.edf, direct.edf);
2654 assert_eq!(with_config.nbasis, direct.nbasis);
2655 }
2656
2657 #[test]
2658 fn test_smooth_basis_gcv_with_config_fourier() {
2659 let m = 100;
2660 let t = uniform_grid(m);
2661 let mut data = FdMatrix::zeros(2, m);
2662 for i in 0..2 {
2663 for j in 0..m {
2664 data[(i, j)] = (2.0 * PI * t[j]).sin() + (4.0 * PI * t[j]).cos();
2665 }
2666 }
2667 let config = SmoothBasisGcvConfig {
2668 basis_type: BasisType::Fourier { period: 1.0 },
2669 nbasis: 7,
2670 n_grid: 20,
2671 ..SmoothBasisGcvConfig::default()
2672 };
2673 let result = smooth_basis_gcv_with_config(&data, &t, &config);
2674 assert!(result.is_ok());
2675 }
2676
2677 #[test]
2680 fn test_basis_nbasis_cv_config_default() {
2681 let config = BasisNbasisCvConfig::default();
2682 assert_eq!(config.basis_type, BasisType::Bspline { order: 4 });
2683 assert_eq!(config.nbasis_range, (5, 30));
2684 assert!((config.lambda - 1e-4).abs() < 1e-15);
2685 assert_eq!(config.lfd_order, 2);
2686 assert_eq!(config.n_folds, 5);
2687 assert_eq!(config.criterion, BasisCriterion::Gcv);
2688 }
2689
2690 #[test]
2691 fn test_basis_nbasis_cv_config_clone_eq() {
2692 let config = BasisNbasisCvConfig {
2693 nbasis_range: (4, 15),
2694 ..BasisNbasisCvConfig::default()
2695 };
2696 let cloned = config.clone();
2697 assert_eq!(config, cloned);
2698 }
2699
2700 #[test]
2701 fn test_basis_nbasis_cv_config_debug() {
2702 let config = BasisNbasisCvConfig::default();
2703 let debug_str = format!("{:?}", config);
2704 assert!(debug_str.contains("BasisNbasisCvConfig"));
2705 assert!(debug_str.contains("nbasis_range"));
2706 }
2707
2708 #[test]
2709 fn test_basis_nbasis_cv_config_partial_override() {
2710 let config = BasisNbasisCvConfig {
2711 criterion: BasisCriterion::Aic,
2712 lambda: 1e-2,
2713 ..BasisNbasisCvConfig::default()
2714 };
2715 assert_eq!(config.criterion, BasisCriterion::Aic);
2716 assert!((config.lambda - 1e-2).abs() < 1e-15);
2717 assert_eq!(config.nbasis_range, (5, 30));
2719 assert_eq!(config.n_folds, 5);
2720 }
2721
2722 #[test]
2723 fn test_basis_nbasis_cv_with_config_default() {
2724 let (data, t) = make_test_data(5, 51);
2725 let config = BasisNbasisCvConfig {
2726 nbasis_range: (5, 12),
2727 ..BasisNbasisCvConfig::default()
2728 };
2729 let result = basis_nbasis_cv_with_config(&data, &t, &config);
2730 assert!(
2731 result.is_ok(),
2732 "nbasis CV with default config should succeed"
2733 );
2734 let res = result.unwrap();
2735 assert!(res.optimal_nbasis >= 5 && res.optimal_nbasis <= 12);
2736 assert_eq!(res.scores.len(), 8); assert_eq!(res.criterion, BasisCriterion::Gcv);
2738 }
2739
2740 #[test]
2741 fn test_basis_nbasis_cv_with_config_aic() {
2742 let (data, t) = make_test_data(5, 51);
2743 let config = BasisNbasisCvConfig {
2744 nbasis_range: (5, 10),
2745 criterion: BasisCriterion::Aic,
2746 ..BasisNbasisCvConfig::default()
2747 };
2748 let result = basis_nbasis_cv_with_config(&data, &t, &config);
2749 assert!(result.is_ok());
2750 assert_eq!(result.unwrap().criterion, BasisCriterion::Aic);
2751 }
2752
2753 #[test]
2754 fn test_basis_nbasis_cv_with_config_cv_folds() {
2755 let (data, t) = make_test_data(10, 51);
2756 let config = BasisNbasisCvConfig {
2757 nbasis_range: (5, 9),
2758 criterion: BasisCriterion::Cv,
2759 n_folds: 3,
2760 ..BasisNbasisCvConfig::default()
2761 };
2762 let result = basis_nbasis_cv_with_config(&data, &t, &config);
2763 assert!(result.is_ok());
2764 assert_eq!(result.unwrap().criterion, BasisCriterion::Cv);
2765 }
2766
2767 #[test]
2768 fn test_basis_nbasis_cv_with_config_matches_direct() {
2769 let (data, t) = make_test_data(5, 51);
2770 let config = BasisNbasisCvConfig {
2771 nbasis_range: (5, 10),
2772 criterion: BasisCriterion::Bic,
2773 lambda: 1e-3,
2774 ..BasisNbasisCvConfig::default()
2775 };
2776 let with_config = basis_nbasis_cv_with_config(&data, &t, &config).unwrap();
2777 let nbasis_range: Vec<usize> = (5..=10).collect();
2778 let direct = basis_nbasis_cv(
2779 &data,
2780 &t,
2781 &nbasis_range,
2782 &config.basis_type,
2783 config.criterion,
2784 config.n_folds,
2785 config.lambda,
2786 )
2787 .unwrap();
2788 assert_eq!(with_config.optimal_nbasis, direct.optimal_nbasis);
2789 assert_eq!(with_config.scores, direct.scores);
2790 assert_eq!(with_config.nbasis_range, direct.nbasis_range);
2791 }
2792
2793 #[test]
2794 fn test_basis_nbasis_cv_with_config_nbasis_range_expansion() {
2795 let (data, t) = make_test_data(5, 51);
2796 let config = BasisNbasisCvConfig {
2797 nbasis_range: (7, 7), ..BasisNbasisCvConfig::default()
2799 };
2800 let result = basis_nbasis_cv_with_config(&data, &t, &config);
2801 assert!(result.is_ok());
2802 let res = result.unwrap();
2803 assert_eq!(res.optimal_nbasis, 7);
2804 assert_eq!(res.scores.len(), 1);
2805 }
2806}