Skip to main content

fdars_core/alignment/
lambda_cv.rs

1//! Cross-validation for the elastic alignment regularisation parameter lambda.
2
3use super::karcher::karcher_mean;
4use super::pairwise::elastic_distance;
5use crate::cv::{create_folds, fold_indices, subset_rows};
6use crate::error::FdarError;
7use crate::matrix::FdMatrix;
8
9// ─── Config / Result ─────────────────────────────────────────────────────────
10
11/// Configuration for lambda cross-validation.
12///
13/// Construct via `LambdaCvConfig::default()`, then assign the fields you need (e.g. `let mut c = LambdaCvConfig::default(); c.field = …;`). This struct is `#[non_exhaustive]`, so external crates cannot build it with a struct literal — not even functional-update `..Default::default()` form.
14#[non_exhaustive]
15#[derive(Debug, Clone, PartialEq)]
16pub struct LambdaCvConfig {
17    /// Candidate lambda values to evaluate.
18    pub lambdas: Vec<f64>,
19    /// Number of folds (0 = leave-one-out).
20    pub n_folds: usize,
21    /// Maximum Karcher iterations per fold.
22    pub max_iter: usize,
23    /// Karcher convergence tolerance.
24    pub tol: f64,
25    /// RNG seed for fold assignment.
26    pub seed: u64,
27}
28
29impl Default for LambdaCvConfig {
30    fn default() -> Self {
31        Self {
32            lambdas: vec![0.0, 0.01, 0.1, 1.0, 10.0],
33            n_folds: 5,
34            max_iter: 15,
35            tol: 1e-3,
36            seed: 42,
37        }
38    }
39}
40
41/// Result of lambda cross-validation.
42#[derive(Debug, Clone, PartialEq)]
43#[non_exhaustive]
44pub struct LambdaCvResult {
45    /// Lambda with the lowest mean CV score.
46    pub best_lambda: f64,
47    /// Mean CV score for each candidate lambda (same order as `lambdas`).
48    pub cv_scores: Vec<f64>,
49    /// Candidate lambda values (copied from config).
50    pub lambdas: Vec<f64>,
51}
52
53// ─── Cross-validation ────────────────────────────────────────────────────────
54
55/// Select the best elastic-alignment regularisation parameter via K-fold
56/// cross-validation.
57///
58/// For each candidate lambda the data are split into K folds. A Karcher mean
59/// is computed on the training set and every held-out curve is scored by its
60/// elastic distance to that mean. The lambda with the lowest average
61/// held-out distance wins.
62///
63/// # Arguments
64/// * `data`    — Functional data matrix (n x m).
65/// * `argvals` — Evaluation grid (length m).
66/// * `config`  — Cross-validation settings (lambdas, folds, iterations, …).
67///
68/// # Errors
69/// Returns `FdarError::InvalidDimension` if `data` has fewer than 4 rows
70/// or `argvals` length does not match `data.ncols()`.
71/// Returns `FdarError::InvalidParameter` if any lambda is negative or
72/// `n_folds` is 1.
73#[must_use = "expensive computation whose result should not be discarded"]
74pub fn lambda_cv(
75    data: &FdMatrix,
76    argvals: &[f64],
77    config: &LambdaCvConfig,
78) -> Result<LambdaCvResult, FdarError> {
79    let n = data.nrows();
80    let m = data.ncols();
81
82    // ── Validation ──────────────────────────────────────────────────────
83    if n < 4 {
84        return Err(FdarError::InvalidDimension {
85            parameter: "data",
86            expected: "at least 4 rows".to_string(),
87            actual: format!("{n} rows"),
88        });
89    }
90    if argvals.len() != m {
91        return Err(FdarError::InvalidDimension {
92            parameter: "argvals",
93            expected: format!("{m}"),
94            actual: format!("{}", argvals.len()),
95        });
96    }
97    if config.lambdas.iter().any(|&l| l < 0.0) {
98        return Err(FdarError::InvalidParameter {
99            parameter: "lambdas",
100            message: "all lambda values must be >= 0".to_string(),
101        });
102    }
103    if config.n_folds == 1 {
104        return Err(FdarError::InvalidParameter {
105            parameter: "n_folds",
106            message: "n_folds must be > 1 or 0 (leave-one-out)".to_string(),
107        });
108    }
109
110    let actual_folds = if config.n_folds == 0 {
111        n
112    } else {
113        config.n_folds
114    };
115    let folds = create_folds(n, actual_folds, config.seed);
116
117    // Number of distinct fold labels actually produced.
118    let k_max = *folds.iter().max().unwrap_or(&0) + 1;
119
120    // ── Evaluate each lambda ────────────────────────────────────────────
121    let mut cv_scores = Vec::with_capacity(config.lambdas.len());
122
123    for &lambda in &config.lambdas {
124        let mut fold_scores = Vec::with_capacity(k_max);
125
126        for k in 0..k_max {
127            let (train_idx, test_idx) = fold_indices(&folds, k);
128            if train_idx.is_empty() || test_idx.is_empty() {
129                continue;
130            }
131
132            let train_data = subset_rows(data, &train_idx);
133            let km = karcher_mean(&train_data, argvals, config.max_iter, config.tol, lambda);
134
135            let fold_dist: f64 = test_idx
136                .iter()
137                .map(|&idx| {
138                    let test_curve = data.row(idx);
139                    elastic_distance(&test_curve, &km.mean, argvals, lambda)
140                })
141                .sum::<f64>()
142                / test_idx.len() as f64;
143
144            fold_scores.push(fold_dist);
145        }
146
147        let mean_score = if fold_scores.is_empty() {
148            f64::INFINITY
149        } else {
150            fold_scores.iter().sum::<f64>() / fold_scores.len() as f64
151        };
152        cv_scores.push(mean_score);
153    }
154
155    // ── Pick best lambda ────────────────────────────────────────────────
156    let best_idx = cv_scores
157        .iter()
158        .enumerate()
159        .min_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
160        .map(|(i, _)| i)
161        .unwrap_or(0);
162
163    Ok(LambdaCvResult {
164        best_lambda: config.lambdas[best_idx],
165        cv_scores,
166        lambdas: config.lambdas.clone(),
167    })
168}
169
170#[cfg(test)]
171mod tests {
172    use super::*;
173    use crate::simulation::{sim_fundata, EFunType, EValType};
174    use crate::test_helpers::uniform_grid;
175
176    fn make_test_data(n: usize, m: usize) -> (FdMatrix, Vec<f64>) {
177        let t = uniform_grid(m);
178        let data = sim_fundata(n, &t, 3, EFunType::Fourier, EValType::Exponential, Some(42));
179        (data, t)
180    }
181
182    #[test]
183    fn lambda_cv_default_config() {
184        let (data, t) = make_test_data(8, 30);
185        let config = LambdaCvConfig {
186            max_iter: 5,
187            tol: 1e-2,
188            ..LambdaCvConfig::default()
189        };
190        let result = lambda_cv(&data, &t, &config).unwrap();
191        assert_eq!(result.cv_scores.len(), config.lambdas.len());
192        assert!(result.best_lambda >= 0.0);
193        assert!(result.cv_scores.iter().all(|&s| s.is_finite()));
194    }
195
196    #[test]
197    fn lambda_cv_loo() {
198        let (data, t) = make_test_data(6, 25);
199        let config = LambdaCvConfig {
200            lambdas: vec![0.0, 1.0],
201            n_folds: 0,
202            max_iter: 3,
203            tol: 1e-2,
204            seed: 7,
205        };
206        let result = lambda_cv(&data, &t, &config).unwrap();
207        assert_eq!(result.cv_scores.len(), 2);
208    }
209
210    #[test]
211    fn lambda_cv_rejects_too_few_rows() {
212        let t = uniform_grid(10);
213        let data = sim_fundata(3, &t, 2, EFunType::Fourier, EValType::Exponential, Some(0));
214        let config = LambdaCvConfig::default();
215        assert!(lambda_cv(&data, &t, &config).is_err());
216    }
217
218    #[test]
219    fn lambda_cv_rejects_negative_lambda() {
220        let (data, t) = make_test_data(8, 20);
221        let config = LambdaCvConfig {
222            lambdas: vec![-1.0, 0.0],
223            ..LambdaCvConfig::default()
224        };
225        assert!(lambda_cv(&data, &t, &config).is_err());
226    }
227
228    #[test]
229    fn lambda_cv_rejects_one_fold() {
230        let (data, t) = make_test_data(8, 20);
231        let config = LambdaCvConfig {
232            n_folds: 1,
233            ..LambdaCvConfig::default()
234        };
235        assert!(lambda_cv(&data, &t, &config).is_err());
236    }
237
238    #[test]
239    fn lambda_cv_rejects_argval_mismatch() {
240        let (data, _) = make_test_data(8, 20);
241        let bad_t = uniform_grid(15);
242        let config = LambdaCvConfig::default();
243        assert!(lambda_cv(&data, &bad_t, &config).is_err());
244    }
245}