Skip to main content

fdars_core/
smooth_basis.rs

1//! Basis-penalized smoothing with continuous derivative penalties.
2//!
3//! This module implements `smooth.basis` from R's fda package. Unlike the
4//! discrete difference penalty used in P-splines (`basis.rs`), this uses
5//! continuous derivative penalties: `min ||y - Φc||² + λ·∫(Lf)² dt`.
6//!
7//! Key capabilities:
8//! - [`smooth_basis`] — Penalized least squares with continuous roughness penalty
9//! - [`smooth_basis_gcv`] — GCV-optimal smoothing parameter selection
10//! - [`bspline_penalty_matrix`] / [`fourier_penalty_matrix`] — Roughness penalty matrices
11
12use 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// ─── Types ──────────────────────────────────────────────────────────────────
19
20/// Basis type for penalized smoothing.
21#[derive(Debug, Clone, PartialEq)]
22pub enum BasisType {
23    /// B-spline basis with given order (typically 4 for cubic).
24    Bspline { order: usize },
25    /// Fourier basis with given period.
26    Fourier { period: f64 },
27}
28
29/// Functional data parameter object (basis + penalty specification).
30#[derive(Debug, Clone, PartialEq)]
31pub struct FdPar {
32    /// Type of basis system.
33    pub basis_type: BasisType,
34    /// Number of basis functions.
35    pub nbasis: usize,
36    /// Smoothing parameter.
37    pub lambda: f64,
38    /// Derivative order for the penalty (default: 2).
39    pub lfd_order: usize,
40    /// Precomputed K×K penalty matrix (column-major).
41    pub penalty_matrix: Vec<f64>,
42}
43
44/// Result of basis-penalized smoothing.
45#[derive(Debug, Clone, PartialEq)]
46#[non_exhaustive]
47pub struct SmoothBasisResult {
48    /// Basis coefficients (n × K).
49    pub coefficients: FdMatrix,
50    /// Fitted values (n × m).
51    pub fitted: FdMatrix,
52    /// Effective degrees of freedom.
53    pub edf: f64,
54    /// Generalized cross-validation score.
55    pub gcv: f64,
56    /// AIC.
57    pub aic: f64,
58    /// BIC.
59    pub bic: f64,
60    /// Roughness penalty matrix (K × K, column-major).
61    pub penalty_matrix: Vec<f64>,
62    /// Number of basis functions used.
63    pub nbasis: usize,
64}
65
66// ─── Penalty Matrices ───────────────────────────────────────────────────────
67
68/// Compute the roughness penalty matrix for B-splines via numerical quadrature.
69///
70/// R\[j,k\] = ∫ D^m B_j(t) · D^m B_k(t) dt
71///
72/// Uses Simpson's rule on a fine sub-grid for each knot interval.
73///
74/// # Arguments
75/// * `argvals` — Evaluation points (length m)
76/// * `nbasis` — Number of basis functions
77/// * `order` — B-spline order (typically 4 for cubic)
78/// * `lfd_order` — Derivative order for penalty (typically 2)
79///
80/// # Returns
81/// K × K penalty matrix in column-major layout (K = nbasis)
82pub 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    // Create a fine quadrature grid (10 sub-points per original interval)
95    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    // Evaluate B-spline basis on fine grid
104    let basis_fine = bspline_basis(&quad_t, nknots, order);
105    let actual_nbasis = basis_fine.len() / n_quad;
106
107    // Compute derivatives of B-spline basis numerically
108    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    // Integration weights on fine grid
112    let weights = simpsons_weights(&quad_t);
113
114    // Compute penalty matrix: R[j,k] = ∫ D^m B_j · D^m B_k dt
115    integrate_symmetric_penalty(&deriv_basis, &weights, actual_nbasis, n_quad)
116}
117
118/// Compute the roughness penalty matrix for a Fourier basis.
119///
120/// For Fourier basis, the penalty is diagonal with eigenvalues `(2πk/T)^(2m)`.
121///
122/// # Arguments
123/// * `nbasis` — Number of basis functions
124/// * `period` — Period of the Fourier basis
125/// * `lfd_order` — Derivative order for penalty
126///
127/// # Returns
128/// K × K penalty matrix in column-major layout
129pub 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    // First basis function is constant → lfd_order-th derivative is 0
134    // penalty[0] = 0 (already zero)
135
136    // For sin/cos pairs: eigenvalue is (2πk/T)^(2m)
137    // Matches R's fda package convention (sqrt(2)-normalized basis)
138    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        // sin component
145        if idx < k {
146            penalty[idx + idx * k] = eigenval;
147            idx += 1;
148        }
149        // cos component
150        if idx < k {
151            penalty[idx + idx * k] = eigenval;
152            idx += 1;
153        }
154        freq += 1;
155    }
156
157    penalty
158}
159
160// ─── Smoothing Functions ────────────────────────────────────────────────────
161
162/// Perform basis-penalized smoothing.
163///
164/// Solves `(Φ'Φ + λR)c = Φ'y` per curve via Cholesky decomposition.
165/// This implements `smooth.basis` from R's fda package.
166///
167/// # Arguments
168/// * `data` — Functional data matrix (n × m)
169/// * `argvals` — Evaluation points (length m)
170/// * `fdpar` — Functional parameter object specifying basis and penalty
171///
172/// # Returns
173/// [`SmoothBasisResult`] with coefficients, fitted values, and diagnostics.
174pub 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    // Evaluate basis on argvals
195    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    // (Φ'Φ + λR + εI) — small ridge ensures positive definiteness
202    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    // Invert the penalized system
208    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    // Hat matrix: H = Φ (Φ'Φ + λR)^{-1} Φ'  →  EDF = tr(H)
215    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    // Project all curves
219    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    // Total effective degrees of freedom = n curves * per-curve edf
226    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
242/// Perform basis-penalized smoothing with GCV-optimal lambda.
243///
244/// Searches over a log-lambda grid and selects the lambda minimizing GCV.
245///
246/// # Arguments
247/// * `data` — Functional data matrix (n × m)
248/// * `argvals` — Evaluation points (length m)
249/// * `basis_type` — Type of basis system
250/// * `nbasis` — Number of basis functions
251/// * `lfd_order` — Derivative order for penalty
252/// * `log_lambda_range` — Range of log10(lambda) to search, e.g. (-8.0, 4.0)
253/// * `n_grid` — Number of grid points for the search
254pub 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    // Compute penalty matrix once
269    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// ─── Config Structs ─────────────────────────────────────────────────────────
302
303/// Configuration for GCV-based smoothing parameter selection.
304///
305/// Collects all tuning parameters for [`smooth_basis_gcv_with_config`], with
306/// sensible defaults obtained via [`SmoothBasisGcvConfig::default()`].
307///
308/// # Example
309/// ```no_run
310/// use fdars_core::smooth_basis::{SmoothBasisGcvConfig, BasisType};
311///
312/// let config = SmoothBasisGcvConfig {
313///     nbasis: 20,
314///     n_grid: 100,
315///     ..SmoothBasisGcvConfig::default()
316/// };
317/// ```
318#[derive(Debug, Clone, PartialEq)]
319pub struct SmoothBasisGcvConfig {
320    /// Basis type (BSpline or Fourier).
321    pub basis_type: BasisType,
322    /// Number of basis functions (default: 15).
323    pub nbasis: usize,
324    /// Order of the roughness penalty differential operator (default: 2).
325    pub lfd_order: usize,
326    /// Range of log10(lambda) values to search (default: (-10.0, 2.0)).
327    pub log_lambda_range: (f64, f64),
328    /// Number of grid points in the lambda search (default: 50).
329    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/// Perform basis-penalized smoothing with GCV-optimal lambda using a config struct.
345///
346/// This is the config-based alternative to [`smooth_basis_gcv`]. It takes data
347/// parameters directly and reads all tuning parameters from the config.
348///
349/// # Arguments
350/// * `data` — Functional data matrix (n × m)
351/// * `argvals` — Evaluation points (length m)
352/// * `config` — Tuning parameters
353///
354/// # Errors
355///
356/// Returns [`crate::FdarError::ComputationFailed`] if no valid smoothing result
357/// is found for any lambda in the search grid.
358#[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/// Configuration for cross-validation-based basis selection.
380///
381/// Collects all tuning parameters for [`basis_nbasis_cv_with_config`], with
382/// sensible defaults obtained via [`BasisNbasisCvConfig::default()`].
383///
384/// # Example
385/// ```no_run
386/// use fdars_core::smooth_basis::{BasisNbasisCvConfig, BasisType, BasisCriterion};
387///
388/// let config = BasisNbasisCvConfig {
389///     nbasis_range: (5, 25),
390///     criterion: BasisCriterion::Aic,
391///     ..BasisNbasisCvConfig::default()
392/// };
393/// ```
394#[derive(Debug, Clone, PartialEq)]
395pub struct BasisNbasisCvConfig {
396    /// Basis type (default: BSpline with order 4).
397    pub basis_type: BasisType,
398    /// Range of nbasis values to try, inclusive (default: (5, 30)).
399    pub nbasis_range: (usize, usize),
400    /// Roughness penalty lambda (default: 1e-4).
401    pub lambda: f64,
402    /// Penalty order (default: 2).
403    pub lfd_order: usize,
404    /// Number of CV folds (default: 5). Only used when `criterion` is `Cv`.
405    pub n_folds: usize,
406    /// Selection criterion (default: `Gcv`).
407    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/// Select the optimal number of basis functions using a config struct.
424///
425/// This is the config-based alternative to [`basis_nbasis_cv`]. It takes data
426/// parameters directly and reads all tuning parameters from the config.
427///
428/// The `nbasis_range` tuple `(lo, hi)` is expanded to `lo..=hi` to form the
429/// candidate set.
430///
431/// # Arguments
432/// * `data` — Functional data matrix (n × m)
433/// * `argvals` — Evaluation points (length m)
434/// * `config` — Tuning parameters
435///
436/// # Errors
437///
438/// Returns [`crate::FdarError::ComputationFailed`] if no valid result is found
439/// for any nbasis in the search range.
440#[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
462// ─── Internal Helpers ───────────────────────────────────────────────────────
463
464/// Differentiate column-major basis matrix `lfd_order` times using gradient_uniform.
465fn 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
487/// Integrate symmetric penalty: R[j,k] = ∫ D^m B_j · D^m B_k dt.
488fn 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
508/// Evaluate basis functions on argvals, returning (flat column-major, actual_nbasis).
509fn 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
525/// Invert the penalized system matrix via Cholesky or SVD pseudoinverse.
526fn 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    // SVD fallback
531    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
551/// Project all curves onto basis, returning (coefficients, fitted, total_rss).
552fn 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
583/// Compute GCV score.
584fn 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// ─── Nbasis Selection via CV ────────────────────────────────────────────────
594
595/// Criterion for nbasis selection.
596#[derive(Debug, Clone, Copy, PartialEq)]
597pub enum BasisCriterion {
598    /// Generalized cross-validation.
599    Gcv,
600    /// Leave-one-out cross-validation (k-fold).
601    Cv,
602    /// Akaike Information Criterion.
603    Aic,
604    /// Bayesian Information Criterion.
605    Bic,
606}
607
608/// Result of nbasis selection.
609#[derive(Debug, Clone, PartialEq)]
610#[non_exhaustive]
611pub struct BasisNbasisCvResult {
612    /// Optimal number of basis functions.
613    pub optimal_nbasis: usize,
614    /// Score for each nbasis tested.
615    pub scores: Vec<f64>,
616    /// Range of nbasis values tested.
617    pub nbasis_range: Vec<usize>,
618    /// Criterion used.
619    pub criterion: BasisCriterion,
620}
621
622/// Evaluate information criterion (GCV/AIC/BIC) for a range of nbasis values.
623fn 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
664/// Evaluate nbasis via k-fold cross-validation of reconstruction error.
665fn 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    // Cross-validate over TIME POINTS, not curves. Leaving out curves is
675    // ill-posed for per-curve basis fitting — each curve has its own
676    // coefficients, so a held-out curve can only be scored against its own data,
677    // which is an in-sample residual that decreases monotonically in `nbasis`
678    // and always selects the largest candidate (GH #33). Holding out points and
679    // predicting them from a fit on the remaining points gives a genuine
680    // predictive score that penalizes overfitting and shows an interior minimum.
681    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            // Basis rows at the train/test points, and the projection operator
707            // built from the training points only (independent of any curve).
708            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(); // (k x |train|)
718
719            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; // predictions at held-out points
726                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
743/// Select the optimal number of basis functions using multiple criteria
744/// (R's `fdata2basis_cv`).
745pub 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; // may differ from actual due to knot construction
799        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        // Diagonal elements should be non-negative
818        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        // Should be diagonal
832        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        // Constant term should have zero penalty
846        assert!(penalty[0].abs() < 1e-10);
847        // Higher frequency terms should have larger penalties
848        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        // Generate noisy sine curves
859        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        // Fourier basis should fit periodic data well
918        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    // ============== basis_nbasis_cv tests ==============
993
994    #[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    /// Regression for GH #33: the CV path scored held-out curves against their
1088    /// own data (no true hold-out), so scores fell monotonically in `n_basis`
1089    /// and it always selected the maximum candidate. With point-wise CV,
1090    /// overfitting the noise is penalized and the maximum is not chosen.
1091    #[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                // Smooth signal + deterministic pseudo-noise.
1100                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    // ============== Comprehensive additional tests ==============
1131
1132    // Helper: generate standard test data (sine + high-freq component)
1133    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    // Helper: create an FdPar for B-spline smoothing
1147    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    // Helper: create an FdPar for Fourier smoothing
1159    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    // ─── BasisType enum tests ───────────────────────────────────────────────
1171
1172    #[test]
1173    fn test_basis_type_bspline_variant() {
1174        let bt = BasisType::Bspline { order: 4 };
1175        assert_eq!(bt, BasisType::Bspline { order: 4 });
1176        // Different orders are not equal
1177        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    // ─── FdPar struct tests ─────────────────────────────────────────────────
1205
1206    #[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    // ─── BasisCriterion enum tests ──────────────────────────────────────────
1233
1234    #[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; // Copy
1248        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    // ─── SmoothBasisResult tests ────────────────────────────────────────────
1258
1259    #[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        // coefficients: n curves x k basis functions
1266        assert_eq!(res.coefficients.nrows(), 3);
1267        assert!(res.coefficients.ncols() > 0);
1268        assert_eq!(res.nbasis, res.coefficients.ncols());
1269        // fitted: n x m
1270        assert_eq!(res.fitted.shape(), (3, 50));
1271        // edf should be between 1 and nbasis
1272        assert!(res.edf > 0.0 && res.edf <= res.nbasis as f64);
1273        // gcv, aic, bic should be finite
1274        assert!(res.gcv.is_finite());
1275        assert!(res.aic.is_finite());
1276        assert!(res.bic.is_finite());
1277        // penalty_matrix should be k x k
1278        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    // ─── smooth_basis: B-spline detailed tests ──────────────────────────────
1292
1293    #[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        // actual nbasis may differ from requested due to knot construction
1301        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        // With lambda=0, the smoother should nearly interpolate the data
1318        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        // Residuals should be very small (near interpolation)
1325        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        // With very large lambda, the fit should be much smoother (lower variance)
1342        // than with small lambda
1343        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        // Compare roughness of fits with small vs large lambda
1372        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        // Measure roughness as sum of squared second differences
1383        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        // nbasis = 2 is the minimum allowed
1427        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        // Should succeed (or at least not panic); the fit may be poor
1432        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        // Order 3 (quadratic B-splines)
1440        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        // Order 5 (quartic B-splines)
1452        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    // ─── smooth_basis: Fourier detailed tests ───────────────────────────────
1465
1466    #[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        // Fourier basis should perfectly fit a pure sine with enough basis fns
1488        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        // Period = 1.0 (matches the data)
1519        let fdpar1 = make_fourier_fdpar(7, 1.0, 1e-6);
1520        let res1 = smooth_basis(&data, &t, &fdpar1).unwrap();
1521
1522        // Period = 2.0 (mismatch, but should still produce a result)
1523        let fdpar2 = make_fourier_fdpar(7, 2.0, 1e-6);
1524        let res2 = smooth_basis(&data, &t, &fdpar2).unwrap();
1525
1526        // Both should succeed and have valid shapes
1527        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        // EDF should be close to nbasis with zero penalty
1543        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        // EDF should be very small with huge penalty
1557        assert!(
1558            res.edf < 5.0,
1559            "Large lambda should reduce EDF; edf={}",
1560            res.edf
1561        );
1562    }
1563
1564    // ─── smooth_basis: Lambda comparison tests ──────────────────────────────
1565
1566    #[test]
1567    fn test_smooth_basis_lambda_gradient_edf() {
1568        // EDF should monotonically decrease with increasing lambda
1569        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        // RSS should monotonically increase with increasing lambda
1590        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    // ─── smooth_basis: Error cases ──────────────────────────────────────────
1616
1617    #[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); // m=40 but argvals has 50
1644        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        // nbasis = 1, which is below minimum of 2
1654        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 { .. } => {} // expected
1673            other => panic!("Expected InvalidDimension, got {:?}", other),
1674        }
1675    }
1676
1677    // ─── Penalty matrix detailed tests ──────────────────────────────────────
1678
1679    #[test]
1680    fn test_bspline_penalty_matrix_different_orders() {
1681        let t = uniform_grid(101);
1682        // Order 1 penalty (penalize derivatives)
1683        let p1 = bspline_penalty_matrix(&t, 10, 4, 1);
1684        // Order 2 penalty (penalize curvature)
1685        let p2 = bspline_penalty_matrix(&t, 10, 4, 2);
1686        // Both should be square and same size
1687        assert_eq!(p1.len(), p2.len());
1688        // But they should differ
1689        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        // Too few argvals
1699        let t = vec![0.0];
1700        let p = bspline_penalty_matrix(&t, 10, 4, 2);
1701        // Should return zero matrix
1702        assert!(p.iter().all(|&v| v == 0.0));
1703
1704        // nbasis < 2
1705        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        // lfd_order >= order
1710        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        // Constant term is zero
1738        assert!(penalty[0].abs() < 1e-15);
1739        // Pairs: (1,2) -> freq 1, (3,4) -> freq 2, etc.
1740        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            // cos and sin of same frequency should have same penalty
1753            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        // Longer period -> smaller omega -> smaller penalty eigenvalues
1769        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        // lfd_order = 1: penalize first derivative
1781        let p = fourier_penalty_matrix(5, 1.0, 1);
1782        // Eigenvalues: (2*pi*freq)^2 for lfd_order=1
1783        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); // constant term has zero penalty
1804    }
1805
1806    // ─── smooth_basis_gcv detailed tests ────────────────────────────────────
1807
1808    #[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        // Should return None since smooth_basis will error for empty data
1854        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        // Very narrow search range
1886        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        // Very wide search range
1895        let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-12.0, 8.0), 30);
1896        assert!(result.is_some());
1897    }
1898
1899    // ─── basis_nbasis_cv detailed tests ─────────────────────────────────────
1900
1901    #[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        // Range includes nbasis = 1 which is invalid
2003        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        // Score for nbasis=1 should be infinity, so optimal should be 5 or 10
2016        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); // mismatch
2061        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        // BIC penalizes complexity more heavily than AIC, so it should generally
2095        // select the same or fewer basis functions
2096        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        // BIC should select at most as many basis functions as AIC
2120        // (not guaranteed in all cases, but typical behavior)
2121        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    // ─── Fitted values quality tests ────────────────────────────────────────
2130
2131    #[test]
2132    fn test_smooth_basis_fitted_close_to_data() {
2133        // With moderate penalty and enough basis functions, fitted should be close to data
2134        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        // Constant data should be fit exactly
2163        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        // Linear data should be fit well with cubic B-splines
2190        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    // ─── EDF and diagnostic tests ───────────────────────────────────────────
2211
2212    #[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        // EDF should be between 1 and m (evaluation points)
2219        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    // ─── Penalty matrix size consistency tests ──────────────────────────────
2238
2239    #[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    // ─── Regression: multiple identical curves ──────────────────────────────
2258
2259    #[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        // All curves should have the same coefficients
2275        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    // ─── Cross-validation: different numbers of folds ───────────────────────
2289
2290    #[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    // ─── Large nbasis / more basis than reasonable ──────────────────────────
2311
2312    #[test]
2313    fn test_smooth_basis_many_basis_functions() {
2314        let m = 100;
2315        let (data, t) = make_test_data(2, m);
2316        // Many basis functions relative to data points
2317        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    // ─── evaluate_basis internal function (indirectly tested) ───────────────
2326
2327    #[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        // Results should differ between the two basis types
2336        let diff: f64 = (0..m)
2337            .map(|j| (res_bs.fitted[(0, j)] - res_f.fitted[(0, j)]).abs())
2338            .sum();
2339        // They fit the same data, so some difference is expected but not huge
2340        assert!(
2341            diff > 1e-10,
2342            "B-spline and Fourier fits should differ for the same data"
2343        );
2344    }
2345
2346    // ─── compute_gcv edge cases (indirectly tested) ─────────────────────────
2347
2348    #[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            // Noisy data
2355            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    // ─── Penalty order (lfd_order) tests ────────────────────────────────────
2363
2364    #[test]
2365    fn test_smooth_basis_different_lfd_orders() {
2366        let m = 50;
2367        let (data, t) = make_test_data(2, m);
2368
2369        // lfd_order = 1 (penalize first derivative)
2370        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        // lfd_order = 2 (penalize second derivative)
2382        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        // Different penalty orders should produce different fitted values
2394        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    // ─── BasisNbasisCvResult field tests ────────────────────────────────────
2406
2407    #[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        // optimal_nbasis should correspond to minimum score
2427        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    // ─── Non-uniform argvals ────────────────────────────────────────────────
2455
2456    #[test]
2457    fn test_smooth_basis_nonuniform_argvals() {
2458        let m = 50;
2459        // Non-uniform grid: denser at the ends
2460        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    // ─── Numerical stability with extreme lambda ────────────────────────────
2480
2481    #[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    // ─── Multiple curves consistency ────────────────────────────────────────
2500
2501    #[test]
2502    fn test_smooth_basis_multi_curve_vs_single_curve() {
2503        // Smoothing multiple curves at once should give the same result as smoothing each individually
2504        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        // All at once
2510        let res_all = smooth_basis(&data, &t, &fdpar).unwrap();
2511
2512        // One at a time
2513        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    // ─── BasisCriterion comparison: all criteria produce finite scores ──────
2531
2532    #[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            // At least some scores should be finite (valid nbasis values)
2554            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    // ─── SmoothBasisGcvConfig tests ────────────────────────────────────────
2564
2565    #[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        // defaults preserved
2603        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    // ─── BasisNbasisCvConfig tests ─────────────────────────────────────────
2678
2679    #[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        // defaults preserved
2718        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); // 5..=12 = 8 values
2737        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), // single value
2798            ..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}