use super::cross_validation::CvScores;
use crate::error::{Error, Result};
use crate::resampling::LeaveOneOutCrossValidation;
pub fn loo_indices(n: usize) -> Result<Vec<(Vec<usize>, Vec<usize>)>> {
if n < 2 {
return Err(Error::InsufficientData);
}
let folds = (0..n)
.map(|i| {
let train: Vec<usize> = (0..n).filter(|&j| j != i).collect();
(train, vec![i])
})
.collect();
Ok(folds)
}
pub fn loo_cross_validate(
n: usize,
mut fit_score: impl FnMut(&[usize], &[usize]) -> f64,
) -> Result<CvScores> {
let folds = loo_indices(n)?;
let fold_scores: Vec<f64> = folds
.iter()
.map(|(train, test)| fit_score(train, test))
.collect();
Ok(CvScores::new(fold_scores))
}
impl LeaveOneOutCrossValidation {
pub fn run(
&self,
n: usize,
fit_score: impl FnMut(&[usize], &[usize]) -> f64,
) -> Result<CvScores> {
loo_cross_validate(n, fit_score)
}
}
#[cfg(kani)]
mod verification {
use super::{Error, loo_indices};
#[kani::proof]
#[kani::unwind(2)]
fn resampling_loo_rejects_small_n() {
let n: usize = kani::any();
kani::assume(n < 2);
let result = loo_indices(n);
assert!(
matches!(result, Err(Error::InsufficientData)),
"n < 2 must be rejected with InsufficientData"
);
}
#[kani::proof]
#[kani::unwind(5)]
fn resampling_loo_indices_partition() {
const N: usize = 3;
let result = loo_indices(N);
assert!(result.is_ok(), "n >= 2 must produce LOO folds");
if let Ok(folds) = result {
assert!(folds.len() == N, "expected one fold per observation");
let mut seen = [0u8; N];
for (i, (train, test)) in folds.iter().enumerate() {
assert!(test.len() == 1, "each test set must be a singleton");
for &t in test {
assert!(t == i, "fold {i} must test its own index");
assert!(t < N, "test index {t} escaped 0..N");
seen[t] += 1;
}
assert!(train.len() == N - 1, "train must be the complement");
for &tr in train {
assert!(tr < N, "train index {tr} escaped 0..N");
assert!(tr != i, "train must not contain the held-out index");
}
}
assert!(
seen.iter().all(|&c| c == 1),
"test singletons must partition 0..N"
);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::Error;
use crate::resampling::LeaveOneOutCrossValidation;
#[test]
fn loo_indices_rejects_fewer_than_two() {
assert_eq!(
loo_indices(1),
Err(Error::InsufficientData),
"n < 2 must be rejected: a singleton has no held-out complement"
);
}
#[test]
fn loo_indices_three_gives_each_singleton_test() -> Result<()> {
let folds = loo_indices(3)?;
assert_eq!(
folds,
vec![
(vec![1, 2], vec![0]),
(vec![0, 2], vec![1]),
(vec![0, 1], vec![2]),
],
"each fold i must test [i] and train on the ordered complement"
);
Ok(())
}
#[test]
fn cross_validate_calls_each_index_once_as_the_test_point() -> Result<()> {
let mut seen: Vec<(Vec<usize>, Vec<usize>)> = Vec::new();
let scores = loo_cross_validate(4, |train, test| {
seen.push((train.to_vec(), test.to_vec()));
0.0
})?;
assert_eq!(scores.fold_scores().len(), 4, "one score per fold");
assert_eq!(
seen,
vec![
(vec![1, 2, 3], vec![0]),
(vec![0, 2, 3], vec![1]),
(vec![0, 1, 3], vec![2]),
(vec![0, 1, 2], vec![3]),
],
"evaluator must receive each index once with its ordered complement for training"
);
Ok(())
}
#[test]
fn cross_validate_computes_mean_and_std_error() -> Result<()> {
let predetermined = [0.5_f64, 1.5, 2.5, 3.5];
let mut next = predetermined.into_iter();
let scores = loo_cross_validate(4, |_train, _test| next.next().unwrap_or(f64::NAN))?;
for (i, (&got, &want)) in scores
.fold_scores()
.iter()
.zip(predetermined.iter())
.enumerate()
{
assert!(
(got - want).abs() < 1e-12,
"fold {i} score was {got}, want {want}"
);
}
assert!(
(scores.mean() - 2.0).abs() < 1e-12,
"mean was {}",
scores.mean()
);
assert!(
(scores.std_error() - 0.645_497_224_367_902_8).abs() < 1e-12,
"std_error was {}",
scores.std_error()
);
Ok(())
}
#[test]
fn cross_validate_predict_train_mean_squared_error_golden() -> Result<()> {
let data = [2.0_f64, 4.0, 6.0, 8.0, 10.0];
let scores = loo_cross_validate(data.len(), |train, test| {
let train_sum: f64 = train
.iter()
.map(|&j| data.get(j).copied().unwrap_or(f64::NAN))
.sum();
let train_count = f64::from(u32::try_from(train.len()).unwrap_or(0));
let prediction = train_sum / train_count;
let held_out = test
.first()
.and_then(|&i| data.get(i).copied())
.unwrap_or(f64::NAN);
(prediction - held_out).powi(2)
})?;
let expected = [25.0_f64, 6.25, 0.0, 6.25, 25.0];
for (i, (&got, &want)) in scores.fold_scores().iter().zip(expected.iter()).enumerate() {
assert!(
(got - want).abs() < 1e-10,
"fold {i} error was {got}, want {want}"
);
}
assert!(
(scores.mean() - 12.5).abs() < 1e-10,
"mean was {}",
scores.mean()
);
assert!(
(scores.std_error() - 5.229_125_165_837_972).abs() < 1e-10,
"std_error was {}",
scores.std_error()
);
Ok(())
}
#[test]
fn loo_cross_validate_returns_unified_cv_scores() -> Result<()> {
let scores: CvScores = loo_cross_validate(4, |_train, _test| 1.0)?;
assert_eq!(scores.fold_scores().len(), 4, "one score per fold");
assert!(
(scores.mean() - 1.0).abs() < 1e-12,
"mean was {}",
scores.mean()
);
Ok(())
}
#[test]
fn run_delegates_to_cross_validate() -> Result<()> {
let scheme = LeaveOneOutCrossValidation::default();
let scores = scheme.run(5, |_train, _test| 3.0)?;
assert_eq!(scores.fold_scores().len(), 5, "one score per fold");
assert!(
scores
.fold_scores()
.iter()
.all(|&s| (s - 3.0).abs() < 1e-12),
"every fold score should be the constant 3.0"
);
assert!(
(scores.mean() - 3.0).abs() < 1e-12,
"mean was {}",
scores.mean()
);
assert!(
scores.std_error().abs() < 1e-12,
"std_error was {}",
scores.std_error()
);
Ok(())
}
#[test]
fn accessors_expose_the_stored_summaries() -> Result<()> {
let scores = loo_cross_validate(5, |_train, _test| 3.0)?;
assert_eq!(
scores.fold_scores().len(),
5,
"fold_scores accessor exposes one score per fold"
);
assert!(
scores
.fold_scores()
.iter()
.all(|&s| (s - 3.0).abs() < 1e-12),
"fold_scores accessor returns the stored slice"
);
assert!(
(scores.mean() - 3.0).abs() < 1e-12,
"mean accessor was {}",
scores.mean()
);
assert!(
scores.std_error().abs() < 1e-12,
"std_error accessor was {}",
scores.std_error()
);
Ok(())
}
#[test]
fn loo_indices_two_gives_both_singleton_folds() -> Result<()> {
assert_eq!(
loo_indices(2)?,
vec![(vec![1], vec![0]), (vec![0], vec![1])],
"n=2 LOO folds must be ([1],[0]) then ([0],[1])"
);
Ok(())
}
#[test]
fn n_two_aggregates_mean_and_standard_error() -> Result<()> {
let predetermined = [1.0_f64, 3.0];
let mut next = predetermined.into_iter();
let scores = loo_cross_validate(2, |_train, _test| next.next().unwrap_or(f64::NAN))?;
assert_eq!(
scores.fold_scores(),
&predetermined,
"fold scores recorded in order"
);
assert!(
(scores.mean() - 2.0).abs() < 1e-12,
"n=2 mean was {}, expected 2.0",
scores.mean()
);
assert!(
(scores.std_error() - 1.0).abs() < 1e-12,
"n=2 std_error was {}, expected sd(ddof=1)/sqrt(2) = 1.0",
scores.std_error()
);
Ok(())
}
}