use crate::error::{DatarustError, Result};
use crate::matrix::Matrix;
use crate::model_selection::kfold::KFold;
use crate::traits::Predictor;
pub fn cross_val_score<T, F>(
estimator: &T,
x: &Matrix,
y: &[f64],
cv: &KFold,
scorer: F,
) -> Result<Vec<f64>>
where
T: Predictor + Clone,
F: Fn(&[f64], &[f64]) -> Result<f64>,
{
let n = x.nrows();
if y.len() != n {
return Err(DatarustError::ShapeMismatch {
expected: format!("{} targets", n),
actual: format!("{} targets", y.len()),
});
}
let mut scores = Vec::new();
for (train_idx, test_idx) in cv.split(n)? {
let x_train = x.select_rows(&train_idx)?;
let x_test = x.select_rows(&test_idx)?;
let y_train: Vec<f64> = train_idx.iter().map(|&i| y[i]).collect();
let y_test: Vec<f64> = test_idx.iter().map(|&i| y[i]).collect();
let mut model = estimator.clone();
model.fit(&x_train, &y_train)?;
let pred = model.predict(&x_test)?;
scores.push(scorer(&y_test, &pred)?);
}
Ok(scores)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::linear_model::{LinearRegression, LogisticRegression};
use crate::metrics::classification::accuracy_score;
use crate::metrics::regression::r2_score;
fn regression_data(n: usize) -> (Matrix, Vec<f64>) {
let rows: Vec<Vec<f64>> = (0..n)
.map(|i| {
let i = i as f64;
vec![i, i.sin()]
})
.collect();
let y: Vec<f64> = rows.iter().map(|r| 2.0 * r[0] + r[1]).collect();
(Matrix::new(rows).unwrap(), y)
}
#[test]
fn returns_one_score_per_fold() {
let (x, y) = regression_data(30);
let cv = KFold::new().with_n_splits(5);
let scores = cross_val_score(&LinearRegression::new(), &x, &y, &cv, r2_score).unwrap();
assert_eq!(scores.len(), 5);
}
#[test]
fn high_score_on_clean_linear_signal() {
let (x, y) = regression_data(30);
let cv = KFold::new().with_n_splits(3);
let scores = cross_val_score(&LinearRegression::new(), &x, &y, &cv, r2_score).unwrap();
for s in &scores {
assert!(*s > 0.99, "low R² score: {s}");
}
}
#[test]
fn classification_uses_accuracy_scorer() {
let rows: Vec<Vec<f64>> = (-10..=10)
.filter(|&i| i != 0)
.map(|i| vec![i as f64 * 0.5])
.collect();
let x = Matrix::new(rows.clone()).unwrap();
let y: Vec<f64> = rows
.iter()
.map(|r| if r[0] > 0.0 { 1.0 } else { 0.0 })
.collect();
let cv = KFold::new().with_n_splits(3);
let scores =
cross_val_score(&LogisticRegression::new(), &x, &y, &cv, accuracy_score).unwrap();
assert_eq!(scores.len(), 3);
for s in &scores {
assert!((*s - 1.0).abs() < 1e-9, "low accuracy: {s}");
}
}
#[test]
fn works_with_custom_closure_scorer() {
let (x, y) = regression_data(20);
let cv = KFold::new().with_n_splits(2);
let mse_scorer = |y_true: &[f64], y_pred: &[f64]| {
let n = y_true.len() as f64;
let s: f64 = y_true
.iter()
.zip(y_pred.iter())
.map(|(t, p)| (t - p).powi(2))
.sum();
Ok(s / n)
};
let scores = cross_val_score(&LinearRegression::new(), &x, &y, &cv, mse_scorer).unwrap();
assert_eq!(scores.len(), 2);
for s in &scores {
assert!(*s >= 0.0);
}
}
#[test]
fn target_length_mismatch_is_rejected_before_indexing() {
let (x, _) = regression_data(10);
let cv = KFold::new().with_n_splits(2);
let err = cross_val_score(&LinearRegression::new(), &x, &[1.0], &cv, r2_score).unwrap_err();
assert!(matches!(err, DatarustError::ShapeMismatch { .. }));
}
}