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/// Perform basis-penalized smoothing with AIC-optimal lambda.
302///
303/// Mirrors [`smooth_basis_gcv`] but selects the lambda minimizing the `aic`
304/// field of [`SmoothBasisResult`] (which `smooth_basis` computes as
305/// `total_points·ln(mse) + 2·total_edf`, using the hat-matrix trace as the
306/// degrees of freedom — the same df-as-trace convention as the kernel-path
307/// [`aic_smoother`](crate::smoothing::aic_smoother)). It searches the identical
308/// log-lambda grid, so it is a drop-in AIC alternative that leaves
309/// [`smooth_basis_gcv`] untouched.
310///
311/// # Arguments
312/// * `data` — Functional data matrix (n × m)
313/// * `argvals` — Evaluation points (length m)
314/// * `basis_type` — Type of basis system
315/// * `nbasis` — Number of basis functions
316/// * `lfd_order` — Derivative order for penalty
317/// * `log_lambda_range` — Range of log10(lambda) to search, e.g. (-8.0, 4.0)
318/// * `n_grid` — Number of grid points for the search
319///
320/// # Example
321/// ```no_run
322/// use fdars_core::smooth_basis::{smooth_basis_aic, BasisType};
323/// use fdars_core::matrix::FdMatrix;
324///
325/// let argvals: Vec<f64> = (0..20).map(|i| i as f64 / 19.0).collect();
326/// let data = FdMatrix::from_column_major(vec![0.0; 20], 1, 20).unwrap();
327/// let result = smooth_basis_aic(
328///     &data,
329///     &argvals,
330///     &BasisType::Bspline { order: 4 },
331///     8,
332///     2,
333///     (-8.0, 4.0),
334///     30,
335/// );
336/// ```
337pub fn smooth_basis_aic(
338    data: &FdMatrix,
339    argvals: &[f64],
340    basis_type: &BasisType,
341    nbasis: usize,
342    lfd_order: usize,
343    log_lambda_range: (f64, f64),
344    n_grid: usize,
345) -> Option<SmoothBasisResult> {
346    let m = argvals.len();
347    if m == 0 || nbasis < 2 || n_grid < 2 {
348        return None;
349    }
350
351    // Compute penalty matrix once (same as smooth_basis_gcv).
352    let penalty = match basis_type {
353        BasisType::Bspline { order } => bspline_penalty_matrix(argvals, nbasis, *order, lfd_order),
354        BasisType::Fourier { period } => fourier_penalty_matrix(nbasis, *period, lfd_order),
355    };
356
357    let (lo, hi) = log_lambda_range;
358    let mut best_aic = f64::INFINITY;
359    let mut best_result: Option<SmoothBasisResult> = None;
360
361    for i in 0..n_grid {
362        let log_lam = lo + (hi - lo) * i as f64 / (n_grid - 1) as f64;
363        let lam = 10.0_f64.powf(log_lam);
364
365        let fdpar = FdPar {
366            basis_type: basis_type.clone(),
367            nbasis,
368            lambda: lam,
369            lfd_order,
370            penalty_matrix: penalty.clone(),
371        };
372
373        if let Ok(result) = smooth_basis(data, argvals, &fdpar) {
374            // Reuse the aic field smooth_basis already populates; do NOT
375            // recompute AIC inline, keeping the formula consistent.
376            if result.aic < best_aic {
377                best_aic = result.aic;
378                best_result = Some(result);
379            }
380        }
381    }
382
383    best_result
384}
385
386// ─── Config Structs ─────────────────────────────────────────────────────────
387
388/// Configuration for GCV-based smoothing parameter selection.
389///
390/// Collects all tuning parameters for [`smooth_basis_gcv_with_config`], with
391/// sensible defaults obtained via [`SmoothBasisGcvConfig::default()`].
392///
393/// # Example
394/// ```no_run
395/// use fdars_core::smooth_basis::{SmoothBasisGcvConfig, BasisType};
396///
397/// let config = SmoothBasisGcvConfig {
398///     nbasis: 20,
399///     n_grid: 100,
400///     ..SmoothBasisGcvConfig::default()
401/// };
402/// ```
403#[derive(Debug, Clone, PartialEq)]
404pub struct SmoothBasisGcvConfig {
405    /// Basis type (BSpline or Fourier).
406    pub basis_type: BasisType,
407    /// Number of basis functions (default: 15).
408    pub nbasis: usize,
409    /// Order of the roughness penalty differential operator (default: 2).
410    pub lfd_order: usize,
411    /// Range of log10(lambda) values to search (default: (-10.0, 2.0)).
412    pub log_lambda_range: (f64, f64),
413    /// Number of grid points in the lambda search (default: 50).
414    pub n_grid: usize,
415}
416
417impl Default for SmoothBasisGcvConfig {
418    fn default() -> Self {
419        Self {
420            basis_type: BasisType::Bspline { order: 4 },
421            nbasis: 15,
422            lfd_order: 2,
423            log_lambda_range: (-10.0, 2.0),
424            n_grid: 50,
425        }
426    }
427}
428
429/// Perform basis-penalized smoothing with GCV-optimal lambda using a config struct.
430///
431/// This is the config-based alternative to [`smooth_basis_gcv`]. It takes data
432/// parameters directly and reads all tuning parameters from the config.
433///
434/// # Arguments
435/// * `data` — Functional data matrix (n × m)
436/// * `argvals` — Evaluation points (length m)
437/// * `config` — Tuning parameters
438///
439/// # Errors
440///
441/// Returns [`crate::FdarError::ComputationFailed`] if no valid smoothing result
442/// is found for any lambda in the search grid.
443#[must_use = "expensive computation whose result should not be discarded"]
444pub fn smooth_basis_gcv_with_config(
445    data: &FdMatrix,
446    argvals: &[f64],
447    config: &SmoothBasisGcvConfig,
448) -> Result<SmoothBasisResult, crate::FdarError> {
449    smooth_basis_gcv(
450        data,
451        argvals,
452        &config.basis_type,
453        config.nbasis,
454        config.lfd_order,
455        config.log_lambda_range,
456        config.n_grid,
457    )
458    .ok_or_else(|| crate::FdarError::ComputationFailed {
459        operation: "smooth_basis_gcv_with_config",
460        detail: "no valid smoothing result found in GCV lambda search".to_string(),
461    })
462}
463
464/// Configuration for cross-validation-based basis selection.
465///
466/// Collects all tuning parameters for [`basis_nbasis_cv_with_config`], with
467/// sensible defaults obtained via [`BasisNbasisCvConfig::default()`].
468///
469/// # Example
470/// ```no_run
471/// use fdars_core::smooth_basis::{BasisNbasisCvConfig, BasisType, BasisCriterion};
472///
473/// let config = BasisNbasisCvConfig {
474///     nbasis_range: (5, 25),
475///     criterion: BasisCriterion::Aic,
476///     ..BasisNbasisCvConfig::default()
477/// };
478/// ```
479#[derive(Debug, Clone, PartialEq)]
480pub struct BasisNbasisCvConfig {
481    /// Basis type (default: BSpline with order 4).
482    pub basis_type: BasisType,
483    /// Range of nbasis values to try, inclusive (default: (5, 30)).
484    pub nbasis_range: (usize, usize),
485    /// Roughness penalty lambda (default: 1e-4).
486    pub lambda: f64,
487    /// Penalty order (default: 2).
488    pub lfd_order: usize,
489    /// Number of CV folds (default: 5). Only used when `criterion` is `Cv`.
490    pub n_folds: usize,
491    /// Selection criterion (default: `Gcv`).
492    pub criterion: BasisCriterion,
493}
494
495impl Default for BasisNbasisCvConfig {
496    fn default() -> Self {
497        Self {
498            basis_type: BasisType::Bspline { order: 4 },
499            nbasis_range: (5, 30),
500            lambda: 1e-4,
501            lfd_order: 2,
502            n_folds: 5,
503            criterion: BasisCriterion::Gcv,
504        }
505    }
506}
507
508/// Select the optimal number of basis functions using a config struct.
509///
510/// This is the config-based alternative to [`basis_nbasis_cv`]. It takes data
511/// parameters directly and reads all tuning parameters from the config.
512///
513/// The `nbasis_range` tuple `(lo, hi)` is expanded to `lo..=hi` to form the
514/// candidate set.
515///
516/// # Arguments
517/// * `data` — Functional data matrix (n × m)
518/// * `argvals` — Evaluation points (length m)
519/// * `config` — Tuning parameters
520///
521/// # Errors
522///
523/// Returns [`crate::FdarError::ComputationFailed`] if no valid result is found
524/// for any nbasis in the search range.
525#[must_use = "expensive computation whose result should not be discarded"]
526pub fn basis_nbasis_cv_with_config(
527    data: &FdMatrix,
528    argvals: &[f64],
529    config: &BasisNbasisCvConfig,
530) -> Result<BasisNbasisCvResult, crate::FdarError> {
531    let nbasis_range: Vec<usize> = (config.nbasis_range.0..=config.nbasis_range.1).collect();
532    basis_nbasis_cv(
533        data,
534        argvals,
535        &nbasis_range,
536        &config.basis_type,
537        config.criterion,
538        config.n_folds,
539        config.lambda,
540    )
541    .ok_or_else(|| crate::FdarError::ComputationFailed {
542        operation: "basis_nbasis_cv_with_config",
543        detail: "no valid result found in nbasis CV search".to_string(),
544    })
545}
546
547// ─── Internal Helpers ───────────────────────────────────────────────────────
548
549/// Differentiate column-major basis matrix `lfd_order` times using gradient_uniform.
550fn differentiate_basis_columns(
551    basis: &[f64],
552    n_quad: usize,
553    nbasis: usize,
554    h: f64,
555    lfd_order: usize,
556) -> Vec<f64> {
557    let mut deriv = basis.to_vec();
558    for _ in 0..lfd_order {
559        let mut new_deriv = vec![0.0; n_quad * nbasis];
560        for j in 0..nbasis {
561            let col: Vec<f64> = (0..n_quad).map(|i| deriv[i + j * n_quad]).collect();
562            let grad = crate::helpers::gradient_uniform(&col, h);
563            for i in 0..n_quad {
564                new_deriv[i + j * n_quad] = grad[i];
565            }
566        }
567        deriv = new_deriv;
568    }
569    deriv
570}
571
572/// Integrate symmetric penalty: R[j,k] = ∫ D^m B_j · D^m B_k dt.
573fn integrate_symmetric_penalty(
574    deriv_basis: &[f64],
575    weights: &[f64],
576    k: usize,
577    n_quad: usize,
578) -> Vec<f64> {
579    let mut penalty = vec![0.0; k * k];
580    for j in 0..k {
581        for l in j..k {
582            let mut val = 0.0;
583            for i in 0..n_quad {
584                val += deriv_basis[i + j * n_quad] * deriv_basis[i + l * n_quad] * weights[i];
585            }
586            penalty[j + l * k] = val;
587            penalty[l + j * k] = val;
588        }
589    }
590    penalty
591}
592
593/// Evaluate basis functions on argvals, returning (flat column-major, actual_nbasis).
594fn evaluate_basis(argvals: &[f64], basis_type: &BasisType, nbasis: usize) -> (Vec<f64>, usize) {
595    let m = argvals.len();
596    match basis_type {
597        BasisType::Bspline { order } => {
598            let nknots = nbasis.saturating_sub(*order).max(2);
599            let basis = bspline_basis(argvals, nknots, *order);
600            let actual = basis.len() / m;
601            (basis, actual)
602        }
603        BasisType::Fourier { period } => {
604            let basis = fourier_basis_with_period(argvals, nbasis, *period);
605            (basis, nbasis)
606        }
607    }
608}
609
610/// Invert the penalized system matrix via Cholesky or SVD pseudoinverse.
611fn invert_penalized_system(system: &DMatrix<f64>, k: usize) -> Option<DMatrix<f64>> {
612    if let Some(chol) = system.clone().cholesky() {
613        return Some(chol.inverse());
614    }
615    // SVD fallback
616    let svd = nalgebra::SVD::new(system.clone(), true, true);
617    let u = svd.u.as_ref()?;
618    let v_t = svd.v_t.as_ref()?;
619    let max_sv: f64 = svd.singular_values.iter().copied().fold(0.0_f64, f64::max);
620    let eps = 1e-10 * max_sv;
621    let mut inv = DMatrix::<f64>::zeros(k, k);
622    for ii in 0..k {
623        for jj in 0..k {
624            let mut sum = 0.0;
625            for s in 0..k.min(svd.singular_values.len()) {
626                if svd.singular_values[s] > eps {
627                    sum += v_t[(s, ii)] / svd.singular_values[s] * u[(jj, s)];
628                }
629            }
630            inv[(ii, jj)] = sum;
631        }
632    }
633    Some(inv)
634}
635
636/// Project all curves onto basis, returning (coefficients, fitted, total_rss).
637fn project_all_curves(
638    data: &FdMatrix,
639    b_mat: &DMatrix<f64>,
640    proj: &DMatrix<f64>,
641    n: usize,
642    m: usize,
643    k: usize,
644) -> (FdMatrix, FdMatrix, f64) {
645    let mut all_coefs = FdMatrix::zeros(n, k);
646    let mut all_fitted = FdMatrix::zeros(n, m);
647    let mut total_rss = 0.0;
648
649    for i in 0..n {
650        let curve: Vec<f64> = (0..m).map(|j| data[(i, j)]).collect();
651        let y_vec = nalgebra::DVector::from_vec(curve.clone());
652        let coefs = proj * &y_vec;
653
654        for j in 0..k {
655            all_coefs[(i, j)] = coefs[j];
656        }
657        let fitted = b_mat * &coefs;
658        for j in 0..m {
659            all_fitted[(i, j)] = fitted[j];
660            let resid = curve[j] - fitted[j];
661            total_rss += resid * resid;
662        }
663    }
664
665    (all_coefs, all_fitted, total_rss)
666}
667
668/// Compute GCV score.
669fn compute_gcv(rss: f64, n_points: f64, edf: f64, m: usize) -> f64 {
670    let gcv_denom = 1.0 - edf / m as f64;
671    if gcv_denom.abs() > 1e-10 {
672        (rss / n_points) / (gcv_denom * gcv_denom)
673    } else {
674        f64::INFINITY
675    }
676}
677
678// ─── Nbasis Selection via CV ────────────────────────────────────────────────
679
680/// Criterion for nbasis selection.
681#[derive(Debug, Clone, Copy, PartialEq)]
682pub enum BasisCriterion {
683    /// Generalized cross-validation.
684    Gcv,
685    /// Leave-one-out cross-validation (k-fold).
686    Cv,
687    /// Akaike Information Criterion.
688    Aic,
689    /// Bayesian Information Criterion.
690    Bic,
691}
692
693/// Result of nbasis selection.
694#[derive(Debug, Clone, PartialEq)]
695#[non_exhaustive]
696pub struct BasisNbasisCvResult {
697    /// Optimal number of basis functions.
698    pub optimal_nbasis: usize,
699    /// Score for each nbasis tested.
700    pub scores: Vec<f64>,
701    /// Range of nbasis values tested.
702    pub nbasis_range: Vec<usize>,
703    /// Criterion used.
704    pub criterion: BasisCriterion,
705}
706
707/// Evaluate information criterion (GCV/AIC/BIC) for a range of nbasis values.
708fn evaluate_nbasis_info_criterion(
709    data: &FdMatrix,
710    argvals: &[f64],
711    nbasis_range: &[usize],
712    basis_type: &BasisType,
713    criterion: BasisCriterion,
714    lambda: f64,
715) -> Vec<f64> {
716    let mut scores = Vec::with_capacity(nbasis_range.len());
717    for &nb in nbasis_range {
718        if nb < 2 {
719            scores.push(f64::INFINITY);
720            continue;
721        }
722        let penalty = match basis_type {
723            BasisType::Bspline { order } => bspline_penalty_matrix(argvals, nb, *order, 2),
724            BasisType::Fourier { period } => fourier_penalty_matrix(nb, *period, 2),
725        };
726        let fdpar = FdPar {
727            basis_type: basis_type.clone(),
728            nbasis: nb,
729            lambda,
730            lfd_order: 2,
731            penalty_matrix: penalty,
732        };
733        match smooth_basis(data, argvals, &fdpar) {
734            Ok(result) => {
735                let score = match criterion {
736                    BasisCriterion::Gcv => result.gcv,
737                    BasisCriterion::Aic => result.aic,
738                    BasisCriterion::Bic => result.bic,
739                    BasisCriterion::Cv => unreachable!(),
740                };
741                scores.push(score);
742            }
743            Err(_) => scores.push(f64::INFINITY),
744        }
745    }
746    scores
747}
748
749/// Evaluate nbasis via k-fold cross-validation of reconstruction error.
750fn evaluate_nbasis_cv(
751    data: &FdMatrix,
752    argvals: &[f64],
753    nbasis_range: &[usize],
754    basis_type: &BasisType,
755    lambda: f64,
756    n_folds: usize,
757) -> Vec<f64> {
758    let (n, m) = data.shape();
759    // Cross-validate over TIME POINTS, not curves. Leaving out curves is
760    // ill-posed for per-curve basis fitting — each curve has its own
761    // coefficients, so a held-out curve can only be scored against its own data,
762    // which is an in-sample residual that decreases monotonically in `nbasis`
763    // and always selects the largest candidate (GH #33). Holding out points and
764    // predicting them from a fit on the remaining points gives a genuine
765    // predictive score that penalizes overfitting and shows an interior minimum.
766    let n_folds = n_folds.max(2).min(m);
767    let point_folds = crate::cv::create_folds(m, n_folds, 42);
768    let mut scores = Vec::with_capacity(nbasis_range.len());
769
770    for &nb in nbasis_range {
771        if nb < 2 {
772            scores.push(f64::INFINITY);
773            continue;
774        }
775        let penalty = match basis_type {
776            BasisType::Bspline { order } => bspline_penalty_matrix(argvals, nb, *order, 2),
777            BasisType::Fourier { period } => fourier_penalty_matrix(nb, *period, 2),
778        };
779        let (basis_flat, actual_k) = evaluate_basis(argvals, basis_type, nb);
780        let b_full = DMatrix::from_column_slice(m, actual_k, &basis_flat);
781        let r_mat = DMatrix::from_column_slice(actual_k, actual_k, &penalty);
782
783        let mut total_se = 0.0;
784        let mut count = 0usize;
785
786        for fold in 0..n_folds {
787            let (train_pts, test_pts) = crate::cv::fold_indices(&point_folds, fold);
788            if train_pts.is_empty() || test_pts.is_empty() {
789                continue;
790            }
791            // Basis rows at the train/test points, and the projection operator
792            // built from the training points only (independent of any curve).
793            let b_train = b_full.select_rows(train_pts.iter());
794            let b_test = b_full.select_rows(test_pts.iter());
795            let btb = b_train.transpose() * &b_train;
796            let ridge_eps = 1e-10;
797            let system: DMatrix<f64> =
798                &btb + lambda * &r_mat + ridge_eps * DMatrix::<f64>::identity(actual_k, actual_k);
799            let Some(system_inv) = invert_penalized_system(&system, actual_k) else {
800                continue;
801            };
802            let proj = &system_inv * b_train.transpose(); // (k x |train|)
803
804            for i in 0..n {
805                let y_train = nalgebra::DVector::from_iterator(
806                    train_pts.len(),
807                    train_pts.iter().map(|&j| data[(i, j)]),
808                );
809                let coefs = &proj * &y_train;
810                let pred = &b_test * &coefs; // predictions at held-out points
811                for (t_idx, &j) in test_pts.iter().enumerate() {
812                    let err = data[(i, j)] - pred[t_idx];
813                    total_se += err * err;
814                    count += 1;
815                }
816            }
817        }
818
819        if count > 0 {
820            scores.push(total_se / count as f64);
821        } else {
822            scores.push(f64::INFINITY);
823        }
824    }
825    scores
826}
827
828/// Select the optimal number of basis functions using multiple criteria
829/// (R's `fdata2basis_cv`).
830pub fn basis_nbasis_cv(
831    data: &FdMatrix,
832    argvals: &[f64],
833    nbasis_range: &[usize],
834    basis_type: &BasisType,
835    criterion: BasisCriterion,
836    n_folds: usize,
837    lambda: f64,
838) -> Option<BasisNbasisCvResult> {
839    let (n, m) = data.shape();
840    if n == 0 || m == 0 || argvals.len() != m || nbasis_range.is_empty() {
841        return None;
842    }
843
844    let scores = match criterion {
845        BasisCriterion::Gcv | BasisCriterion::Aic | BasisCriterion::Bic => {
846            evaluate_nbasis_info_criterion(
847                data,
848                argvals,
849                nbasis_range,
850                basis_type,
851                criterion,
852                lambda,
853            )
854        }
855        BasisCriterion::Cv => {
856            evaluate_nbasis_cv(data, argvals, nbasis_range, basis_type, lambda, n_folds)
857        }
858    };
859
860    let (best_idx, _) = scores
861        .iter()
862        .enumerate()
863        .min_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))?;
864
865    Some(BasisNbasisCvResult {
866        optimal_nbasis: nbasis_range[best_idx],
867        scores,
868        nbasis_range: nbasis_range.to_vec(),
869        criterion,
870    })
871}
872
873#[cfg(test)]
874mod tests {
875    use super::*;
876    use crate::test_helpers::uniform_grid;
877    use std::f64::consts::PI;
878
879    #[test]
880    fn test_bspline_penalty_matrix_symmetric() {
881        let t = uniform_grid(101);
882        let penalty = bspline_penalty_matrix(&t, 15, 4, 2);
883        let _k = 15; // may differ from actual due to knot construction
884        let actual_k = (penalty.len() as f64).sqrt() as usize;
885        for i in 0..actual_k {
886            for j in 0..actual_k {
887                assert!(
888                    (penalty[i + j * actual_k] - penalty[j + i * actual_k]).abs() < 1e-10,
889                    "Penalty matrix not symmetric at ({}, {})",
890                    i,
891                    j
892                );
893            }
894        }
895    }
896
897    #[test]
898    fn test_bspline_penalty_matrix_positive_semidefinite() {
899        let t = uniform_grid(101);
900        let penalty = bspline_penalty_matrix(&t, 10, 4, 2);
901        let k = (penalty.len() as f64).sqrt() as usize;
902        // Diagonal elements should be non-negative
903        for i in 0..k {
904            assert!(
905                penalty[i + i * k] >= -1e-10,
906                "Diagonal element {} is negative: {}",
907                i,
908                penalty[i + i * k]
909            );
910        }
911    }
912
913    #[test]
914    fn test_fourier_penalty_diagonal() {
915        let penalty = fourier_penalty_matrix(7, 1.0, 2);
916        // Should be diagonal
917        for i in 0..7 {
918            for j in 0..7 {
919                if i != j {
920                    assert!(
921                        penalty[i + j * 7].abs() < 1e-10,
922                        "Off-diagonal ({},{}) = {}",
923                        i,
924                        j,
925                        penalty[i + j * 7]
926                    );
927                }
928            }
929        }
930        // Constant term should have zero penalty
931        assert!(penalty[0].abs() < 1e-10);
932        // Higher frequency terms should have larger penalties
933        assert!(penalty[1 + 7] > 0.0);
934        assert!(penalty[3 + 3 * 7] > penalty[1 + 7]);
935    }
936
937    #[test]
938    fn test_smooth_basis_bspline() {
939        let m = 101;
940        let n = 5;
941        let t = uniform_grid(m);
942
943        // Generate noisy sine curves
944        let mut data = FdMatrix::zeros(n, m);
945        for i in 0..n {
946            for j in 0..m {
947                data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.1 * (i as f64 * 0.3 + j as f64 * 0.01);
948            }
949        }
950
951        let nbasis = 15;
952        let penalty = bspline_penalty_matrix(&t, nbasis, 4, 2);
953        let _actual_k = (penalty.len() as f64).sqrt() as usize;
954
955        let fdpar = FdPar {
956            basis_type: BasisType::Bspline { order: 4 },
957            nbasis,
958            lambda: 1e-4,
959            lfd_order: 2,
960            penalty_matrix: penalty,
961        };
962
963        let result = smooth_basis(&data, &t, &fdpar);
964        assert!(result.is_ok(), "smooth_basis should succeed");
965
966        let res = result.unwrap();
967        assert_eq!(res.fitted.shape(), (n, m));
968        assert_eq!(res.coefficients.nrows(), n);
969        assert!(res.edf > 0.0, "EDF should be positive");
970        assert!(res.gcv > 0.0, "GCV should be positive");
971    }
972
973    #[test]
974    fn test_smooth_basis_fourier() {
975        let m = 101;
976        let n = 3;
977        let t = uniform_grid(m);
978
979        let mut data = FdMatrix::zeros(n, m);
980        for i in 0..n {
981            for j in 0..m {
982                data[(i, j)] = (2.0 * PI * t[j]).sin() + (4.0 * PI * t[j]).cos();
983            }
984        }
985
986        let nbasis = 7;
987        let period = 1.0;
988        let penalty = fourier_penalty_matrix(nbasis, period, 2);
989
990        let fdpar = FdPar {
991            basis_type: BasisType::Fourier { period },
992            nbasis,
993            lambda: 1e-6,
994            lfd_order: 2,
995            penalty_matrix: penalty,
996        };
997
998        let result = smooth_basis(&data, &t, &fdpar);
999        assert!(result.is_ok());
1000
1001        let res = result.unwrap();
1002        // Fourier basis should fit periodic data well
1003        for j in 0..m {
1004            let expected = (2.0 * PI * t[j]).sin() + (4.0 * PI * t[j]).cos();
1005            assert!(
1006                (res.fitted[(0, j)] - expected).abs() < 0.1,
1007                "Fourier fit poor at j={}: got {}, expected {}",
1008                j,
1009                res.fitted[(0, j)],
1010                expected
1011            );
1012        }
1013    }
1014
1015    #[test]
1016    fn test_smooth_basis_gcv_selects_reasonable_lambda() {
1017        let m = 101;
1018        let n = 5;
1019        let t = uniform_grid(m);
1020
1021        let mut data = FdMatrix::zeros(n, m);
1022        for i in 0..n {
1023            for j in 0..m {
1024                data[(i, j)] =
1025                    (2.0 * PI * t[j]).sin() + 0.1 * ((i * 37 + j * 13) % 20) as f64 / 20.0;
1026            }
1027        }
1028
1029        let basis_type = BasisType::Bspline { order: 4 };
1030        let result = smooth_basis_gcv(&data, &t, &basis_type, 15, 2, (-8.0, 4.0), 25);
1031        assert!(result.is_some(), "GCV search should succeed");
1032    }
1033
1034    #[test]
1035    fn test_smooth_basis_aic_matches_brute_force_grid() {
1036        // smooth_basis_aic must return the SmoothBasisResult whose aic equals the
1037        // argmin of an explicit AIC grid built by calling smooth_basis at each
1038        // grid lambda and reading result.aic over the identical log-lambda grid.
1039        let m = 81;
1040        let n = 4;
1041        let t = uniform_grid(m);
1042        let mut data = FdMatrix::zeros(n, m);
1043        for i in 0..n {
1044            for j in 0..m {
1045                data[(i, j)] =
1046                    (2.0 * PI * t[j]).sin() + 0.15 * ((i * 41 + j * 7) % 23) as f64 / 23.0;
1047            }
1048        }
1049
1050        let basis_type = BasisType::Bspline { order: 4 };
1051        let nbasis = 12;
1052        let lfd_order = 2;
1053        let range = (-8.0, 4.0);
1054        let n_grid = 25;
1055
1056        // Brute-force the AIC grid using the same penalty + grid as the selector.
1057        let penalty = bspline_penalty_matrix(&t, nbasis, 4, lfd_order);
1058        let (lo, hi) = range;
1059        let mut brute_best_aic = f64::INFINITY;
1060        for k in 0..n_grid {
1061            let log_lam = lo + (hi - lo) * k as f64 / (n_grid - 1) as f64;
1062            let lam = 10.0_f64.powf(log_lam);
1063            let fdpar = FdPar {
1064                basis_type: basis_type.clone(),
1065                nbasis,
1066                lambda: lam,
1067                lfd_order,
1068                penalty_matrix: penalty.clone(),
1069            };
1070            if let Ok(result) = smooth_basis(&data, &t, &fdpar) {
1071                if result.aic < brute_best_aic {
1072                    brute_best_aic = result.aic;
1073                }
1074            }
1075        }
1076
1077        let selected =
1078            smooth_basis_aic(&data, &t, &basis_type, nbasis, lfd_order, range, n_grid).unwrap();
1079        assert!(
1080            (selected.aic - brute_best_aic).abs() < 1e-9,
1081            "selected aic={}, brute-force min aic={}",
1082            selected.aic,
1083            brute_best_aic
1084        );
1085    }
1086
1087    #[test]
1088    fn test_smooth_basis_aic_prefers_smoother_fit_than_smallest_lambda() {
1089        // On noisy data, the AIC-selected fit must have lower edf (be smoother)
1090        // than the smallest-lambda candidate from the same grid, which over-fits.
1091        let m = 81;
1092        let n = 4;
1093        let t = uniform_grid(m);
1094        let mut data = FdMatrix::zeros(n, m);
1095        // Smooth trend + deterministic pseudo-noise.
1096        for i in 0..n {
1097            for j in 0..m {
1098                let noise = ((((i * 91 + j * 53) % 101) as f64) / 101.0) - 0.5;
1099                data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.6 * noise;
1100            }
1101        }
1102
1103        let basis_type = BasisType::Bspline { order: 4 };
1104        let nbasis = 15;
1105        let lfd_order = 2;
1106        let range = (-8.0, 4.0);
1107        let n_grid = 25;
1108
1109        // Smallest-lambda candidate (most flexible / highest edf).
1110        let penalty = bspline_penalty_matrix(&t, nbasis, 4, lfd_order);
1111        let smallest_lam = 10.0_f64.powf(range.0);
1112        let fdpar_small = FdPar {
1113            basis_type: basis_type.clone(),
1114            nbasis,
1115            lambda: smallest_lam,
1116            lfd_order,
1117            penalty_matrix: penalty.clone(),
1118        };
1119        let overfit = smooth_basis(&data, &t, &fdpar_small).unwrap();
1120
1121        let selected =
1122            smooth_basis_aic(&data, &t, &basis_type, nbasis, lfd_order, range, n_grid).unwrap();
1123
1124        assert!(
1125            selected.edf < overfit.edf,
1126            "AIC-selected edf ({}) should be smaller (smoother) than the smallest-lambda edf ({})",
1127            selected.edf,
1128            overfit.edf
1129        );
1130    }
1131
1132    #[test]
1133    fn test_smooth_basis_large_lambda_reduces_edf() {
1134        let m = 101;
1135        let n = 3;
1136        let t = uniform_grid(m);
1137
1138        let mut data = FdMatrix::zeros(n, m);
1139        for i in 0..n {
1140            for j in 0..m {
1141                data[(i, j)] = (2.0 * PI * t[j]).sin();
1142            }
1143        }
1144
1145        let nbasis = 15;
1146        let penalty = bspline_penalty_matrix(&t, nbasis, 4, 2);
1147        let _actual_k = (penalty.len() as f64).sqrt() as usize;
1148
1149        let fdpar_small = FdPar {
1150            basis_type: BasisType::Bspline { order: 4 },
1151            nbasis,
1152            lambda: 1e-8,
1153            lfd_order: 2,
1154            penalty_matrix: penalty.clone(),
1155        };
1156        let fdpar_large = FdPar {
1157            basis_type: BasisType::Bspline { order: 4 },
1158            nbasis,
1159            lambda: 1e2,
1160            lfd_order: 2,
1161            penalty_matrix: penalty,
1162        };
1163
1164        let res_small = smooth_basis(&data, &t, &fdpar_small).unwrap();
1165        let res_large = smooth_basis(&data, &t, &fdpar_large).unwrap();
1166
1167        assert!(
1168            res_large.edf < res_small.edf,
1169            "Larger lambda should reduce EDF: {} vs {}",
1170            res_large.edf,
1171            res_small.edf
1172        );
1173    }
1174
1175    // ============== basis_nbasis_cv tests ==============
1176
1177    #[test]
1178    fn test_basis_nbasis_cv_gcv() {
1179        let m = 101;
1180        let n = 5;
1181        let t = uniform_grid(m);
1182        let mut data = FdMatrix::zeros(n, m);
1183        for i in 0..n {
1184            for j in 0..m {
1185                data[(i, j)] =
1186                    (2.0 * PI * t[j]).sin() + 0.1 * ((i * 37 + j * 13) % 20) as f64 / 20.0;
1187            }
1188        }
1189
1190        let nbasis_range: Vec<usize> = (4..=20).step_by(2).collect();
1191        let result = basis_nbasis_cv(
1192            &data,
1193            &t,
1194            &nbasis_range,
1195            &BasisType::Bspline { order: 4 },
1196            BasisCriterion::Gcv,
1197            5,
1198            1e-4,
1199        );
1200        assert!(result.is_some());
1201        let res = result.unwrap();
1202        assert!(nbasis_range.contains(&res.optimal_nbasis));
1203        assert_eq!(res.scores.len(), nbasis_range.len());
1204        assert_eq!(res.criterion, BasisCriterion::Gcv);
1205    }
1206
1207    #[test]
1208    fn test_basis_nbasis_cv_aic_bic() {
1209        let m = 51;
1210        let n = 5;
1211        let t = uniform_grid(m);
1212        let mut data = FdMatrix::zeros(n, m);
1213        for i in 0..n {
1214            for j in 0..m {
1215                data[(i, j)] = (2.0 * PI * t[j]).sin();
1216            }
1217        }
1218
1219        let nbasis_range: Vec<usize> = vec![5, 7, 9, 11];
1220        let aic_result = basis_nbasis_cv(
1221            &data,
1222            &t,
1223            &nbasis_range,
1224            &BasisType::Bspline { order: 4 },
1225            BasisCriterion::Aic,
1226            5,
1227            0.0,
1228        );
1229        let bic_result = basis_nbasis_cv(
1230            &data,
1231            &t,
1232            &nbasis_range,
1233            &BasisType::Bspline { order: 4 },
1234            BasisCriterion::Bic,
1235            5,
1236            0.0,
1237        );
1238        assert!(aic_result.is_some());
1239        assert!(bic_result.is_some());
1240    }
1241
1242    #[test]
1243    fn test_basis_nbasis_cv_kfold() {
1244        let m = 51;
1245        let n = 10;
1246        let t = uniform_grid(m);
1247        let mut data = FdMatrix::zeros(n, m);
1248        for i in 0..n {
1249            for j in 0..m {
1250                data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.05 * ((i * 7 + j * 3) % 10) as f64;
1251            }
1252        }
1253
1254        let nbasis_range: Vec<usize> = vec![5, 7, 9];
1255        let result = basis_nbasis_cv(
1256            &data,
1257            &t,
1258            &nbasis_range,
1259            &BasisType::Bspline { order: 4 },
1260            BasisCriterion::Cv,
1261            5,
1262            1e-4,
1263        );
1264        assert!(result.is_some());
1265        let res = result.unwrap();
1266        assert!(nbasis_range.contains(&res.optimal_nbasis));
1267        assert_eq!(res.criterion, BasisCriterion::Cv);
1268    }
1269
1270    /// Regression for GH #33: the CV path scored held-out curves against their
1271    /// own data (no true hold-out), so scores fell monotonically in `n_basis`
1272    /// and it always selected the maximum candidate. With point-wise CV,
1273    /// overfitting the noise is penalized and the maximum is not chosen.
1274    #[test]
1275    fn test_basis_nbasis_cv_penalizes_overfitting() {
1276        let m = 120;
1277        let n = 6;
1278        let t = uniform_grid(m);
1279        let mut data = FdMatrix::zeros(n, m);
1280        for i in 0..n {
1281            for j in 0..m {
1282                // Smooth signal + deterministic pseudo-noise.
1283                let noise = 0.2 * (((i * 31 + j * 17) % 13) as f64 / 13.0 - 0.5);
1284                data[(i, j)] = (2.0 * PI * t[j]).sin() + noise;
1285            }
1286        }
1287
1288        let nbasis_range: Vec<usize> = vec![5, 8, 12, 20, 30];
1289        let res = basis_nbasis_cv(
1290            &data,
1291            &t,
1292            &nbasis_range,
1293            &BasisType::Bspline { order: 4 },
1294            BasisCriterion::Cv,
1295            5,
1296            1e-6,
1297        )
1298        .unwrap();
1299
1300        assert_ne!(
1301            res.optimal_nbasis, 30,
1302            "CV must not always select the maximum n_basis (GH #33); scores={:?}",
1303            res.scores
1304        );
1305        let monotone_decreasing = res.scores.windows(2).all(|w| w[1] <= w[0] + 1e-12);
1306        assert!(
1307            !monotone_decreasing,
1308            "CV scores must not be monotone-decreasing in n_basis; scores={:?}",
1309            res.scores
1310        );
1311    }
1312
1313    // ============== Comprehensive additional tests ==============
1314
1315    // Helper: generate standard test data (sine + high-freq component)
1316    fn make_test_data(n: usize, m: usize) -> (FdMatrix, Vec<f64>) {
1317        let t: Vec<f64> = (0..m).map(|i| i as f64 / (m - 1) as f64).collect();
1318        let mut data = FdMatrix::zeros(n, m);
1319        for i in 0..n {
1320            for j in 0..m {
1321                data[(i, j)] = (2.0 * PI * t[j]).sin()
1322                    + 0.1 * (10.0 * t[j]).sin()
1323                    + 0.05 * ((i * 37 + j * 13) % 20) as f64 / 20.0;
1324            }
1325        }
1326        (data, t)
1327    }
1328
1329    // Helper: create an FdPar for B-spline smoothing
1330    fn make_bspline_fdpar(argvals: &[f64], nbasis: usize, lambda: f64) -> FdPar {
1331        let penalty = bspline_penalty_matrix(argvals, nbasis, 4, 2);
1332        FdPar {
1333            basis_type: BasisType::Bspline { order: 4 },
1334            nbasis,
1335            lambda,
1336            lfd_order: 2,
1337            penalty_matrix: penalty,
1338        }
1339    }
1340
1341    // Helper: create an FdPar for Fourier smoothing
1342    fn make_fourier_fdpar(nbasis: usize, period: f64, lambda: f64) -> FdPar {
1343        let penalty = fourier_penalty_matrix(nbasis, period, 2);
1344        FdPar {
1345            basis_type: BasisType::Fourier { period },
1346            nbasis,
1347            lambda,
1348            lfd_order: 2,
1349            penalty_matrix: penalty,
1350        }
1351    }
1352
1353    // ─── BasisType enum tests ───────────────────────────────────────────────
1354
1355    #[test]
1356    fn test_basis_type_bspline_variant() {
1357        let bt = BasisType::Bspline { order: 4 };
1358        assert_eq!(bt, BasisType::Bspline { order: 4 });
1359        // Different orders are not equal
1360        assert_ne!(bt, BasisType::Bspline { order: 3 });
1361    }
1362
1363    #[test]
1364    fn test_basis_type_fourier_variant() {
1365        let bt = BasisType::Fourier { period: 1.0 };
1366        assert_eq!(bt, BasisType::Fourier { period: 1.0 });
1367        assert_ne!(bt, BasisType::Fourier { period: 2.0 });
1368    }
1369
1370    #[test]
1371    fn test_basis_type_cross_variant_inequality() {
1372        let bspline = BasisType::Bspline { order: 4 };
1373        let fourier = BasisType::Fourier { period: 1.0 };
1374        assert_ne!(bspline, fourier);
1375    }
1376
1377    #[test]
1378    fn test_basis_type_clone_and_debug() {
1379        let bt = BasisType::Bspline { order: 4 };
1380        let cloned = bt.clone();
1381        assert_eq!(bt, cloned);
1382        let debug_str = format!("{:?}", bt);
1383        assert!(debug_str.contains("Bspline"));
1384        assert!(debug_str.contains("4"));
1385    }
1386
1387    // ─── FdPar struct tests ─────────────────────────────────────────────────
1388
1389    #[test]
1390    fn test_fdpar_construction_and_fields() {
1391        let penalty = vec![1.0, 0.0, 0.0, 1.0];
1392        let fdpar = FdPar {
1393            basis_type: BasisType::Bspline { order: 4 },
1394            nbasis: 2,
1395            lambda: 0.01,
1396            lfd_order: 2,
1397            penalty_matrix: penalty.clone(),
1398        };
1399        assert_eq!(fdpar.nbasis, 2);
1400        assert!((fdpar.lambda - 0.01).abs() < 1e-15);
1401        assert_eq!(fdpar.lfd_order, 2);
1402        assert_eq!(fdpar.penalty_matrix.len(), 4);
1403    }
1404
1405    #[test]
1406    fn test_fdpar_clone_and_debug() {
1407        let t = uniform_grid(50);
1408        let fdpar = make_bspline_fdpar(&t, 8, 1e-3);
1409        let cloned = fdpar.clone();
1410        assert_eq!(fdpar, cloned);
1411        let debug_str = format!("{:?}", fdpar);
1412        assert!(debug_str.contains("FdPar"));
1413    }
1414
1415    // ─── BasisCriterion enum tests ──────────────────────────────────────────
1416
1417    #[test]
1418    fn test_basis_criterion_variants() {
1419        assert_eq!(BasisCriterion::Gcv, BasisCriterion::Gcv);
1420        assert_eq!(BasisCriterion::Cv, BasisCriterion::Cv);
1421        assert_eq!(BasisCriterion::Aic, BasisCriterion::Aic);
1422        assert_eq!(BasisCriterion::Bic, BasisCriterion::Bic);
1423        assert_ne!(BasisCriterion::Gcv, BasisCriterion::Aic);
1424        assert_ne!(BasisCriterion::Cv, BasisCriterion::Bic);
1425    }
1426
1427    #[test]
1428    fn test_basis_criterion_copy() {
1429        let c = BasisCriterion::Gcv;
1430        let copied = c; // Copy
1431        assert_eq!(c, copied);
1432    }
1433
1434    #[test]
1435    fn test_basis_criterion_debug() {
1436        let debug_str = format!("{:?}", BasisCriterion::Bic);
1437        assert!(debug_str.contains("Bic"));
1438    }
1439
1440    // ─── SmoothBasisResult tests ────────────────────────────────────────────
1441
1442    #[test]
1443    fn test_smooth_basis_result_all_fields() {
1444        let (data, t) = make_test_data(3, 50);
1445        let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1446        let res = smooth_basis(&data, &t, &fdpar).unwrap();
1447
1448        // coefficients: n curves x k basis functions
1449        assert_eq!(res.coefficients.nrows(), 3);
1450        assert!(res.coefficients.ncols() > 0);
1451        assert_eq!(res.nbasis, res.coefficients.ncols());
1452        // fitted: n x m
1453        assert_eq!(res.fitted.shape(), (3, 50));
1454        // edf should be between 1 and nbasis
1455        assert!(res.edf > 0.0 && res.edf <= res.nbasis as f64);
1456        // gcv, aic, bic should be finite
1457        assert!(res.gcv.is_finite());
1458        assert!(res.aic.is_finite());
1459        assert!(res.bic.is_finite());
1460        // penalty_matrix should be k x k
1461        let k = res.nbasis;
1462        assert_eq!(res.penalty_matrix.len(), k * k);
1463    }
1464
1465    #[test]
1466    fn test_smooth_basis_result_clone() {
1467        let (data, t) = make_test_data(2, 50);
1468        let fdpar = make_bspline_fdpar(&t, 8, 1e-3);
1469        let res = smooth_basis(&data, &t, &fdpar).unwrap();
1470        let cloned = res.clone();
1471        assert_eq!(res, cloned);
1472    }
1473
1474    // ─── smooth_basis: B-spline detailed tests ──────────────────────────────
1475
1476    #[test]
1477    fn test_smooth_basis_bspline_coefficient_shape() {
1478        let (data, t) = make_test_data(4, 50);
1479        let nbasis = 12;
1480        let fdpar = make_bspline_fdpar(&t, nbasis, 1e-4);
1481        let res = smooth_basis(&data, &t, &fdpar).unwrap();
1482        assert_eq!(res.coefficients.nrows(), 4);
1483        // actual nbasis may differ from requested due to knot construction
1484        assert!(res.coefficients.ncols() >= 2);
1485        assert_eq!(res.nbasis, res.coefficients.ncols());
1486    }
1487
1488    #[test]
1489    fn test_smooth_basis_bspline_fitted_values_shape() {
1490        let m = 80;
1491        let n = 6;
1492        let (data, t) = make_test_data(n, m);
1493        let fdpar = make_bspline_fdpar(&t, 15, 1e-4);
1494        let res = smooth_basis(&data, &t, &fdpar).unwrap();
1495        assert_eq!(res.fitted.shape(), (n, m));
1496    }
1497
1498    #[test]
1499    fn test_smooth_basis_bspline_zero_lambda_interpolates() {
1500        // With lambda=0, the smoother should nearly interpolate the data
1501        let m = 30;
1502        let n = 2;
1503        let (data, t) = make_test_data(n, m);
1504        let fdpar = make_bspline_fdpar(&t, 15, 0.0);
1505        let res = smooth_basis(&data, &t, &fdpar).unwrap();
1506
1507        // Residuals should be very small (near interpolation)
1508        let mut max_resid = 0.0_f64;
1509        for i in 0..n {
1510            for j in 0..m {
1511                let resid = (data[(i, j)] - res.fitted[(i, j)]).abs();
1512                max_resid = max_resid.max(resid);
1513            }
1514        }
1515        assert!(
1516            max_resid < 0.5,
1517            "Zero-lambda B-spline should closely interpolate; max_resid = {}",
1518            max_resid
1519        );
1520    }
1521
1522    #[test]
1523    fn test_smooth_basis_bspline_large_lambda_oversmooths() {
1524        // With very large lambda, the fit should be much smoother (lower variance)
1525        // than with small lambda
1526        let m = 50;
1527        let n = 1;
1528        let (data, t) = make_test_data(n, m);
1529
1530        let fdpar_small = make_bspline_fdpar(&t, 15, 1e-6);
1531        let res_small = smooth_basis(&data, &t, &fdpar_small).unwrap();
1532
1533        let fdpar_large = make_bspline_fdpar(&t, 15, 1e6);
1534        let res_large = smooth_basis(&data, &t, &fdpar_large).unwrap();
1535
1536        let compute_variance = |fitted: &FdMatrix, row: usize, ncols: usize| -> f64 {
1537            let vals: Vec<f64> = (0..ncols).map(|j| fitted[(row, j)]).collect();
1538            let mean = vals.iter().sum::<f64>() / ncols as f64;
1539            vals.iter().map(|&v| (v - mean).powi(2)).sum::<f64>() / ncols as f64
1540        };
1541
1542        let var_small = compute_variance(&res_small.fitted, 0, m);
1543        let var_large = compute_variance(&res_large.fitted, 0, m);
1544        assert!(
1545            var_large < var_small,
1546            "Large lambda should yield lower variance fit: var_large={}, var_small={}",
1547            var_large,
1548            var_small
1549        );
1550    }
1551
1552    #[test]
1553    fn test_smooth_basis_bspline_penalty_effect_on_smoothness() {
1554        // Compare roughness of fits with small vs large lambda
1555        let m = 50;
1556        let n = 1;
1557        let (data, t) = make_test_data(n, m);
1558
1559        let fdpar_small = make_bspline_fdpar(&t, 15, 1e-8);
1560        let fdpar_large = make_bspline_fdpar(&t, 15, 1.0);
1561
1562        let res_small = smooth_basis(&data, &t, &fdpar_small).unwrap();
1563        let res_large = smooth_basis(&data, &t, &fdpar_large).unwrap();
1564
1565        // Measure roughness as sum of squared second differences
1566        let roughness = |fitted: &FdMatrix, row: usize, ncols: usize| -> f64 {
1567            (1..ncols - 1)
1568                .map(|j| {
1569                    let d2 = fitted[(row, j + 1)] - 2.0 * fitted[(row, j)] + fitted[(row, j - 1)];
1570                    d2 * d2
1571                })
1572                .sum::<f64>()
1573        };
1574
1575        let r_small = roughness(&res_small.fitted, 0, m);
1576        let r_large = roughness(&res_large.fitted, 0, m);
1577        assert!(
1578            r_large < r_small,
1579            "Larger lambda should produce smoother fit: roughness_large={}, roughness_small={}",
1580            r_large,
1581            r_small
1582        );
1583    }
1584
1585    #[test]
1586    fn test_smooth_basis_bspline_single_curve() {
1587        let m = 50;
1588        let (data, t) = make_test_data(1, m);
1589        let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1590        let res = smooth_basis(&data, &t, &fdpar).unwrap();
1591        assert_eq!(res.fitted.nrows(), 1);
1592        assert_eq!(res.fitted.ncols(), m);
1593        assert!(res.gcv.is_finite());
1594    }
1595
1596    #[test]
1597    fn test_smooth_basis_bspline_many_curves() {
1598        let m = 50;
1599        let n = 20;
1600        let (data, t) = make_test_data(n, m);
1601        let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1602        let res = smooth_basis(&data, &t, &fdpar).unwrap();
1603        assert_eq!(res.fitted.nrows(), n);
1604        assert_eq!(res.coefficients.nrows(), n);
1605    }
1606
1607    #[test]
1608    fn test_smooth_basis_bspline_minimal_nbasis() {
1609        // nbasis = 2 is the minimum allowed
1610        let m = 50;
1611        let (data, t) = make_test_data(1, m);
1612        let fdpar = make_bspline_fdpar(&t, 2, 1e-4);
1613        let res = smooth_basis(&data, &t, &fdpar);
1614        // Should succeed (or at least not panic); the fit may be poor
1615        assert!(res.is_ok());
1616    }
1617
1618    #[test]
1619    fn test_smooth_basis_bspline_different_orders() {
1620        let m = 50;
1621        let (data, t) = make_test_data(2, m);
1622        // Order 3 (quadratic B-splines)
1623        let penalty3 = bspline_penalty_matrix(&t, 10, 3, 2);
1624        let fdpar3 = FdPar {
1625            basis_type: BasisType::Bspline { order: 3 },
1626            nbasis: 10,
1627            lambda: 1e-4,
1628            lfd_order: 2,
1629            penalty_matrix: penalty3,
1630        };
1631        let res3 = smooth_basis(&data, &t, &fdpar3);
1632        assert!(res3.is_ok());
1633
1634        // Order 5 (quartic B-splines)
1635        let penalty5 = bspline_penalty_matrix(&t, 10, 5, 2);
1636        let fdpar5 = FdPar {
1637            basis_type: BasisType::Bspline { order: 5 },
1638            nbasis: 10,
1639            lambda: 1e-4,
1640            lfd_order: 2,
1641            penalty_matrix: penalty5,
1642        };
1643        let res5 = smooth_basis(&data, &t, &fdpar5);
1644        assert!(res5.is_ok());
1645    }
1646
1647    // ─── smooth_basis: Fourier detailed tests ───────────────────────────────
1648
1649    #[test]
1650    fn test_smooth_basis_fourier_coefficient_shape() {
1651        let m = 50;
1652        let n = 3;
1653        let t = uniform_grid(m);
1654        let mut data = FdMatrix::zeros(n, m);
1655        for i in 0..n {
1656            for j in 0..m {
1657                data[(i, j)] = (2.0 * PI * t[j]).sin();
1658            }
1659        }
1660        let nbasis = 7;
1661        let fdpar = make_fourier_fdpar(nbasis, 1.0, 1e-6);
1662        let res = smooth_basis(&data, &t, &fdpar).unwrap();
1663        assert_eq!(res.coefficients.nrows(), n);
1664        assert_eq!(res.coefficients.ncols(), nbasis);
1665        assert_eq!(res.nbasis, nbasis);
1666    }
1667
1668    #[test]
1669    fn test_smooth_basis_fourier_fits_pure_sine() {
1670        // Fourier basis should perfectly fit a pure sine with enough basis fns
1671        let m = 100;
1672        let t = uniform_grid(m);
1673        let mut data = FdMatrix::zeros(1, m);
1674        for j in 0..m {
1675            data[(0, j)] = (2.0 * PI * t[j]).sin();
1676        }
1677        let fdpar = make_fourier_fdpar(5, 1.0, 1e-8);
1678        let res = smooth_basis(&data, &t, &fdpar).unwrap();
1679
1680        for j in 0..m {
1681            let expected = (2.0 * PI * t[j]).sin();
1682            assert!(
1683                (res.fitted[(0, j)] - expected).abs() < 0.05,
1684                "Fourier should fit pure sine; j={}, got={}, expected={}",
1685                j,
1686                res.fitted[(0, j)],
1687                expected
1688            );
1689        }
1690    }
1691
1692    #[test]
1693    fn test_smooth_basis_fourier_different_periods() {
1694        let m = 50;
1695        let t = uniform_grid(m);
1696        let mut data = FdMatrix::zeros(1, m);
1697        for j in 0..m {
1698            data[(0, j)] = (2.0 * PI * t[j]).sin();
1699        }
1700
1701        // Period = 1.0 (matches the data)
1702        let fdpar1 = make_fourier_fdpar(7, 1.0, 1e-6);
1703        let res1 = smooth_basis(&data, &t, &fdpar1).unwrap();
1704
1705        // Period = 2.0 (mismatch, but should still produce a result)
1706        let fdpar2 = make_fourier_fdpar(7, 2.0, 1e-6);
1707        let res2 = smooth_basis(&data, &t, &fdpar2).unwrap();
1708
1709        // Both should succeed and have valid shapes
1710        assert_eq!(res1.fitted.shape(), (1, m));
1711        assert_eq!(res2.fitted.shape(), (1, m));
1712    }
1713
1714    #[test]
1715    fn test_smooth_basis_fourier_zero_lambda() {
1716        let m = 50;
1717        let t = uniform_grid(m);
1718        let mut data = FdMatrix::zeros(1, m);
1719        for j in 0..m {
1720            data[(0, j)] = (2.0 * PI * t[j]).sin() + (4.0 * PI * t[j]).cos();
1721        }
1722        let fdpar = make_fourier_fdpar(9, 1.0, 0.0);
1723        let res = smooth_basis(&data, &t, &fdpar).unwrap();
1724        assert_eq!(res.fitted.shape(), (1, m));
1725        // EDF should be close to nbasis with zero penalty
1726        assert!(res.edf > 1.0);
1727    }
1728
1729    #[test]
1730    fn test_smooth_basis_fourier_large_lambda() {
1731        let m = 50;
1732        let t = uniform_grid(m);
1733        let mut data = FdMatrix::zeros(1, m);
1734        for j in 0..m {
1735            data[(0, j)] = (2.0 * PI * t[j]).sin();
1736        }
1737        let fdpar = make_fourier_fdpar(9, 1.0, 1e6);
1738        let res = smooth_basis(&data, &t, &fdpar).unwrap();
1739        // EDF should be very small with huge penalty
1740        assert!(
1741            res.edf < 5.0,
1742            "Large lambda should reduce EDF; edf={}",
1743            res.edf
1744        );
1745    }
1746
1747    // ─── smooth_basis: Lambda comparison tests ──────────────────────────────
1748
1749    #[test]
1750    fn test_smooth_basis_lambda_gradient_edf() {
1751        // EDF should monotonically decrease with increasing lambda
1752        let m = 50;
1753        let (data, t) = make_test_data(3, m);
1754        let lambdas = [1e-8, 1e-4, 1e-2, 1.0, 1e2];
1755        let mut prev_edf = f64::INFINITY;
1756        for &lam in &lambdas {
1757            let fdpar = make_bspline_fdpar(&t, 12, lam);
1758            let res = smooth_basis(&data, &t, &fdpar).unwrap();
1759            assert!(
1760                res.edf <= prev_edf + 0.01,
1761                "EDF should decrease: lambda={}, edf={}, prev_edf={}",
1762                lam,
1763                res.edf,
1764                prev_edf
1765            );
1766            prev_edf = res.edf;
1767        }
1768    }
1769
1770    #[test]
1771    fn test_smooth_basis_lambda_gradient_rss() {
1772        // RSS should monotonically increase with increasing lambda
1773        let m = 50;
1774        let n = 2;
1775        let (data, t) = make_test_data(n, m);
1776        let lambdas = [0.0, 1e-6, 1e-2, 1.0, 1e4];
1777        let mut prev_rss = -1.0;
1778        for &lam in &lambdas {
1779            let fdpar = make_bspline_fdpar(&t, 12, lam);
1780            let res = smooth_basis(&data, &t, &fdpar).unwrap();
1781            let mut rss = 0.0;
1782            for i in 0..n {
1783                for j in 0..m {
1784                    rss += (data[(i, j)] - res.fitted[(i, j)]).powi(2);
1785                }
1786            }
1787            assert!(
1788                rss >= prev_rss - 1e-8,
1789                "RSS should increase: lambda={}, rss={}, prev_rss={}",
1790                lam,
1791                rss,
1792                prev_rss
1793            );
1794            prev_rss = rss;
1795        }
1796    }
1797
1798    // ─── smooth_basis: Error cases ──────────────────────────────────────────
1799
1800    #[test]
1801    fn test_smooth_basis_empty_data_rows() {
1802        let t = uniform_grid(50);
1803        let data = FdMatrix::zeros(0, 50);
1804        let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1805        let res = smooth_basis(&data, &t, &fdpar);
1806        assert!(res.is_err());
1807    }
1808
1809    #[test]
1810    fn test_smooth_basis_empty_data_cols() {
1811        let data = FdMatrix::zeros(5, 0);
1812        let fdpar = FdPar {
1813            basis_type: BasisType::Bspline { order: 4 },
1814            nbasis: 10,
1815            lambda: 1e-4,
1816            lfd_order: 2,
1817            penalty_matrix: vec![0.0; 100],
1818        };
1819        let res = smooth_basis(&data, &[], &fdpar);
1820        assert!(res.is_err());
1821    }
1822
1823    #[test]
1824    fn test_smooth_basis_mismatched_argvals() {
1825        let t = uniform_grid(50);
1826        let data = FdMatrix::zeros(3, 40); // m=40 but argvals has 50
1827        let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1828        let res = smooth_basis(&data, &t, &fdpar);
1829        assert!(res.is_err());
1830    }
1831
1832    #[test]
1833    fn test_smooth_basis_nbasis_too_small() {
1834        let t = uniform_grid(50);
1835        let data = FdMatrix::zeros(3, 50);
1836        // nbasis = 1, which is below minimum of 2
1837        let fdpar = FdPar {
1838            basis_type: BasisType::Bspline { order: 4 },
1839            nbasis: 1,
1840            lambda: 1e-4,
1841            lfd_order: 2,
1842            penalty_matrix: vec![0.0; 1],
1843        };
1844        let res = smooth_basis(&data, &t, &fdpar);
1845        assert!(res.is_err());
1846    }
1847
1848    #[test]
1849    fn test_smooth_basis_error_is_invalid_dimension() {
1850        let t = uniform_grid(50);
1851        let data = FdMatrix::zeros(0, 50);
1852        let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
1853        let err = smooth_basis(&data, &t, &fdpar).unwrap_err();
1854        match err {
1855            crate::FdarError::InvalidDimension { .. } => {} // expected
1856            other => panic!("Expected InvalidDimension, got {:?}", other),
1857        }
1858    }
1859
1860    // ─── Penalty matrix detailed tests ──────────────────────────────────────
1861
1862    #[test]
1863    fn test_bspline_penalty_matrix_different_orders() {
1864        let t = uniform_grid(101);
1865        // Order 1 penalty (penalize derivatives)
1866        let p1 = bspline_penalty_matrix(&t, 10, 4, 1);
1867        // Order 2 penalty (penalize curvature)
1868        let p2 = bspline_penalty_matrix(&t, 10, 4, 2);
1869        // Both should be square and same size
1870        assert_eq!(p1.len(), p2.len());
1871        // But they should differ
1872        let diff: f64 = p1.iter().zip(p2.iter()).map(|(a, b)| (a - b).abs()).sum();
1873        assert!(
1874            diff > 1e-10,
1875            "Different lfd_orders should produce different penalties"
1876        );
1877    }
1878
1879    #[test]
1880    fn test_bspline_penalty_matrix_edge_cases() {
1881        // Too few argvals
1882        let t = vec![0.0];
1883        let p = bspline_penalty_matrix(&t, 10, 4, 2);
1884        // Should return zero matrix
1885        assert!(p.iter().all(|&v| v == 0.0));
1886
1887        // nbasis < 2
1888        let t2 = uniform_grid(50);
1889        let p2 = bspline_penalty_matrix(&t2, 1, 4, 2);
1890        assert!(p2.iter().all(|&v| v == 0.0));
1891
1892        // lfd_order >= order
1893        let p3 = bspline_penalty_matrix(&t2, 10, 4, 4);
1894        assert!(p3.iter().all(|&v| v == 0.0));
1895    }
1896
1897    #[test]
1898    fn test_bspline_penalty_nonnegative_diagonal() {
1899        let t = uniform_grid(101);
1900        for nbasis in [5, 10, 20] {
1901            let p = bspline_penalty_matrix(&t, nbasis, 4, 2);
1902            let k = (p.len() as f64).sqrt() as usize;
1903            for i in 0..k {
1904                assert!(
1905                    p[i + i * k] >= -1e-10,
1906                    "Diagonal ({},{}) negative for nbasis={}: {}",
1907                    i,
1908                    i,
1909                    nbasis,
1910                    p[i + i * k]
1911                );
1912            }
1913        }
1914    }
1915
1916    #[test]
1917    fn test_fourier_penalty_increasing_with_frequency() {
1918        let penalty = fourier_penalty_matrix(11, 1.0, 2);
1919        let k = 11;
1920        // Constant term is zero
1921        assert!(penalty[0].abs() < 1e-15);
1922        // Pairs: (1,2) -> freq 1, (3,4) -> freq 2, etc.
1923        let mut prev_eigenval = 0.0;
1924        for freq in 1..=5 {
1925            let idx_sin = 2 * freq - 1;
1926            let eigenval = penalty[idx_sin + idx_sin * k];
1927            assert!(
1928                eigenval > prev_eigenval,
1929                "Higher frequency should have larger penalty: freq={}, eigenval={}, prev={}",
1930                freq,
1931                eigenval,
1932                prev_eigenval
1933            );
1934            prev_eigenval = eigenval;
1935            // cos and sin of same frequency should have same penalty
1936            let idx_cos = 2 * freq;
1937            if idx_cos < k {
1938                assert!(
1939                    (penalty[idx_cos + idx_cos * k] - eigenval).abs() < 1e-10,
1940                    "Sin and cos penalty should match at freq {}",
1941                    freq
1942                );
1943            }
1944        }
1945    }
1946
1947    #[test]
1948    fn test_fourier_penalty_different_periods() {
1949        let p1 = fourier_penalty_matrix(7, 1.0, 2);
1950        let p2 = fourier_penalty_matrix(7, 2.0, 2);
1951        // Longer period -> smaller omega -> smaller penalty eigenvalues
1952        for i in 1..7 {
1953            assert!(
1954                p2[i + i * 7] < p1[i + i * 7] || (p1[i + i * 7] == 0.0 && p2[i + i * 7] == 0.0),
1955                "Longer period should have smaller penalties at i={}",
1956                i
1957            );
1958        }
1959    }
1960
1961    #[test]
1962    fn test_fourier_penalty_first_order() {
1963        // lfd_order = 1: penalize first derivative
1964        let p = fourier_penalty_matrix(5, 1.0, 1);
1965        // Eigenvalues: (2*pi*freq)^2 for lfd_order=1
1966        let omega1 = 2.0 * PI;
1967        let expected1 = omega1.powi(2);
1968        assert!(
1969            (p[1 + 5] - expected1).abs() < 1e-6,
1970            "First-order penalty eigenval: got {}, expected {}",
1971            p[1 + 5],
1972            expected1
1973        );
1974    }
1975
1976    #[test]
1977    fn test_fourier_penalty_zero_nbasis() {
1978        let p = fourier_penalty_matrix(0, 1.0, 2);
1979        assert!(p.is_empty());
1980    }
1981
1982    #[test]
1983    fn test_fourier_penalty_nbasis_one() {
1984        let p = fourier_penalty_matrix(1, 1.0, 2);
1985        assert_eq!(p.len(), 1);
1986        assert!(p[0].abs() < 1e-15); // constant term has zero penalty
1987    }
1988
1989    // ─── smooth_basis_gcv detailed tests ────────────────────────────────────
1990
1991    #[test]
1992    fn test_smooth_basis_gcv_returns_valid_result() {
1993        let (data, t) = make_test_data(5, 50);
1994        let bt = BasisType::Bspline { order: 4 };
1995        let result = smooth_basis_gcv(&data, &t, &bt, 12, 2, (-6.0, 2.0), 20);
1996        assert!(result.is_some());
1997        let res = result.unwrap();
1998        assert_eq!(res.fitted.shape(), (5, 50));
1999        assert!(res.gcv.is_finite());
2000        assert!(res.edf > 0.0);
2001    }
2002
2003    #[test]
2004    fn test_smooth_basis_gcv_fourier() {
2005        let m = 80;
2006        let t = uniform_grid(m);
2007        let mut data = FdMatrix::zeros(3, m);
2008        for i in 0..3 {
2009            for j in 0..m {
2010                data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.5 * (4.0 * PI * t[j]).cos();
2011            }
2012        }
2013        let bt = BasisType::Fourier { period: 1.0 };
2014        let result = smooth_basis_gcv(&data, &t, &bt, 9, 2, (-8.0, 4.0), 25);
2015        assert!(result.is_some());
2016        let res = result.unwrap();
2017        assert_eq!(res.fitted.nrows(), 3);
2018        assert_eq!(res.nbasis, 9);
2019    }
2020
2021    #[test]
2022    fn test_smooth_basis_gcv_selects_finite_gcv() {
2023        let (data, t) = make_test_data(5, 60);
2024        let bt = BasisType::Bspline { order: 4 };
2025        let res = smooth_basis_gcv(&data, &t, &bt, 12, 2, (-6.0, 2.0), 15).unwrap();
2026        assert!(res.gcv.is_finite());
2027        assert!(res.gcv > 0.0);
2028    }
2029
2030    #[test]
2031    fn test_smooth_basis_gcv_empty_data() {
2032        let data = FdMatrix::zeros(0, 50);
2033        let t = uniform_grid(50);
2034        let bt = BasisType::Bspline { order: 4 };
2035        let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-6.0, 2.0), 10);
2036        // Should return None since smooth_basis will error for empty data
2037        assert!(result.is_none());
2038    }
2039
2040    #[test]
2041    fn test_smooth_basis_gcv_empty_argvals() {
2042        let data = FdMatrix::zeros(5, 0);
2043        let bt = BasisType::Bspline { order: 4 };
2044        let result = smooth_basis_gcv(&data, &[], &bt, 10, 2, (-6.0, 2.0), 10);
2045        assert!(result.is_none());
2046    }
2047
2048    #[test]
2049    fn test_smooth_basis_gcv_nbasis_too_small() {
2050        let (data, t) = make_test_data(5, 50);
2051        let bt = BasisType::Bspline { order: 4 };
2052        let result = smooth_basis_gcv(&data, &t, &bt, 1, 2, (-6.0, 2.0), 10);
2053        assert!(result.is_none());
2054    }
2055
2056    #[test]
2057    fn test_smooth_basis_gcv_ngrid_too_small() {
2058        let (data, t) = make_test_data(5, 50);
2059        let bt = BasisType::Bspline { order: 4 };
2060        let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-6.0, 2.0), 1);
2061        assert!(result.is_none());
2062    }
2063
2064    #[test]
2065    fn test_smooth_basis_gcv_narrow_range() {
2066        let (data, t) = make_test_data(3, 50);
2067        let bt = BasisType::Bspline { order: 4 };
2068        // Very narrow search range
2069        let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-3.0, -2.0), 5);
2070        assert!(result.is_some());
2071    }
2072
2073    #[test]
2074    fn test_smooth_basis_gcv_wide_range() {
2075        let (data, t) = make_test_data(3, 50);
2076        let bt = BasisType::Bspline { order: 4 };
2077        // Very wide search range
2078        let result = smooth_basis_gcv(&data, &t, &bt, 10, 2, (-12.0, 8.0), 30);
2079        assert!(result.is_some());
2080    }
2081
2082    // ─── basis_nbasis_cv detailed tests ─────────────────────────────────────
2083
2084    #[test]
2085    fn test_basis_nbasis_cv_scores_length() {
2086        let (data, t) = make_test_data(5, 50);
2087        let nbasis_range: Vec<usize> = vec![4, 6, 8, 10, 12];
2088        let res = basis_nbasis_cv(
2089            &data,
2090            &t,
2091            &nbasis_range,
2092            &BasisType::Bspline { order: 4 },
2093            BasisCriterion::Gcv,
2094            5,
2095            1e-4,
2096        )
2097        .unwrap();
2098        assert_eq!(res.scores.len(), 5);
2099        assert_eq!(res.nbasis_range.len(), 5);
2100        assert_eq!(res.nbasis_range, nbasis_range);
2101    }
2102
2103    #[test]
2104    fn test_basis_nbasis_cv_optimal_within_range() {
2105        let (data, t) = make_test_data(8, 50);
2106        let nbasis_range: Vec<usize> = vec![5, 7, 9, 11, 13, 15];
2107        for criterion in [
2108            BasisCriterion::Gcv,
2109            BasisCriterion::Aic,
2110            BasisCriterion::Bic,
2111        ] {
2112            let res = basis_nbasis_cv(
2113                &data,
2114                &t,
2115                &nbasis_range,
2116                &BasisType::Bspline { order: 4 },
2117                criterion,
2118                5,
2119                1e-4,
2120            )
2121            .unwrap();
2122            assert!(
2123                nbasis_range.contains(&res.optimal_nbasis),
2124                "optimal_nbasis {} not in range for {:?}",
2125                res.optimal_nbasis,
2126                criterion
2127            );
2128        }
2129    }
2130
2131    #[test]
2132    fn test_basis_nbasis_cv_fourier_gcv() {
2133        let m = 80;
2134        let t = uniform_grid(m);
2135        let mut data = FdMatrix::zeros(5, m);
2136        for i in 0..5 {
2137            for j in 0..m {
2138                data[(i, j)] = (2.0 * PI * t[j]).sin()
2139                    + 0.3 * (4.0 * PI * t[j]).cos()
2140                    + 0.02 * ((i * 7 + j * 3) % 10) as f64;
2141            }
2142        }
2143        let nbasis_range: Vec<usize> = vec![5, 7, 9, 11];
2144        let res = basis_nbasis_cv(
2145            &data,
2146            &t,
2147            &nbasis_range,
2148            &BasisType::Fourier { period: 1.0 },
2149            BasisCriterion::Gcv,
2150            5,
2151            1e-4,
2152        )
2153        .unwrap();
2154        assert!(nbasis_range.contains(&res.optimal_nbasis));
2155    }
2156
2157    #[test]
2158    fn test_basis_nbasis_cv_fourier_cv() {
2159        let m = 60;
2160        let t = uniform_grid(m);
2161        let n = 10;
2162        let mut data = FdMatrix::zeros(n, m);
2163        for i in 0..n {
2164            for j in 0..m {
2165                data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.02 * ((i * 11 + j) % 15) as f64;
2166            }
2167        }
2168        let nbasis_range: Vec<usize> = vec![5, 7, 9];
2169        let res = basis_nbasis_cv(
2170            &data,
2171            &t,
2172            &nbasis_range,
2173            &BasisType::Fourier { period: 1.0 },
2174            BasisCriterion::Cv,
2175            5,
2176            1e-4,
2177        )
2178        .unwrap();
2179        assert!(nbasis_range.contains(&res.optimal_nbasis));
2180        assert_eq!(res.criterion, BasisCriterion::Cv);
2181    }
2182
2183    #[test]
2184    fn test_basis_nbasis_cv_with_nbasis_below_minimum() {
2185        // Range includes nbasis = 1 which is invalid
2186        let (data, t) = make_test_data(5, 50);
2187        let nbasis_range: Vec<usize> = vec![1, 5, 10];
2188        let res = basis_nbasis_cv(
2189            &data,
2190            &t,
2191            &nbasis_range,
2192            &BasisType::Bspline { order: 4 },
2193            BasisCriterion::Gcv,
2194            5,
2195            1e-4,
2196        )
2197        .unwrap();
2198        // Score for nbasis=1 should be infinity, so optimal should be 5 or 10
2199        assert!(
2200            res.optimal_nbasis >= 5,
2201            "Should skip invalid nbasis=1, got optimal={}",
2202            res.optimal_nbasis
2203        );
2204        assert!(res.scores[0].is_infinite());
2205    }
2206
2207    #[test]
2208    fn test_basis_nbasis_cv_empty_range() {
2209        let (data, t) = make_test_data(5, 50);
2210        let nbasis_range: Vec<usize> = vec![];
2211        let result = basis_nbasis_cv(
2212            &data,
2213            &t,
2214            &nbasis_range,
2215            &BasisType::Bspline { order: 4 },
2216            BasisCriterion::Gcv,
2217            5,
2218            1e-4,
2219        );
2220        assert!(result.is_none());
2221    }
2222
2223    #[test]
2224    fn test_basis_nbasis_cv_empty_data() {
2225        let data = FdMatrix::zeros(0, 50);
2226        let t = uniform_grid(50);
2227        let nbasis_range: Vec<usize> = vec![5, 10];
2228        let result = basis_nbasis_cv(
2229            &data,
2230            &t,
2231            &nbasis_range,
2232            &BasisType::Bspline { order: 4 },
2233            BasisCriterion::Gcv,
2234            5,
2235            1e-4,
2236        );
2237        assert!(result.is_none());
2238    }
2239
2240    #[test]
2241    fn test_basis_nbasis_cv_mismatched_argvals() {
2242        let data = FdMatrix::zeros(5, 50);
2243        let t = uniform_grid(40); // mismatch
2244        let nbasis_range: Vec<usize> = vec![5, 10];
2245        let result = basis_nbasis_cv(
2246            &data,
2247            &t,
2248            &nbasis_range,
2249            &BasisType::Bspline { order: 4 },
2250            BasisCriterion::Gcv,
2251            5,
2252            1e-4,
2253        );
2254        assert!(result.is_none());
2255    }
2256
2257    #[test]
2258    fn test_basis_nbasis_cv_single_nbasis() {
2259        let (data, t) = make_test_data(5, 50);
2260        let nbasis_range: Vec<usize> = vec![10];
2261        let res = basis_nbasis_cv(
2262            &data,
2263            &t,
2264            &nbasis_range,
2265            &BasisType::Bspline { order: 4 },
2266            BasisCriterion::Gcv,
2267            5,
2268            1e-4,
2269        )
2270        .unwrap();
2271        assert_eq!(res.optimal_nbasis, 10);
2272        assert_eq!(res.scores.len(), 1);
2273    }
2274
2275    #[test]
2276    fn test_basis_nbasis_cv_bic_penalizes_more_than_aic() {
2277        // BIC penalizes complexity more heavily than AIC, so it should generally
2278        // select the same or fewer basis functions
2279        let (data, t) = make_test_data(5, 80);
2280        let nbasis_range: Vec<usize> = (4..=20).step_by(2).collect();
2281
2282        let aic_res = basis_nbasis_cv(
2283            &data,
2284            &t,
2285            &nbasis_range,
2286            &BasisType::Bspline { order: 4 },
2287            BasisCriterion::Aic,
2288            5,
2289            1e-4,
2290        )
2291        .unwrap();
2292        let bic_res = basis_nbasis_cv(
2293            &data,
2294            &t,
2295            &nbasis_range,
2296            &BasisType::Bspline { order: 4 },
2297            BasisCriterion::Bic,
2298            5,
2299            1e-4,
2300        )
2301        .unwrap();
2302        // BIC should select at most as many basis functions as AIC
2303        // (not guaranteed in all cases, but typical behavior)
2304        assert!(
2305            bic_res.optimal_nbasis <= aic_res.optimal_nbasis + 4,
2306            "BIC selected {} vs AIC selected {} -- BIC should not select much more than AIC",
2307            bic_res.optimal_nbasis,
2308            aic_res.optimal_nbasis
2309        );
2310    }
2311
2312    // ─── Fitted values quality tests ────────────────────────────────────────
2313
2314    #[test]
2315    fn test_smooth_basis_fitted_close_to_data() {
2316        // With moderate penalty and enough basis functions, fitted should be close to data
2317        let m = 50;
2318        let n = 3;
2319        let t = uniform_grid(m);
2320        let mut data = FdMatrix::zeros(n, m);
2321        for i in 0..n {
2322            for j in 0..m {
2323                data[(i, j)] = (2.0 * PI * t[j]).sin();
2324            }
2325        }
2326        let fdpar = make_bspline_fdpar(&t, 15, 1e-6);
2327        let res = smooth_basis(&data, &t, &fdpar).unwrap();
2328
2329        let mut max_err = 0.0_f64;
2330        for i in 0..n {
2331            for j in 0..m {
2332                let err = (data[(i, j)] - res.fitted[(i, j)]).abs();
2333                max_err = max_err.max(err);
2334            }
2335        }
2336        assert!(
2337            max_err < 0.1,
2338            "Fitted should be close to smooth data; max_err={}",
2339            max_err
2340        );
2341    }
2342
2343    #[test]
2344    fn test_smooth_basis_constant_data() {
2345        // Constant data should be fit exactly
2346        let m = 50;
2347        let n = 2;
2348        let t = uniform_grid(m);
2349        let mut data = FdMatrix::zeros(n, m);
2350        for i in 0..n {
2351            for j in 0..m {
2352                data[(i, j)] = 3.15;
2353            }
2354        }
2355        let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2356        let res = smooth_basis(&data, &t, &fdpar).unwrap();
2357        for i in 0..n {
2358            for j in 0..m {
2359                assert!(
2360                    (res.fitted[(i, j)] - 3.15).abs() < 0.01,
2361                    "Constant data should be fit well at ({},{}): got {}",
2362                    i,
2363                    j,
2364                    res.fitted[(i, j)]
2365                );
2366            }
2367        }
2368    }
2369
2370    #[test]
2371    fn test_smooth_basis_linear_data() {
2372        // Linear data should be fit well with cubic B-splines
2373        let m = 50;
2374        let t = uniform_grid(m);
2375        let mut data = FdMatrix::zeros(1, m);
2376        for j in 0..m {
2377            data[(0, j)] = 2.0 * t[j] + 1.0;
2378        }
2379        let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2380        let res = smooth_basis(&data, &t, &fdpar).unwrap();
2381        for j in 0..m {
2382            let expected = 2.0 * t[j] + 1.0;
2383            assert!(
2384                (res.fitted[(0, j)] - expected).abs() < 0.05,
2385                "Linear data should be fit well at j={}: got {}, expected {}",
2386                j,
2387                res.fitted[(0, j)],
2388                expected
2389            );
2390        }
2391    }
2392
2393    // ─── EDF and diagnostic tests ───────────────────────────────────────────
2394
2395    #[test]
2396    fn test_smooth_basis_edf_bounded() {
2397        let m = 50;
2398        let (data, t) = make_test_data(3, m);
2399        let fdpar = make_bspline_fdpar(&t, 12, 1e-4);
2400        let res = smooth_basis(&data, &t, &fdpar).unwrap();
2401        // EDF should be between 1 and m (evaluation points)
2402        assert!(
2403            res.edf > 0.0 && res.edf <= m as f64,
2404            "EDF should be in (0, {}]; got {}",
2405            m,
2406            res.edf
2407        );
2408    }
2409
2410    #[test]
2411    fn test_smooth_basis_gcv_aic_bic_all_finite() {
2412        let (data, t) = make_test_data(4, 60);
2413        let fdpar = make_bspline_fdpar(&t, 12, 1e-3);
2414        let res = smooth_basis(&data, &t, &fdpar).unwrap();
2415        assert!(res.gcv.is_finite(), "GCV should be finite: {}", res.gcv);
2416        assert!(res.aic.is_finite(), "AIC should be finite: {}", res.aic);
2417        assert!(res.bic.is_finite(), "BIC should be finite: {}", res.bic);
2418    }
2419
2420    // ─── Penalty matrix size consistency tests ──────────────────────────────
2421
2422    #[test]
2423    fn test_smooth_basis_penalty_matrix_in_result() {
2424        let (data, t) = make_test_data(3, 50);
2425        let nbasis = 10;
2426        let fdpar = make_bspline_fdpar(&t, nbasis, 1e-4);
2427        let res = smooth_basis(&data, &t, &fdpar).unwrap();
2428        let k = res.nbasis;
2429        assert_eq!(
2430            res.penalty_matrix.len(),
2431            k * k,
2432            "Penalty matrix should be k*k = {}*{} = {}; got {}",
2433            k,
2434            k,
2435            k * k,
2436            res.penalty_matrix.len()
2437        );
2438    }
2439
2440    // ─── Regression: multiple identical curves ──────────────────────────────
2441
2442    #[test]
2443    fn test_smooth_basis_identical_curves_same_coefficients() {
2444        let m = 50;
2445        let t = uniform_grid(m);
2446        let curve: Vec<f64> = (0..m).map(|j| (2.0 * PI * t[j]).sin()).collect();
2447        let n = 4;
2448        let mut data = FdMatrix::zeros(n, m);
2449        for i in 0..n {
2450            for j in 0..m {
2451                data[(i, j)] = curve[j];
2452            }
2453        }
2454        let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2455        let res = smooth_basis(&data, &t, &fdpar).unwrap();
2456
2457        // All curves should have the same coefficients
2458        let k = res.coefficients.ncols();
2459        for i in 1..n {
2460            for j in 0..k {
2461                assert!(
2462                    (res.coefficients[(i, j)] - res.coefficients[(0, j)]).abs() < 1e-10,
2463                    "Identical curves should have identical coefficients: curve {} col {} differs",
2464                    i,
2465                    j
2466                );
2467            }
2468        }
2469    }
2470
2471    // ─── Cross-validation: different numbers of folds ───────────────────────
2472
2473    #[test]
2474    fn test_basis_nbasis_cv_different_nfolds() {
2475        let (data, t) = make_test_data(12, 50);
2476        let nbasis_range: Vec<usize> = vec![5, 8, 11];
2477        for nfolds in [2, 3, 5, 10] {
2478            let res = basis_nbasis_cv(
2479                &data,
2480                &t,
2481                &nbasis_range,
2482                &BasisType::Bspline { order: 4 },
2483                BasisCriterion::Cv,
2484                nfolds,
2485                1e-4,
2486            );
2487            assert!(res.is_some(), "CV should succeed with nfolds={}", nfolds);
2488            let r = res.unwrap();
2489            assert!(nbasis_range.contains(&r.optimal_nbasis));
2490        }
2491    }
2492
2493    // ─── Large nbasis / more basis than reasonable ──────────────────────────
2494
2495    #[test]
2496    fn test_smooth_basis_many_basis_functions() {
2497        let m = 100;
2498        let (data, t) = make_test_data(2, m);
2499        // Many basis functions relative to data points
2500        let fdpar = make_bspline_fdpar(&t, 40, 1e-2);
2501        let res = smooth_basis(&data, &t, &fdpar);
2502        assert!(
2503            res.is_ok(),
2504            "Should handle many basis functions with penalty"
2505        );
2506    }
2507
2508    // ─── evaluate_basis internal function (indirectly tested) ───────────────
2509
2510    #[test]
2511    fn test_smooth_basis_bspline_vs_fourier_different_results() {
2512        let m = 50;
2513        let (data, t) = make_test_data(2, m);
2514        let fdpar_bs = make_bspline_fdpar(&t, 9, 1e-4);
2515        let fdpar_f = make_fourier_fdpar(9, 1.0, 1e-4);
2516        let res_bs = smooth_basis(&data, &t, &fdpar_bs).unwrap();
2517        let res_f = smooth_basis(&data, &t, &fdpar_f).unwrap();
2518        // Results should differ between the two basis types
2519        let diff: f64 = (0..m)
2520            .map(|j| (res_bs.fitted[(0, j)] - res_f.fitted[(0, j)]).abs())
2521            .sum();
2522        // They fit the same data, so some difference is expected but not huge
2523        assert!(
2524            diff > 1e-10,
2525            "B-spline and Fourier fits should differ for the same data"
2526        );
2527    }
2528
2529    // ─── compute_gcv edge cases (indirectly tested) ─────────────────────────
2530
2531    #[test]
2532    fn test_smooth_basis_gcv_positive_for_noisy_data() {
2533        let m = 50;
2534        let t = uniform_grid(m);
2535        let mut data = FdMatrix::zeros(1, m);
2536        for j in 0..m {
2537            // Noisy data
2538            data[(0, j)] = (2.0 * PI * t[j]).sin() + 0.5 * ((j * 37) % 20) as f64 / 20.0 - 0.25;
2539        }
2540        let fdpar = make_bspline_fdpar(&t, 10, 1e-3);
2541        let res = smooth_basis(&data, &t, &fdpar).unwrap();
2542        assert!(res.gcv > 0.0, "GCV should be positive for noisy data");
2543    }
2544
2545    // ─── Penalty order (lfd_order) tests ────────────────────────────────────
2546
2547    #[test]
2548    fn test_smooth_basis_different_lfd_orders() {
2549        let m = 50;
2550        let (data, t) = make_test_data(2, m);
2551
2552        // lfd_order = 1 (penalize first derivative)
2553        let penalty1 = bspline_penalty_matrix(&t, 10, 4, 1);
2554        let fdpar1 = FdPar {
2555            basis_type: BasisType::Bspline { order: 4 },
2556            nbasis: 10,
2557            lambda: 1e-2,
2558            lfd_order: 1,
2559            penalty_matrix: penalty1,
2560        };
2561        let res1 = smooth_basis(&data, &t, &fdpar1);
2562        assert!(res1.is_ok());
2563
2564        // lfd_order = 2 (penalize second derivative)
2565        let penalty2 = bspline_penalty_matrix(&t, 10, 4, 2);
2566        let fdpar2 = FdPar {
2567            basis_type: BasisType::Bspline { order: 4 },
2568            nbasis: 10,
2569            lambda: 1e-2,
2570            lfd_order: 2,
2571            penalty_matrix: penalty2,
2572        };
2573        let res2 = smooth_basis(&data, &t, &fdpar2);
2574        assert!(res2.is_ok());
2575
2576        // Different penalty orders should produce different fitted values
2577        let r1 = res1.unwrap();
2578        let r2 = res2.unwrap();
2579        let diff: f64 = (0..m)
2580            .map(|j| (r1.fitted[(0, j)] - r2.fitted[(0, j)]).abs())
2581            .sum();
2582        assert!(
2583            diff > 1e-10,
2584            "Different lfd_orders should produce different fits"
2585        );
2586    }
2587
2588    // ─── BasisNbasisCvResult field tests ────────────────────────────────────
2589
2590    #[test]
2591    fn test_basis_nbasis_cv_result_fields() {
2592        let (data, t) = make_test_data(6, 50);
2593        let nbasis_range: Vec<usize> = vec![5, 7, 9, 11, 13];
2594        let res = basis_nbasis_cv(
2595            &data,
2596            &t,
2597            &nbasis_range,
2598            &BasisType::Bspline { order: 4 },
2599            BasisCriterion::Aic,
2600            5,
2601            1e-4,
2602        )
2603        .unwrap();
2604
2605        assert!(nbasis_range.contains(&res.optimal_nbasis));
2606        assert_eq!(res.scores.len(), nbasis_range.len());
2607        assert_eq!(res.nbasis_range, nbasis_range);
2608        assert_eq!(res.criterion, BasisCriterion::Aic);
2609        // optimal_nbasis should correspond to minimum score
2610        let min_score = res.scores.iter().copied().fold(f64::INFINITY, f64::min);
2611        let best_idx = res
2612            .scores
2613            .iter()
2614            .position(|&s| (s - min_score).abs() < 1e-15)
2615            .unwrap();
2616        assert_eq!(res.optimal_nbasis, nbasis_range[best_idx]);
2617    }
2618
2619    #[test]
2620    fn test_basis_nbasis_cv_result_clone() {
2621        let (data, t) = make_test_data(5, 50);
2622        let nbasis_range: Vec<usize> = vec![5, 10];
2623        let res = basis_nbasis_cv(
2624            &data,
2625            &t,
2626            &nbasis_range,
2627            &BasisType::Bspline { order: 4 },
2628            BasisCriterion::Gcv,
2629            5,
2630            1e-4,
2631        )
2632        .unwrap();
2633        let cloned = res.clone();
2634        assert_eq!(res, cloned);
2635    }
2636
2637    // ─── Non-uniform argvals ────────────────────────────────────────────────
2638
2639    #[test]
2640    fn test_smooth_basis_nonuniform_argvals() {
2641        let m = 50;
2642        // Non-uniform grid: denser at the ends
2643        let t: Vec<f64> = (0..m)
2644            .map(|i| {
2645                let x = i as f64 / (m - 1) as f64;
2646                0.5 * (1.0 - (PI * x).cos())
2647            })
2648            .collect();
2649        let mut data = FdMatrix::zeros(2, m);
2650        for i in 0..2 {
2651            for j in 0..m {
2652                data[(i, j)] = (2.0 * PI * t[j]).sin() + 0.1 * i as f64;
2653            }
2654        }
2655        let fdpar = make_bspline_fdpar(&t, 10, 1e-4);
2656        let res = smooth_basis(&data, &t, &fdpar);
2657        assert!(res.is_ok(), "Should handle non-uniform argvals");
2658        let r = res.unwrap();
2659        assert_eq!(r.fitted.shape(), (2, m));
2660    }
2661
2662    // ─── Numerical stability with extreme lambda ────────────────────────────
2663
2664    #[test]
2665    fn test_smooth_basis_very_small_lambda() {
2666        let m = 50;
2667        let (data, t) = make_test_data(2, m);
2668        let fdpar = make_bspline_fdpar(&t, 10, 1e-15);
2669        let res = smooth_basis(&data, &t, &fdpar);
2670        assert!(res.is_ok(), "Should handle very small lambda");
2671    }
2672
2673    #[test]
2674    fn test_smooth_basis_very_large_lambda() {
2675        let m = 50;
2676        let (data, t) = make_test_data(2, m);
2677        let fdpar = make_bspline_fdpar(&t, 10, 1e10);
2678        let res = smooth_basis(&data, &t, &fdpar);
2679        assert!(res.is_ok(), "Should handle very large lambda");
2680    }
2681
2682    // ─── Multiple curves consistency ────────────────────────────────────────
2683
2684    #[test]
2685    fn test_smooth_basis_multi_curve_vs_single_curve() {
2686        // Smoothing multiple curves at once should give the same result as smoothing each individually
2687        let m = 50;
2688        let n = 3;
2689        let (data, t) = make_test_data(n, m);
2690        let fdpar = make_bspline_fdpar(&t, 10, 1e-3);
2691
2692        // All at once
2693        let res_all = smooth_basis(&data, &t, &fdpar).unwrap();
2694
2695        // One at a time
2696        for i in 0..n {
2697            let mut single = FdMatrix::zeros(1, m);
2698            for j in 0..m {
2699                single[(0, j)] = data[(i, j)];
2700            }
2701            let res_single = smooth_basis(&single, &t, &fdpar).unwrap();
2702            for j in 0..m {
2703                assert!(
2704                    (res_all.fitted[(i, j)] - res_single.fitted[(0, j)]).abs() < 1e-10,
2705                    "Multi-curve fit should match single-curve fit: curve {} point {}",
2706                    i,
2707                    j
2708                );
2709            }
2710        }
2711    }
2712
2713    // ─── BasisCriterion comparison: all criteria produce finite scores ──────
2714
2715    #[test]
2716    fn test_basis_nbasis_cv_all_criteria_finite_scores() {
2717        let (data, t) = make_test_data(10, 60);
2718        let nbasis_range: Vec<usize> = vec![5, 7, 9, 11];
2719
2720        for criterion in [
2721            BasisCriterion::Gcv,
2722            BasisCriterion::Aic,
2723            BasisCriterion::Bic,
2724            BasisCriterion::Cv,
2725        ] {
2726            let res = basis_nbasis_cv(
2727                &data,
2728                &t,
2729                &nbasis_range,
2730                &BasisType::Bspline { order: 4 },
2731                criterion,
2732                5,
2733                1e-4,
2734            )
2735            .unwrap();
2736            // At least some scores should be finite (valid nbasis values)
2737            let finite_count = res.scores.iter().filter(|s| s.is_finite()).count();
2738            assert!(
2739                finite_count > 0,
2740                "At least one score should be finite for {:?}",
2741                criterion
2742            );
2743        }
2744    }
2745
2746    // ─── SmoothBasisGcvConfig tests ────────────────────────────────────────
2747
2748    #[test]
2749    fn test_smooth_basis_gcv_config_default() {
2750        let config = SmoothBasisGcvConfig::default();
2751        assert_eq!(config.basis_type, BasisType::Bspline { order: 4 });
2752        assert_eq!(config.nbasis, 15);
2753        assert_eq!(config.lfd_order, 2);
2754        assert_eq!(config.log_lambda_range, (-10.0, 2.0));
2755        assert_eq!(config.n_grid, 50);
2756    }
2757
2758    #[test]
2759    fn test_smooth_basis_gcv_config_clone_eq() {
2760        let config = SmoothBasisGcvConfig {
2761            nbasis: 20,
2762            ..SmoothBasisGcvConfig::default()
2763        };
2764        let cloned = config.clone();
2765        assert_eq!(config, cloned);
2766    }
2767
2768    #[test]
2769    fn test_smooth_basis_gcv_config_debug() {
2770        let config = SmoothBasisGcvConfig::default();
2771        let debug_str = format!("{:?}", config);
2772        assert!(debug_str.contains("SmoothBasisGcvConfig"));
2773        assert!(debug_str.contains("nbasis"));
2774    }
2775
2776    #[test]
2777    fn test_smooth_basis_gcv_config_partial_override() {
2778        let config = SmoothBasisGcvConfig {
2779            basis_type: BasisType::Fourier { period: 2.0 },
2780            n_grid: 100,
2781            ..SmoothBasisGcvConfig::default()
2782        };
2783        assert_eq!(config.basis_type, BasisType::Fourier { period: 2.0 });
2784        assert_eq!(config.n_grid, 100);
2785        // defaults preserved
2786        assert_eq!(config.nbasis, 15);
2787        assert_eq!(config.lfd_order, 2);
2788    }
2789
2790    #[test]
2791    fn test_smooth_basis_gcv_with_config_default() {
2792        let (data, t) = make_test_data(5, 101);
2793        let config = SmoothBasisGcvConfig::default();
2794        let result = smooth_basis_gcv_with_config(&data, &t, &config);
2795        assert!(result.is_ok(), "GCV with default config should succeed");
2796        let res = result.unwrap();
2797        assert_eq!(res.fitted.shape(), (5, 101));
2798        assert!(res.edf > 0.0);
2799        assert!(res.gcv.is_finite());
2800    }
2801
2802    #[test]
2803    fn test_smooth_basis_gcv_with_config_custom() {
2804        let (data, t) = make_test_data(3, 50);
2805        let config = SmoothBasisGcvConfig {
2806            nbasis: 10,
2807            log_lambda_range: (-6.0, 0.0),
2808            n_grid: 15,
2809            ..SmoothBasisGcvConfig::default()
2810        };
2811        let result = smooth_basis_gcv_with_config(&data, &t, &config);
2812        assert!(result.is_ok());
2813    }
2814
2815    #[test]
2816    fn test_smooth_basis_gcv_with_config_matches_direct() {
2817        let (data, t) = make_test_data(3, 50);
2818        let config = SmoothBasisGcvConfig {
2819            nbasis: 10,
2820            log_lambda_range: (-6.0, 0.0),
2821            n_grid: 20,
2822            ..SmoothBasisGcvConfig::default()
2823        };
2824        let with_config = smooth_basis_gcv_with_config(&data, &t, &config).unwrap();
2825        let direct = smooth_basis_gcv(
2826            &data,
2827            &t,
2828            &config.basis_type,
2829            config.nbasis,
2830            config.lfd_order,
2831            config.log_lambda_range,
2832            config.n_grid,
2833        )
2834        .unwrap();
2835        assert_eq!(with_config.gcv, direct.gcv);
2836        assert_eq!(with_config.edf, direct.edf);
2837        assert_eq!(with_config.nbasis, direct.nbasis);
2838    }
2839
2840    #[test]
2841    fn test_smooth_basis_gcv_with_config_fourier() {
2842        let m = 100;
2843        let t = uniform_grid(m);
2844        let mut data = FdMatrix::zeros(2, m);
2845        for i in 0..2 {
2846            for j in 0..m {
2847                data[(i, j)] = (2.0 * PI * t[j]).sin() + (4.0 * PI * t[j]).cos();
2848            }
2849        }
2850        let config = SmoothBasisGcvConfig {
2851            basis_type: BasisType::Fourier { period: 1.0 },
2852            nbasis: 7,
2853            n_grid: 20,
2854            ..SmoothBasisGcvConfig::default()
2855        };
2856        let result = smooth_basis_gcv_with_config(&data, &t, &config);
2857        assert!(result.is_ok());
2858    }
2859
2860    // ─── BasisNbasisCvConfig tests ─────────────────────────────────────────
2861
2862    #[test]
2863    fn test_basis_nbasis_cv_config_default() {
2864        let config = BasisNbasisCvConfig::default();
2865        assert_eq!(config.basis_type, BasisType::Bspline { order: 4 });
2866        assert_eq!(config.nbasis_range, (5, 30));
2867        assert!((config.lambda - 1e-4).abs() < 1e-15);
2868        assert_eq!(config.lfd_order, 2);
2869        assert_eq!(config.n_folds, 5);
2870        assert_eq!(config.criterion, BasisCriterion::Gcv);
2871    }
2872
2873    #[test]
2874    fn test_basis_nbasis_cv_config_clone_eq() {
2875        let config = BasisNbasisCvConfig {
2876            nbasis_range: (4, 15),
2877            ..BasisNbasisCvConfig::default()
2878        };
2879        let cloned = config.clone();
2880        assert_eq!(config, cloned);
2881    }
2882
2883    #[test]
2884    fn test_basis_nbasis_cv_config_debug() {
2885        let config = BasisNbasisCvConfig::default();
2886        let debug_str = format!("{:?}", config);
2887        assert!(debug_str.contains("BasisNbasisCvConfig"));
2888        assert!(debug_str.contains("nbasis_range"));
2889    }
2890
2891    #[test]
2892    fn test_basis_nbasis_cv_config_partial_override() {
2893        let config = BasisNbasisCvConfig {
2894            criterion: BasisCriterion::Aic,
2895            lambda: 1e-2,
2896            ..BasisNbasisCvConfig::default()
2897        };
2898        assert_eq!(config.criterion, BasisCriterion::Aic);
2899        assert!((config.lambda - 1e-2).abs() < 1e-15);
2900        // defaults preserved
2901        assert_eq!(config.nbasis_range, (5, 30));
2902        assert_eq!(config.n_folds, 5);
2903    }
2904
2905    #[test]
2906    fn test_basis_nbasis_cv_with_config_default() {
2907        let (data, t) = make_test_data(5, 51);
2908        let config = BasisNbasisCvConfig {
2909            nbasis_range: (5, 12),
2910            ..BasisNbasisCvConfig::default()
2911        };
2912        let result = basis_nbasis_cv_with_config(&data, &t, &config);
2913        assert!(
2914            result.is_ok(),
2915            "nbasis CV with default config should succeed"
2916        );
2917        let res = result.unwrap();
2918        assert!(res.optimal_nbasis >= 5 && res.optimal_nbasis <= 12);
2919        assert_eq!(res.scores.len(), 8); // 5..=12 = 8 values
2920        assert_eq!(res.criterion, BasisCriterion::Gcv);
2921    }
2922
2923    #[test]
2924    fn test_basis_nbasis_cv_with_config_aic() {
2925        let (data, t) = make_test_data(5, 51);
2926        let config = BasisNbasisCvConfig {
2927            nbasis_range: (5, 10),
2928            criterion: BasisCriterion::Aic,
2929            ..BasisNbasisCvConfig::default()
2930        };
2931        let result = basis_nbasis_cv_with_config(&data, &t, &config);
2932        assert!(result.is_ok());
2933        assert_eq!(result.unwrap().criterion, BasisCriterion::Aic);
2934    }
2935
2936    #[test]
2937    fn test_basis_nbasis_cv_with_config_cv_folds() {
2938        let (data, t) = make_test_data(10, 51);
2939        let config = BasisNbasisCvConfig {
2940            nbasis_range: (5, 9),
2941            criterion: BasisCriterion::Cv,
2942            n_folds: 3,
2943            ..BasisNbasisCvConfig::default()
2944        };
2945        let result = basis_nbasis_cv_with_config(&data, &t, &config);
2946        assert!(result.is_ok());
2947        assert_eq!(result.unwrap().criterion, BasisCriterion::Cv);
2948    }
2949
2950    #[test]
2951    fn test_basis_nbasis_cv_with_config_matches_direct() {
2952        let (data, t) = make_test_data(5, 51);
2953        let config = BasisNbasisCvConfig {
2954            nbasis_range: (5, 10),
2955            criterion: BasisCriterion::Bic,
2956            lambda: 1e-3,
2957            ..BasisNbasisCvConfig::default()
2958        };
2959        let with_config = basis_nbasis_cv_with_config(&data, &t, &config).unwrap();
2960        let nbasis_range: Vec<usize> = (5..=10).collect();
2961        let direct = basis_nbasis_cv(
2962            &data,
2963            &t,
2964            &nbasis_range,
2965            &config.basis_type,
2966            config.criterion,
2967            config.n_folds,
2968            config.lambda,
2969        )
2970        .unwrap();
2971        assert_eq!(with_config.optimal_nbasis, direct.optimal_nbasis);
2972        assert_eq!(with_config.scores, direct.scores);
2973        assert_eq!(with_config.nbasis_range, direct.nbasis_range);
2974    }
2975
2976    #[test]
2977    fn test_basis_nbasis_cv_with_config_nbasis_range_expansion() {
2978        let (data, t) = make_test_data(5, 51);
2979        let config = BasisNbasisCvConfig {
2980            nbasis_range: (7, 7), // single value
2981            ..BasisNbasisCvConfig::default()
2982        };
2983        let result = basis_nbasis_cv_with_config(&data, &t, &config);
2984        assert!(result.is_ok());
2985        let res = result.unwrap();
2986        assert_eq!(res.optimal_nbasis, 7);
2987        assert_eq!(res.scores.len(), 1);
2988    }
2989}