Skip to main content

model_selection_rs/splitters/
leave_one_out.rs

1//! Leave-one-out cross-validation.
2
3use super::CvSplitter;
4use crate::error::{ModelSelectionError, Result};
5
6/// Leave-one-out cross-validation (LOO).
7///
8/// Each sample serves as the sole test point exactly once, with every other
9/// sample forming the training set — equivalent to [`KFold`](super::KFold) with
10/// `n_splits == n_samples`. Included for API completeness.
11///
12/// **Cost warning:** LOO fits the model `n_samples` times. On anything but small
13/// datasets this is very expensive; prefer a modest `KFold` (5 or 10) unless you
14/// specifically need LOO's near-unbiased (but high-variance) estimate.
15///
16/// ```
17/// use model_selection_rs::splitters::{CvSplitter, LeaveOneOut};
18///
19/// let loo = LeaveOneOut;
20/// let splits = loo.split(4).unwrap();
21/// assert_eq!(splits.len(), 4);
22/// assert_eq!(splits[0], (vec![1, 2, 3], vec![0]));
23/// ```
24#[derive(Debug, Clone, Copy, Default)]
25pub struct LeaveOneOut;
26
27impl CvSplitter for LeaveOneOut {
28    fn split(&self, n_samples: usize) -> Result<Vec<(Vec<usize>, Vec<usize>)>> {
29        if n_samples < 2 {
30            return Err(ModelSelectionError::NotEnoughSamples {
31                needed: 2,
32                got: n_samples,
33            });
34        }
35        let splits = (0..n_samples)
36            .map(|test_idx| {
37                let train: Vec<usize> = (0..n_samples).filter(|&i| i != test_idx).collect();
38                (train, vec![test_idx])
39            })
40            .collect();
41        Ok(splits)
42    }
43
44    fn n_splits(&self) -> usize {
45        // Unknown without n_samples; LOO's split count equals n_samples, which
46        // is only known at `split` time. Report 0 as "depends on data".
47        0
48    }
49}
50
51#[cfg(test)]
52mod tests {
53    use super::*;
54
55    #[test]
56    fn each_sample_is_test_once() {
57        let splits = LeaveOneOut.split(5).unwrap();
58        assert_eq!(splits.len(), 5);
59        for (i, (train, test)) in splits.iter().enumerate() {
60            assert_eq!(test, &vec![i]);
61            assert_eq!(train.len(), 4);
62            assert!(!train.contains(&i));
63        }
64    }
65
66    #[test]
67    fn errors_on_tiny_input() {
68        assert!(LeaveOneOut.split(1).is_err());
69    }
70}