mod fold;
pub use fold::Fold;
use crate::config::TrainingParams;
use crate::data::DMatrix;
use crate::data::target_stats::OrderedTargetEncoder;
use crate::error::{HessboostError, Result};
use crate::training::Trainer;
use crate::training::eval::{EarlyStopping, configured_metrics};
use std::num::NonZeroUsize;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct CvResult {
pub metric: String,
pub rounds: Vec<CvRound>,
}
#[derive(Debug, Clone, Copy, PartialEq)]
#[non_exhaustive]
pub struct CvRound {
pub mean: f64,
pub std: f64,
}
#[derive(Debug)]
pub struct CrossValidation<'a> {
params: &'a TrainingParams,
data: &'a DMatrix,
num_boost_round: usize,
folds: Vec<Fold>,
early_stopping_rounds: Option<NonZeroUsize>,
target_stats: Option<(OrderedTargetEncoder, Vec<usize>)>,
}
impl<'a> CrossValidation<'a> {
pub fn new(
params: &'a TrainingParams,
data: &'a DMatrix,
num_boost_round: usize,
folds: Vec<Fold>,
) -> Self {
CrossValidation {
params,
data,
num_boost_round,
folds,
early_stopping_rounds: None,
target_stats: None,
}
}
#[must_use]
pub fn early_stopping_rounds(mut self, rounds: NonZeroUsize) -> Self {
self.early_stopping_rounds = Some(rounds);
self
}
#[must_use]
pub fn target_stats(mut self, encoder: OrderedTargetEncoder, columns: Vec<usize>) -> Self {
self.target_stats = Some((encoder, columns));
self
}
pub fn run(self) -> Result<Vec<CvResult>> {
let CrossValidation {
params,
data,
num_boost_round,
folds,
early_stopping_rounds,
target_stats,
} = self;
if folds.is_empty() {
return Err(HessboostError::invalid_param("folds", "no folds"));
}
validate_folds(&folds, data)?;
let objective = params.loss(data.n_targets())?;
let metrics = configured_metrics(params, objective.as_ref())?;
let maximize = metrics.last().is_some_and(|m| m.maximize());
let run = FoldRun {
params,
data,
num_boost_round,
target_stats: target_stats.as_ref(),
};
let values = run.scores(&folds, metrics.len())?;
let mut out: Vec<CvResult> = metrics
.iter()
.zip(values)
.map(|(metric, per_round)| aggregate(metric.name(), &per_round))
.collect();
if let Some(patience) = early_stopping_rounds
&& let Some(watched) = out.last()
&& !watched.rounds.is_empty()
{
let means = watched.rounds.iter().map(|round| round.mean);
let end = best_round(means, patience, maximize) + 1;
for result in &mut out {
result.rounds.truncate(end);
}
}
Ok(out)
}
}
fn validate_folds(folds: &[Fold], data: &DMatrix) -> Result<()> {
let n = data.n_rows();
for (f, fold) in folds.iter().enumerate() {
for (name, rows) in [("training", &fold.train), ("test", &fold.test)] {
if rows.is_empty() {
return Err(HessboostError::invalid_param(
"folds",
format!("fold {f} has no {name} rows"),
));
}
if let Some(&row) = rows.iter().find(|&&row| row >= n) {
return Err(HessboostError::invalid_param(
"folds",
format!("fold {f}: {name} row {row} is out of bounds for {n} rows"),
));
}
data.selected_group_sizes(rows).map_err(|reason| {
HessboostError::invalid_param(
"folds",
format!("fold {f}: its {name} rows {reason}"),
)
})?;
}
}
Ok(())
}
struct FoldRun<'a> {
params: &'a TrainingParams,
data: &'a DMatrix,
num_boost_round: usize,
target_stats: Option<&'a (OrderedTargetEncoder, Vec<usize>)>,
}
impl FoldRun<'_> {
fn scores(&self, folds: &[Fold], n_metrics: usize) -> Result<Vec<Vec<Vec<f64>>>> {
let mut values: Vec<Vec<Vec<f64>>> = vec![Vec::new(); n_metrics];
for fold in folds {
let (dtrain, dtest) = self.fold_data(fold)?;
let res = Trainer::new(self.params, &dtrain, self.num_boost_round)
.eval(&dtest, "test")
.train()?;
for (round, eval) in res.history.rounds().enumerate() {
for (per_round, &value) in values.iter_mut().zip(eval.values()) {
if per_round.len() == round {
per_round.push(Vec::with_capacity(folds.len()));
}
per_round[round].push(value);
}
}
}
Ok(values)
}
fn fold_data(&self, fold: &Fold) -> Result<(DMatrix, DMatrix)> {
let dtrain = self.data.select_rows(&fold.train)?;
let dtest = self.data.select_rows(&fold.test)?;
let Some((encoder, columns)) = self.target_stats else {
return Ok((dtrain, dtest));
};
let (dtrain, fitted) = encoder.fit_transform(&dtrain, columns)?;
let dtest = fitted.transform(&dtest)?;
Ok((dtrain, dtest))
}
}
fn aggregate(metric: &str, per_round: &[Vec<f64>]) -> CvResult {
let rounds = per_round
.iter()
.map(|vals| {
let len = vals.len() as f64;
let mean = vals.iter().sum::<f64>() / len;
let var = vals.iter().map(|v| (v - mean) * (v - mean)).sum::<f64>() / len;
CvRound {
mean,
std: var.sqrt(),
}
})
.collect();
CvResult {
metric: metric.to_string(),
rounds,
}
}
fn best_round(scores: impl Iterator<Item = f64>, patience: NonZeroUsize, maximize: bool) -> usize {
let mut stopping = EarlyStopping::new(patience, maximize, 0);
for (round, score) in scores.enumerate() {
if stopping.observe(round, score) {
break;
}
}
stopping.best_round()
}
pub fn cv(
params: &TrainingParams,
data: &DMatrix,
num_boost_round: usize,
nfold: usize,
seed: u64,
) -> Result<Vec<CvResult>> {
let folds = Fold::k_fold(data.n_rows(), nfold, seed)?;
CrossValidation::new(params, data, num_boost_round, folds).run()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::objective::{Objective, RegLoss};
use crate::test_support::labeled_dense;
#[test]
fn cv_reports_decreasing_rmse() {
let n = 200;
let mut x = Vec::new();
let mut y = Vec::new();
for i in 0..n {
let xi = i as f32 / n as f32;
x.push(xi);
y.push(if xi > 0.5 { 1.0 } else { 0.0 });
}
let d = labeled_dense(&x, n, 1, &y);
let params = TrainingParams::builder()
.objective(Objective::SquaredError(RegLoss::default()))
.max_depth(3)
.eta(0.3)
.build()
.unwrap();
let results = cv(¶ms, &d, 30, 5, 42).unwrap();
assert_eq!(results.len(), 1);
let rmse = &results[0];
assert_eq!(rmse.metric, "rmse");
assert_eq!(rmse.rounds.len(), 30);
assert!(rmse.rounds[29].mean < rmse.rounds[0].mean);
assert!(
rmse.rounds
.iter()
.all(|r| r.std.is_finite() && r.std >= 0.0)
);
}
fn step_data(n: usize) -> DMatrix {
let x: Vec<f32> = (0..n).map(|i| i as f32 / n as f32).collect();
let y: Vec<f32> = x.iter().map(|&v| if v > 0.5 { 1.0 } else { 0.0 }).collect();
labeled_dense(&x, n, 1, &y)
}
#[test]
fn forward_chaining_trains_only_on_rows_before_the_gap() {
let folds = Fold::forward_chaining(23, 3, 2).unwrap();
let expect = [(0..6, 8..13), (0..11, 13..18), (0..16, 18..23)];
assert_eq!(folds.len(), 3);
for (fold, (train, test)) in folds.iter().zip(expect) {
assert_eq!(fold.train, train.collect::<Vec<_>>());
assert_eq!(fold.test, test.collect::<Vec<_>>());
}
assert_eq!(Fold::forward_chaining(23, 3, 7).unwrap()[0].train, [0]);
assert!(Fold::forward_chaining(23, 3, 8).is_err());
let tight = Fold::forward_chaining(4, 3, 0).unwrap();
assert_eq!(tight[0], Fold::new(vec![0], vec![1]));
assert_eq!(tight[2], Fold::new(vec![0, 1, 2], vec![3]));
assert!(Fold::forward_chaining(3, 3, 0).is_err());
assert!(Fold::forward_chaining(10, usize::MAX, 0).is_err());
assert!(Fold::forward_chaining(10, 2, usize::MAX).is_err());
assert!(Fold::forward_chaining(10, 0, 0).is_err());
}
#[test]
fn caller_folds_are_checked() {
let d = step_data(20);
let params = TrainingParams::default();
let run = |folds: Vec<Fold>| CrossValidation::new(¶ms, &d, 2, folds).run();
assert!(run(Vec::new()).is_err());
assert!(run(vec![Fold::new(vec![], vec![1])]).is_err());
assert!(run(vec![Fold::new(vec![0], vec![])]).is_err());
assert!(run(vec![Fold::new(vec![0, 20], vec![1])]).is_err());
assert!(run(vec![Fold::new(vec![0, 1], vec![19])]).is_ok());
}
#[test]
fn metrics_keep_their_configured_order() {
let d = step_data(60);
let params = TrainingParams::builder()
.eval_metric(crate::metric::EvalMetric::Rmse)
.eval_metric(crate::metric::EvalMetric::Mae)
.build()
.unwrap();
let results = cv(¶ms, &d, 3, 3, 1).unwrap();
let names: Vec<&str> = results.iter().map(|r| r.metric.as_str()).collect();
assert_eq!(names, ["rmse", "mae"]);
}
#[test]
fn early_stopping_ends_at_the_best_mean_round() {
let d = step_data(120);
let params = TrainingParams::builder()
.eval_metric(crate::metric::EvalMetric::Mae)
.eval_metric(crate::metric::EvalMetric::Rmse)
.max_depth(6)
.eta(0.8)
.build()
.unwrap();
let folds = Fold::k_fold(d.n_rows(), 4, 3).unwrap();
let full = CrossValidation::new(¶ms, &d, 40, folds.clone())
.run()
.unwrap();
let rmse: Vec<f64> = full[1].rounds.iter().map(|r| r.mean).collect();
let best = (0..rmse.len())
.min_by(|&a, &b| rmse[a].total_cmp(&rmse[b]))
.unwrap();
let stopped = CrossValidation::new(¶ms, &d, 40, folds)
.early_stopping_rounds(NonZeroUsize::new(5).unwrap())
.run()
.unwrap();
let end = stopped[1].rounds.len();
assert!(end + 5 <= 40, "ended after {end} rounds");
let b = end - 1;
assert!(rmse[..b].iter().all(|&v| v > rmse[b]));
assert!(rmse[b + 1..=b + 5].iter().all(|&v| v >= rmse[b]));
for (s, f) in stopped.iter().zip(&full) {
assert_eq!(s.rounds, f.rounds[..end]);
}
let folds = Fold::k_fold(d.n_rows(), 4, 3).unwrap();
let long = CrossValidation::new(¶ms, &d, 40, folds)
.early_stopping_rounds(NonZeroUsize::new(100).unwrap())
.run()
.unwrap();
assert_eq!(long[1].rounds, full[1].rounds[..=best]);
}
fn ranking_data(n_groups: usize) -> DMatrix {
let x: Vec<f32> = (0..n_groups * 4).map(|i| ((i * 5) % 7) as f32).collect();
let y: Vec<f32> = x.iter().map(|&v| (v / 2.0).floor()).collect();
labeled_dense(&x, x.len(), 1, &y)
.with_group_sizes(&vec![4; n_groups])
.unwrap()
}
#[test]
fn ranking_folds_keep_whole_query_groups() {
let d = ranking_data(6);
let params = TrainingParams::builder()
.objective(Objective::RankNdcg(crate::objective::LambdaRank::default()))
.max_depth(2)
.build()
.unwrap();
let folds: Vec<Fold> = (0..3)
.map(|f| {
let (test, train) = (0..24).partition(|&row| row / 8 == f);
Fold::new(train, test)
})
.collect();
let results = CrossValidation::new(¶ms, &d, 3, folds).run().unwrap();
assert_eq!(results[0].metric, "ndcg@32");
assert!(results[0].rounds.iter().all(|r| r.mean.is_finite()));
let shuffled = Fold::k_fold(24, 3, 0).unwrap();
assert!(matches!(
CrossValidation::new(¶ms, &d, 3, shuffled).run(),
Err(HessboostError::InvalidParameter { name: "folds", .. })
));
}
#[test]
fn target_stats_are_fitted_on_each_folds_training_rows() {
use crate::data::FeatureType;
let n = 90;
let x: Vec<f32> = (0..n)
.flat_map(|i| [(i % 5) as f32, ((i * 13) % 11) as f32])
.collect();
let y: Vec<f32> = (0..n)
.map(|i| (i % 5) as f32 + ((i * 7) % 3) as f32)
.collect();
let d = labeled_dense(&x, n, 2, &y)
.with_feature_types(&[FeatureType::Categorical, FeatureType::Numerical])
.unwrap();
let params = TrainingParams::builder().max_depth(2).build().unwrap();
let encoder = OrderedTargetEncoder::builder().seed(4).build().unwrap();
let folds = Fold::k_fold(n, 3, 7).unwrap();
let results = CrossValidation::new(¶ms, &d, 5, folds.clone())
.target_stats(encoder.clone(), vec![0])
.run()
.unwrap();
let mut last = Vec::new();
for fold in &folds {
let (dtrain, fitted) = encoder
.fit_transform(&d.select_rows(&fold.train).unwrap(), &[0])
.unwrap();
let dtest = fitted
.transform(&d.select_rows(&fold.test).unwrap())
.unwrap();
let res = Trainer::new(¶ms, &dtrain, 5)
.eval(&dtest, "test")
.train()
.unwrap();
last.push(res.history.rounds().last().unwrap().values()[0]);
}
let mean = last.iter().sum::<f64>() / last.len() as f64;
assert_eq!(results[0].rounds[4].mean, mean);
let folds = Fold::k_fold(n, 3, 7).unwrap();
assert!(matches!(
CrossValidation::new(¶ms, &d, 5, folds)
.target_stats(encoder, vec![1])
.run(),
Err(HessboostError::InvalidParameter {
name: "columns",
..
})
));
}
}