Skip to main content

model_selection_rs/evaluate/
cross_validate.rs

1//! The core cross-validation evaluation loop.
2
3use std::time::{Duration, Instant};
4
5use ndarray::{Array1, Array2, Axis};
6
7use crate::error::Result;
8use crate::scoring::Scorer;
9use crate::splitters::CvSplitter;
10
11/// A scorer usable by the evaluation utilities.
12///
13/// The `Send + Sync` bounds keep the public signatures identical whether or not
14/// the `parallel` feature is enabled (feature flags never change the API), and
15/// let folds be scored across threads when it is. Every built-in scorer, and any
16/// [`make_scorer`](crate::scoring::make_scorer) closure over `Send + Sync` data,
17/// satisfies them.
18pub type BoxedScorer = Box<dyn Scorer + Send + Sync>;
19
20/// Results of a [`cross_validate`] run.
21///
22/// Scores are stored as `[scorer][fold]`. Use [`mean_test_score`] /
23/// [`std_test_score`] (by index or by name) for summaries.
24///
25/// [`mean_test_score`]: CvResults::mean_test_score
26/// [`std_test_score`]: CvResults::std_test_score
27#[derive(Debug, Clone)]
28pub struct CvResults {
29    /// Metric names, in the order they were supplied.
30    pub scorer_names: Vec<String>,
31    /// Test scores as `[scorer][fold]`.
32    pub test_scores: Vec<Vec<f64>>,
33    /// Train scores as `[scorer][fold]`, if `return_train_scores` was set.
34    pub train_scores: Option<Vec<Vec<f64>>>,
35    /// Wall-clock fit time per fold.
36    pub fit_times: Vec<Duration>,
37    /// Wall-clock score time per fold.
38    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    /// Number of folds evaluated.
59    #[must_use]
60    pub fn n_splits(&self) -> usize {
61        self.fit_times.len()
62    }
63
64    /// Mean test score for scorer `idx`.
65    #[must_use]
66    pub fn mean_test_score(&self, idx: usize) -> f64 {
67        mean(&self.test_scores[idx])
68    }
69
70    /// Population standard deviation of the test scores for scorer `idx`.
71    #[must_use]
72    pub fn std_test_score(&self, idx: usize) -> f64 {
73        std(&self.test_scores[idx])
74    }
75
76    /// Mean test score for the scorer named `name`, if present.
77    #[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    /// Total fit time across all folds.
86    #[must_use]
87    pub fn total_fit_time(&self) -> Duration {
88        self.fit_times.iter().sum()
89    }
90}
91
92/// Compute every scorer's score for a single fold. Shared by the serial and
93/// parallel paths so both produce numerically identical results.
94struct 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
151/// Cross-validate a model-fitting closure with one or more scorers.
152///
153/// Ties a [`CvSplitter`], a model-fitting closure, and one or more
154/// [`Scorer`]s together — the utility you actually call day to day. `fit_fn` is
155/// generic (it takes `(x_train, y_train)` and returns a *prediction* closure),
156/// so this works with `smartcore`, `linfa`, or a hand-rolled model without this
157/// crate depending on any of them.
158///
159/// Every scorer is evaluated in the **same pass** over the folds (re-fitting per
160/// metric would be wasteful), so asking for accuracy *and* F1 together costs one
161/// set of fits, not two.
162///
163/// With the `parallel` feature enabled the folds are fit and scored across a
164/// `rayon` thread pool. The results are numerically identical to the serial path
165/// — parallelism changes only wall-clock time, and the per-fold order is
166/// preserved.
167///
168/// # Errors
169///
170/// Propagates any error from `splitter.split(x.nrows())`.
171///
172/// # Example
173///
174/// ```
175/// use ndarray::{array, Array1, Array2};
176/// use model_selection_rs::evaluate::{cross_validate, BoxedScorer};
177/// use model_selection_rs::scoring::MeanSquaredError;
178/// use model_selection_rs::splitters::KFold;
179///
180/// // A trivial "model" that predicts the training mean.
181/// let x: Array2<f64> = Array2::zeros((10, 1));
182/// let y: Array1<f64> = array![1., 2., 3., 4., 5., 6., 7., 8., 9., 10.];
183/// let kf = KFold::new(5).unwrap();
184/// let scorers: Vec<BoxedScorer> = vec![Box::new(MeanSquaredError)];
185/// let res = cross_validate(&kf, &x, &y, |_xt, yt| {
186///     let mean = yt.sum() / yt.len() as f64;
187///     move |xq: &Array2<f64>| Array1::from_elem(xq.nrows(), mean)
188/// }, &scorers, false).unwrap();
189/// assert_eq!(res.n_splits(), 5);
190/// ```
191pub 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
228/// Transpose per-fold results into the `[scorer][fold]` layout of `CvResults`.
229fn 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    /// Fit an ordinary least squares line y = a*x + b on one feature.
274    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        // Perfect linear data -> near-zero MSE, R2 ~ 1.
301        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}