use ndarray::Array1;
use model_selection_rs::splitters::{
CvSplitter, KFold as MsKFold, StratifiedKFold as MsStratifiedKFold,
};
use crate::error::{Error, Result};
use crate::frame::Dataset;
use crate::traits::Model;
use super::scoring::Metric;
fn ms_err(e: impl std::fmt::Display) -> Error {
Error::Backend(format!("model-selection: {e}"))
}
pub trait CrossValidator: CrossValidatorClone + Send + Sync {
fn splits(&self, dataset: &Dataset) -> Result<Vec<(Vec<usize>, Vec<usize>)>>;
fn n_splits(&self) -> usize;
}
pub trait CrossValidatorClone {
fn clone_box(&self) -> Box<dyn CrossValidator>;
}
impl<T> CrossValidatorClone for T
where
T: CrossValidator + Clone + 'static,
{
fn clone_box(&self) -> Box<dyn CrossValidator> {
Box::new(self.clone())
}
}
impl Clone for Box<dyn CrossValidator> {
fn clone(&self) -> Self {
self.clone_box()
}
}
#[derive(Clone, Copy, Debug)]
pub struct KFold {
k: usize,
}
impl KFold {
pub fn new(k: usize) -> Self {
KFold { k }
}
}
impl CrossValidator for KFold {
fn splits(&self, dataset: &Dataset) -> Result<Vec<(Vec<usize>, Vec<usize>)>> {
MsKFold::new(self.k)
.map_err(ms_err)?
.split(dataset.features().nrows())
.map_err(ms_err)
}
fn n_splits(&self) -> usize {
self.k
}
}
#[derive(Clone, Copy, Debug)]
pub struct StratifiedKFold {
k: usize,
}
impl StratifiedKFold {
pub fn new(k: usize) -> Self {
StratifiedKFold { k }
}
}
impl CrossValidator for StratifiedKFold {
fn splits(&self, dataset: &Dataset) -> Result<Vec<(Vec<usize>, Vec<usize>)>> {
let labels: Array1<i64> = Array1::from(
dataset
.target()
.iter()
.map(|v| v.round() as i64)
.collect::<Vec<_>>(),
);
MsStratifiedKFold::new(self.k, &labels)
.map_err(ms_err)?
.split(labels.len())
.map_err(ms_err)
}
fn n_splits(&self) -> usize {
self.k
}
}
pub fn cross_val_score(
model: &dyn Model,
dataset: &Dataset,
cv: &dyn CrossValidator,
metric: Metric,
) -> Result<f64> {
let splits = cv.splits(dataset)?;
if splits.is_empty() {
return Err(Error::Pipeline(
"cross-validation produced no splits".into(),
));
}
use rayon::prelude::*;
let scores: Vec<f64> = splits
.par_iter()
.map(|(train, test)| -> Result<f64> {
let mut m = model.clone_box();
m.fit(&dataset.select(train))?;
let preds = m.predict(&dataset.features().select_rows(test))?;
let truth: Vec<f64> = test.iter().map(|&i| dataset.target()[i]).collect();
Ok(metric.score(&truth, &preds))
})
.collect::<Result<Vec<f64>>>()?;
Ok(scores.iter().sum::<f64>() / scores.len() as f64)
}