use hessboost::config::{BoosterKind, Dart, ProcessType, Refresh};
use hessboost::model::ModelFormat;
use hessboost::prelude::{BoostedModel, Trainer, TrainingParams, train};
use hessboost::training::RoundEval;
use std::num::NonZeroUsize;
use std::ops::ControlFlow;
mod common;
use common::smooth::noisy as data;
fn bytes(model: &BoostedModel) -> Vec<u8> {
model.encode(ModelFormat::Binary).unwrap()
}
fn sampled() -> TrainingParams {
TrainingParams::builder()
.booster(BoosterKind::Dart(
Dart::builder().rate_drop(0.2).build().unwrap(),
))
.subsample(0.7)
.colsample_bynode(0.8)
.seed(9)
.max_depth(3)
.build()
.unwrap()
}
fn stop_at(last: usize) -> impl FnMut(RoundEval<'_>) -> ControlFlow<()> + Send {
move |round| {
if round.iteration() == last {
ControlFlow::Break(())
} else {
ControlFlow::Continue(())
}
}
}
#[test]
fn a_continuing_hook_changes_nothing_and_sees_every_round_in_order() {
let (dtrain, dvalid) = (data(300, 0), data(200, 1));
let params = sampled();
let plain = Trainer::new(¶ms, &dtrain, 12)
.eval(&dvalid, "valid")
.train()
.unwrap();
let mut seen = Vec::new();
let hooked = Trainer::new(¶ms, &dtrain, 12)
.eval(&dvalid, "valid")
.on_round(|round| {
seen.push((round.iteration(), round.values().to_vec()));
ControlFlow::Continue(())
})
.train()
.unwrap();
assert_eq!(bytes(&hooked.model), bytes(&plain.model));
let history: Vec<_> = plain
.history
.rounds()
.map(|round| (round.iteration(), round.values().to_vec()))
.collect();
assert_eq!(seen, history);
let mut iterations = Vec::new();
let unscored = Trainer::new(¶ms, &dtrain, 5)
.on_round(|round| {
assert!(round.values().is_empty());
iterations.push(round.iteration());
ControlFlow::Continue(())
})
.train()
.unwrap();
assert_eq!(iterations, [0, 1, 2, 3, 4]);
assert!(unscored.history.is_empty());
}
#[test]
fn break_keeps_the_rounds_so_far_as_a_shorter_run_would() {
let (dtrain, dvalid) = (data(300, 0), data(200, 1));
let params = sampled();
let stopped = Trainer::new(¶ms, &dtrain, 100)
.eval(&dvalid, "valid")
.on_round(stop_at(6))
.train()
.unwrap();
assert_eq!(stopped.model.num_boost_rounds(), 7);
assert_eq!(stopped.history.len(), 7);
assert_eq!(stopped.model.best_iteration(), None);
assert_eq!(stopped.best_score, None);
let short = Trainer::new(¶ms, &dtrain, 7)
.eval(&dvalid, "valid")
.train()
.unwrap();
assert_eq!(bytes(&stopped.model), bytes(&short.model));
let mut seen = Vec::new();
let continued = Trainer::new(¶ms, &dtrain, 10)
.init_model(&short.model)
.on_round(|round| {
seen.push(round.iteration());
ControlFlow::Break(())
})
.train()
.unwrap();
assert_eq!(seen, [7]);
assert_eq!(continued.model.num_boost_rounds(), 8);
}
#[test]
fn break_under_early_stopping_records_the_best_round_so_far() {
let (dtrain, dvalid) = (data(300, 0), data(200, 1));
let params = TrainingParams::builder()
.eta(0.5)
.max_depth(4)
.build()
.unwrap();
let full = Trainer::new(¶ms, &dtrain, 200)
.eval(&dvalid, "valid")
.early_stopping_rounds(NonZeroUsize::new(3).unwrap())
.train()
.unwrap();
let best = full.model.best_iteration().unwrap();
let last = full.model.num_boost_rounds() - 1;
assert!(best >= 2 && last == best + 3, "best {best}, last {last}");
let mut seen = Vec::new();
let watched = Trainer::new(¶ms, &dtrain, 200)
.eval(&dvalid, "valid")
.early_stopping_rounds(NonZeroUsize::new(3).unwrap())
.on_round(|round| {
seen.push(round.iteration());
ControlFlow::Continue(())
})
.train()
.unwrap();
assert_eq!(seen, (0..=last).collect::<Vec<_>>());
assert_eq!(bytes(&watched.model), bytes(&full.model));
let cut = best - 1;
let stopped = Trainer::new(¶ms, &dtrain, 200)
.eval(&dvalid, "valid")
.early_stopping_rounds(NonZeroUsize::new(3).unwrap())
.on_round(stop_at(cut))
.train()
.unwrap();
assert_eq!(stopped.model.num_boost_rounds(), cut + 1);
let scores: Vec<f64> = stopped
.history
.rounds()
.map(|round| round.values()[0])
.collect();
let best_so_far = (0..scores.len())
.min_by(|&a, &b| scores[a].total_cmp(&scores[b]))
.unwrap();
assert_eq!(stopped.model.best_iteration(), Some(best_so_far));
assert_eq!(stopped.best_score, Some(scores[best_so_far]));
}
#[test]
fn refresh_and_gblinear_stop_on_break_too() {
let dtrain = data(300, 0);
let base = TrainingParams::builder().max_depth(3).build().unwrap();
let old = train(&base, &data(300, 5), 6).unwrap();
let update = TrainingParams::builder()
.max_depth(3)
.process_type(ProcessType::Update(Refresh::default()))
.build()
.unwrap();
let refreshed = Trainer::new(&update, &dtrain, 6)
.init_model(&old)
.on_round(stop_at(2))
.train()
.unwrap();
let short = Trainer::new(&update, &dtrain, 3)
.init_model(&old)
.train()
.unwrap();
assert_eq!(refreshed.model.num_boost_rounds(), 3);
assert_eq!(bytes(&refreshed.model), bytes(&short.model));
let linear = TrainingParams::builder()
.booster(BoosterKind::GbLinear)
.build()
.unwrap();
let mut seen = Vec::new();
let stopped = Trainer::new(&linear, &dtrain, 50)
.on_round(|round| {
seen.push(round.iteration());
if round.iteration() == 3 {
ControlFlow::Break(())
} else {
ControlFlow::Continue(())
}
})
.train()
.unwrap();
assert_eq!(seen, [0, 1, 2, 3]);
let four = train(&linear, &dtrain, 4).unwrap();
assert_eq!(bytes(&stopped.model), bytes(&four));
let all = Trainer::new(&linear, &dtrain, 50)
.on_round(|_| ControlFlow::Continue(()))
.train()
.unwrap();
assert_eq!(
bytes(&all.model),
bytes(&train(&linear, &dtrain, 50).unwrap())
);
}