model_selection_rs/evaluate/
cross_validate.rs1use std::time::{Duration, Instant};
4
5use ndarray::{Array1, Array2, Axis};
6
7use crate::error::Result;
8use crate::scoring::Scorer;
9use crate::splitters::CvSplitter;
10
11pub type BoxedScorer = Box<dyn Scorer + Send + Sync>;
19
20#[derive(Debug, Clone)]
28pub struct CvResults {
29 pub scorer_names: Vec<String>,
31 pub test_scores: Vec<Vec<f64>>,
33 pub train_scores: Option<Vec<Vec<f64>>>,
35 pub fit_times: Vec<Duration>,
37 pub score_times: Vec<Duration>,
39}
40
41fn mean(xs: &[f64]) -> f64 {
42 if xs.is_empty() {
43 return f64::NAN;
44 }
45 xs.iter().sum::<f64>() / xs.len() as f64
46}
47
48fn std(xs: &[f64]) -> f64 {
49 if xs.len() < 2 {
50 return 0.0;
51 }
52 let m = mean(xs);
53 let var = xs.iter().map(|x| (x - m).powi(2)).sum::<f64>() / xs.len() as f64;
54 var.sqrt()
55}
56
57impl CvResults {
58 #[must_use]
60 pub fn n_splits(&self) -> usize {
61 self.fit_times.len()
62 }
63
64 #[must_use]
66 pub fn mean_test_score(&self, idx: usize) -> f64 {
67 mean(&self.test_scores[idx])
68 }
69
70 #[must_use]
72 pub fn std_test_score(&self, idx: usize) -> f64 {
73 std(&self.test_scores[idx])
74 }
75
76 #[must_use]
78 pub fn mean_test_score_by_name(&self, name: &str) -> Option<f64> {
79 self.scorer_names
80 .iter()
81 .position(|n| n == name)
82 .map(|i| self.mean_test_score(i))
83 }
84
85 #[must_use]
87 pub fn total_fit_time(&self) -> Duration {
88 self.fit_times.iter().sum()
89 }
90}
91
92struct FoldScores {
95 test: Vec<f64>,
96 train: Option<Vec<f64>>,
97 fit_time: Duration,
98 score_time: Duration,
99}
100
101fn evaluate_fold<F, M>(
102 x: &Array2<f64>,
103 y: &Array1<f64>,
104 train: &[usize],
105 test: &[usize],
106 fit_fn: &F,
107 scorers: &[BoxedScorer],
108 return_train_scores: bool,
109) -> FoldScores
110where
111 F: Fn(&Array2<f64>, &Array1<f64>) -> M,
112 M: Fn(&Array2<f64>) -> Array1<f64>,
113{
114 let x_train = x.select(Axis(0), train);
115 let y_train = y.select(Axis(0), train);
116 let x_test = x.select(Axis(0), test);
117 let y_test = y.select(Axis(0), test);
118
119 let fit_start = Instant::now();
120 let model = fit_fn(&x_train, &y_train);
121 let fit_time = fit_start.elapsed();
122
123 let score_start = Instant::now();
124 let pred_test = model(&x_test);
125 let test: Vec<f64> = scorers
126 .iter()
127 .map(|s| s.score(&y_test, &pred_test))
128 .collect();
129 let score_time = score_start.elapsed();
130
131 let train = if return_train_scores {
132 let pred_train = model(&x_train);
133 Some(
134 scorers
135 .iter()
136 .map(|s| s.score(&y_train, &pred_train))
137 .collect(),
138 )
139 } else {
140 None
141 };
142
143 FoldScores {
144 test,
145 train,
146 fit_time,
147 score_time,
148 }
149}
150
151pub fn cross_validate<S, F, M>(
192 splitter: &S,
193 x: &Array2<f64>,
194 y: &Array1<f64>,
195 fit_fn: F,
196 scorers: &[BoxedScorer],
197 return_train_scores: bool,
198) -> Result<CvResults>
199where
200 S: CvSplitter,
201 F: Fn(&Array2<f64>, &Array1<f64>) -> M + Sync,
202 M: Fn(&Array2<f64>) -> Array1<f64>,
203{
204 let splits = splitter.split(x.nrows())?;
205
206 #[cfg(feature = "parallel")]
207 let fold_scores: Vec<FoldScores> = {
208 use rayon::prelude::*;
209 splits
210 .par_iter()
211 .map(|(train, test)| {
212 evaluate_fold(x, y, train, test, &fit_fn, scorers, return_train_scores)
213 })
214 .collect()
215 };
216
217 #[cfg(not(feature = "parallel"))]
218 let fold_scores: Vec<FoldScores> = splits
219 .iter()
220 .map(|(train, test)| {
221 evaluate_fold(x, y, train, test, &fit_fn, scorers, return_train_scores)
222 })
223 .collect();
224
225 Ok(assemble(fold_scores, scorers, return_train_scores))
226}
227
228fn assemble(
230 fold_scores: Vec<FoldScores>,
231 scorers: &[BoxedScorer],
232 return_train_scores: bool,
233) -> CvResults {
234 let n_scorers = scorers.len();
235 let mut test_scores = vec![Vec::with_capacity(fold_scores.len()); n_scorers];
236 let mut train_scores = if return_train_scores {
237 Some(vec![Vec::with_capacity(fold_scores.len()); n_scorers])
238 } else {
239 None
240 };
241 let mut fit_times = Vec::with_capacity(fold_scores.len());
242 let mut score_times = Vec::with_capacity(fold_scores.len());
243
244 for fold in fold_scores {
245 for (s, v) in fold.test.into_iter().enumerate() {
246 test_scores[s].push(v);
247 }
248 if let (Some(dst), Some(src)) = (train_scores.as_mut(), fold.train) {
249 for (s, v) in src.into_iter().enumerate() {
250 dst[s].push(v);
251 }
252 }
253 fit_times.push(fold.fit_time);
254 score_times.push(fold.score_time);
255 }
256
257 CvResults {
258 scorer_names: scorers.iter().map(|s| s.name().to_string()).collect(),
259 test_scores,
260 train_scores,
261 fit_times,
262 score_times,
263 }
264}
265
266#[cfg(test)]
267mod tests {
268 use super::*;
269 use crate::scoring::{MeanSquaredError, R2Score};
270 use crate::splitters::KFold;
271 use ndarray::array;
272
273 fn ols_fit(x: &Array2<f64>, y: &Array1<f64>) -> impl Fn(&Array2<f64>) -> Array1<f64> {
275 let n = x.nrows() as f64;
276 let xs: Vec<f64> = x.column(0).to_vec();
277 let ys: Vec<f64> = y.to_vec();
278 let mean_x = xs.iter().sum::<f64>() / n;
279 let mean_y = ys.iter().sum::<f64>() / n;
280 let cov: f64 = xs
281 .iter()
282 .zip(&ys)
283 .map(|(a, b)| (a - mean_x) * (b - mean_y))
284 .sum();
285 let var: f64 = xs.iter().map(|a| (a - mean_x).powi(2)).sum();
286 let slope = if var == 0.0 { 0.0 } else { cov / var };
287 let intercept = mean_y - slope * mean_x;
288 move |xq: &Array2<f64>| xq.column(0).mapv(|v| slope * v + intercept)
289 }
290
291 #[test]
292 fn multiple_scorers_in_one_pass() {
293 let x = Array2::from_shape_fn((20, 1), |(i, _)| i as f64);
294 let y = x.column(0).mapv(|v| 2.0 * v + 1.0);
295 let kf = KFold::new(4).unwrap();
296 let scorers: Vec<BoxedScorer> = vec![Box::new(MeanSquaredError), Box::new(R2Score)];
297 let res = cross_validate(&kf, &x, &y, ols_fit, &scorers, true).unwrap();
298 assert_eq!(res.scorer_names, vec!["mse", "r2"]);
299 assert_eq!(res.n_splits(), 4);
300 assert!(res.mean_test_score(0) < 1e-9);
302 assert!((res.mean_test_score(1) - 1.0).abs() < 1e-9);
303 assert!(res.train_scores.is_some());
304 }
305
306 #[test]
307 fn lookup_by_name() {
308 let x = Array2::from_shape_fn((12, 1), |(i, _)| i as f64);
309 let y = x.column(0).to_owned();
310 let kf = KFold::new(3).unwrap();
311 let scorers: Vec<BoxedScorer> = vec![Box::new(MeanSquaredError)];
312 let res = cross_validate(&kf, &x, &y, ols_fit, &scorers, false).unwrap();
313 assert!(res.mean_test_score_by_name("mse").is_some());
314 assert!(res.mean_test_score_by_name("nope").is_none());
315 }
316
317 #[test]
318 fn array_example_in_docs() {
319 let x = array![[0.0], [1.0], [2.0], [3.0]];
320 let y = array![0.0, 1.0, 2.0, 3.0];
321 let kf = KFold::new(2).unwrap();
322 let scorers: Vec<BoxedScorer> = vec![Box::new(MeanSquaredError)];
323 let res = cross_validate(&kf, &x, &y, ols_fit, &scorers, false).unwrap();
324 assert_eq!(res.n_splits(), 2);
325 }
326}