1use 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#[non_exhaustive]
15#[derive(Debug, Clone, PartialEq)]
16pub struct LambdaCvConfig {
17 pub lambdas: Vec<f64>,
19 pub n_folds: usize,
21 pub max_iter: usize,
23 pub tol: f64,
25 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#[derive(Debug, Clone, PartialEq)]
43#[non_exhaustive]
44pub struct LambdaCvResult {
45 pub best_lambda: f64,
47 pub cv_scores: Vec<f64>,
49 pub lambdas: Vec<f64>,
51}
52
53#[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 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 let k_max = *folds.iter().max().unwrap_or(&0) + 1;
119
120 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 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}