use super::eval::EvalSet;
use super::train::{train_impl, with_thread_pool};
use super::validate::validate_trained_model;
use crate::config::TrainingParams;
use crate::data::DMatrix;
use crate::error::Result;
use crate::metric::Metric;
use crate::model::BoostedModel;
use std::num::NonZeroUsize;
use std::ops::ControlFlow;
#[derive(Debug, Clone, Default, PartialEq)]
pub struct EvalHistory {
datasets: Vec<String>,
metrics: Vec<String>,
first_iteration: usize,
values: Vec<f64>,
}
impl EvalHistory {
pub(super) fn new(datasets: Vec<String>, metrics: Vec<String>, first_iteration: usize) -> Self {
EvalHistory {
datasets,
metrics,
first_iteration,
values: Vec::new(),
}
}
pub(super) fn push_round(&mut self, values: impl IntoIterator<Item = f64>) {
let before = self.values.len();
self.values.extend(values);
debug_assert_eq!(self.values.len() - before, self.round_width());
}
pub fn datasets(&self) -> &[String] {
&self.datasets
}
pub fn metrics(&self) -> &[String] {
&self.metrics
}
pub fn first_iteration(&self) -> usize {
self.first_iteration
}
pub fn len(&self) -> usize {
match self.round_width() {
0 => 0,
width => self.values.len() / width,
}
}
pub fn is_empty(&self) -> bool {
self.values.is_empty()
}
pub fn rounds(&self) -> impl ExactSizeIterator<Item = RoundEval<'_>> + DoubleEndedIterator {
(0..self.len()).map(|round| self.view(round))
}
pub fn round(&self, iteration: usize) -> Option<RoundEval<'_>> {
let round = iteration.checked_sub(self.first_iteration)?;
(round < self.len()).then(|| self.view(round))
}
pub fn last(&self) -> Option<RoundEval<'_>> {
self.len().checked_sub(1).map(|round| self.view(round))
}
pub fn series(
&self,
dataset: &str,
metric: &str,
) -> Option<impl ExactSizeIterator<Item = f64> + '_> {
let cell = self.cell(dataset, metric)?;
let width = self.round_width();
Some((0..self.len()).map(move |round| self.values[round * width + cell]))
}
fn round_width(&self) -> usize {
self.datasets.len() * self.metrics.len()
}
fn cell(&self, dataset: &str, metric: &str) -> Option<usize> {
let d = self.datasets.iter().position(|name| name == dataset)?;
let m = self.metrics.iter().position(|name| name == metric)?;
Some(d * self.metrics.len() + m)
}
fn view(&self, round: usize) -> RoundEval<'_> {
let width = self.round_width();
RoundEval {
iteration: self.first_iteration + round,
history: Some(self),
values: &self.values[round * width..(round + 1) * width],
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct RoundEval<'a> {
iteration: usize,
history: Option<&'a EvalHistory>,
values: &'a [f64],
}
impl<'a> RoundEval<'a> {
pub(super) fn unscored(iteration: usize) -> Self {
RoundEval {
iteration,
history: None,
values: &[],
}
}
pub fn iteration(&self) -> usize {
self.iteration
}
pub fn values(&self) -> &'a [f64] {
self.values
}
pub fn score(&self, dataset: &str, metric: &str) -> Option<f64> {
let cell = self.history?.cell(dataset, metric)?;
Some(self.values[cell])
}
pub fn scores(&self) -> impl Iterator<Item = (&'a str, &'a str, f64)> + 'a {
let (datasets, metrics): (&[String], &[String]) = match self.history {
Some(history) => (&history.datasets, &history.metrics),
None => (&[], &[]),
};
datasets
.iter()
.flat_map(move |dataset| metrics.iter().map(move |metric| (dataset, metric)))
.zip(self.values)
.map(|((dataset, metric), &value)| (dataset.as_str(), metric.as_str(), value))
}
}
#[derive(Debug)]
#[non_exhaustive]
pub struct TrainResult {
pub model: BoostedModel,
pub history: EvalHistory,
pub best_score: Option<f64>,
}
pub fn train(
params: &TrainingParams,
dtrain: &DMatrix,
num_boost_round: usize,
) -> Result<BoostedModel> {
Ok(Trainer::new(params, dtrain, num_boost_round).train()?.model)
}
pub struct Trainer<'a> {
pub(super) params: &'a TrainingParams,
pub(super) dtrain: &'a DMatrix,
pub(super) num_boost_round: usize,
pub(super) evals: Vec<EvalSet<'a>>,
pub(super) early_stopping_rounds: Option<NonZeroUsize>,
pub(super) metric: Option<Box<dyn Metric>>,
pub(super) init_model: Option<&'a BoostedModel>,
pub(super) on_round: Option<RoundHook<'a>>,
}
pub(super) type RoundHook<'a> = Box<dyn FnMut(RoundEval<'_>) -> ControlFlow<()> + Send + 'a>;
impl<'a> Trainer<'a> {
pub fn new(params: &'a TrainingParams, dtrain: &'a DMatrix, num_boost_round: usize) -> Self {
Trainer {
params,
dtrain,
num_boost_round,
evals: Vec::new(),
early_stopping_rounds: None,
metric: None,
init_model: None,
on_round: None,
}
}
#[must_use]
pub fn eval(mut self, data: &'a DMatrix, name: &'a str) -> Self {
self.evals.push(EvalSet { data, name });
self
}
#[must_use]
pub fn early_stopping_rounds(mut self, rounds: NonZeroUsize) -> Self {
self.early_stopping_rounds = Some(rounds);
self
}
#[must_use]
pub fn custom_metric(mut self, metric: Box<dyn Metric>) -> Self {
self.metric = Some(metric);
self
}
#[must_use]
pub fn init_model(mut self, model: &'a BoostedModel) -> Self {
self.init_model = Some(model);
self
}
#[must_use]
pub fn on_round(
mut self,
hook: impl FnMut(RoundEval<'_>) -> ControlFlow<()> + Send + 'a,
) -> Self {
self.on_round = Some(Box::new(hook));
self
}
pub fn train(self) -> Result<TrainResult> {
let loss = self.params.loss(self.dtrain.n_targets())?;
let params = self.params;
let result = with_thread_pool(params, || train_impl(self, loss.as_ref()))?;
validate_trained_model(&result.model)?;
Ok(result)
}
}