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}