use super::api::{EvalHistory, RoundEval, RoundHook, TrainResult};
use super::margins::MarginCaches;
use crate::config::TrainingParams;
use crate::data::{DMatrix, MetaInfo};
use crate::error::{HessboostError, Result};
use crate::metric::Metric;
use crate::model::BoostedModel;
use crate::objective::Loss;
use std::num::NonZeroUsize;
use std::ops::ControlFlow;
#[derive(Clone, Copy)]
pub(super) struct EvalSet<'a> {
pub(super) data: &'a DMatrix,
pub(super) name: &'a str,
}
pub(super) struct EvalPlan<'a> {
objective: &'a dyn Loss,
evals: &'a [EvalSet<'a>],
infos: Vec<MetaInfo<'a>>,
pub(super) metrics: Vec<Box<dyn Metric>>,
preds: Vec<f32>,
scores: Vec<f64>,
}
impl<'a> EvalPlan<'a> {
pub(super) fn new(
params: &TrainingParams,
objective: &'a dyn Loss,
metric_override: Option<Box<dyn Metric>>,
evals: &'a [EvalSet<'a>],
n_targets: usize,
) -> Result<Self> {
let n_out = objective.n_outputs();
let mut metrics = configured_metrics(params, objective)?;
metrics.extend(metric_override);
if n_targets > 1
&& let Some(metric) = metrics.iter().find(|m| !m.supports_label_matrix())
{
return Err(HessboostError::invalid_param(
"eval_metric",
format!(
"metric `{}` does not support multi-target labels",
metric.name()
),
));
}
let infos: Vec<MetaInfo> = evals.iter().map(|set| set.data.info()).collect();
for (info, set) in infos.iter().zip(evals) {
for metric in &metrics {
metric
.validate_info(info)
.map_err(|error| error.in_dataset(set.name))?;
check_prediction_width(metric.as_ref(), info, n_out, set.name)?;
}
}
Ok(EvalPlan {
objective,
evals,
infos,
metrics,
preds: Vec::new(),
scores: Vec::new(),
})
}
pub(super) fn maximize(&self) -> bool {
self.metrics.last().is_some_and(|m| m.maximize())
}
fn history(&self, first_iteration: usize) -> EvalHistory {
EvalHistory::new(
self.evals.iter().map(|set| set.name.to_owned()).collect(),
self.metrics.iter().map(|m| m.name().to_owned()).collect(),
first_iteration,
)
}
fn record(&mut self, margins: &MarginCaches, history: &mut EvalHistory) -> f64 {
self.scores.clear();
let mut last_metric_value = 0.0;
for ei in 0..self.evals.len() {
self.preds.clear();
self.preds.extend_from_slice(&margins.evals[ei]);
self.objective.eval_transform(&mut self.preds);
for m in &self.metrics {
let v = m.eval_info(&self.preds, &self.infos[ei]);
self.scores.push(v);
last_metric_value = v;
}
}
history.push_round(self.scores.iter().copied());
last_metric_value
}
}
pub(super) struct EarlyStopping {
patience: NonZeroUsize,
maximize: bool,
best_score: f64,
best_round: usize,
since_improved: usize,
}
impl EarlyStopping {
pub(crate) fn new(patience: NonZeroUsize, maximize: bool, first_round: usize) -> Self {
EarlyStopping {
patience,
maximize,
best_score: if maximize {
f64::NEG_INFINITY
} else {
f64::INFINITY
},
best_round: first_round,
since_improved: 0,
}
}
pub(crate) fn observe(&mut self, round: usize, score: f64) -> bool {
let improved = if self.maximize {
score > self.best_score
} else {
score < self.best_score
};
if improved {
self.best_score = score;
self.best_round = round;
self.since_improved = 0;
false
} else {
self.since_improved += 1;
self.since_improved >= self.patience.get()
}
}
pub(crate) fn best_round(&self) -> usize {
self.best_round
}
}
pub(super) fn configured_metrics(
params: &TrainingParams,
loss: &dyn Loss,
) -> Result<Vec<Box<dyn Metric>>> {
let n_outputs = loss.n_outputs();
if params.eval_metric.is_empty() {
Ok(vec![loss.default_metric().build(n_outputs)?])
} else {
params
.eval_metric
.iter()
.map(|metric| metric.build(n_outputs))
.collect()
}
}
fn check_prediction_width(
metric: &dyn crate::metric::Metric,
info: &MetaInfo,
n_out: usize,
dataset: &str,
) -> Result<()> {
let reason = match metric.prediction_width(info) {
Some(width) if width != n_out => format!(
"metric `{}` reads {width} prediction(s) per row of dataset `{dataset}`, but the \
model has {n_out} outputs",
metric.name()
),
None if !n_out.is_multiple_of(info.n_targets().max(1)) => format!(
"metric `{}` needs a whole number of the model's {n_out} outputs per label \
column of dataset `{dataset}` ({} columns)",
metric.name(),
info.n_targets()
),
_ => return Ok(()),
};
Err(HessboostError::invalid_param("eval_metric", reason))
}
pub(super) struct RoundReporter<'a> {
eval_plan: Option<EvalPlan<'a>>,
history: EvalHistory,
stopping: Option<EarlyStopping>,
on_round: Option<RoundHook<'a>>,
}
impl<'a> RoundReporter<'a> {
pub(super) fn new(on_round: Option<RoundHook<'a>>) -> Self {
RoundReporter {
eval_plan: None,
history: EvalHistory::default(),
stopping: None,
on_round,
}
}
pub(super) fn watch(
&mut self,
plan: EvalPlan<'a>,
early_stopping_rounds: Option<NonZeroUsize>,
first_round: usize,
) {
self.stopping = early_stopping_rounds
.map(|patience| EarlyStopping::new(patience, plan.maximize(), first_round));
if !plan.evals.is_empty() {
self.history = plan.history(first_round);
self.eval_plan = Some(plan);
}
}
pub(super) fn finish_round(
&mut self,
iteration: usize,
margins: Option<&MarginCaches>,
) -> ControlFlow<()> {
let mut stop = false;
let mut scored = false;
if let (Some(plan), Some(margins)) = (&mut self.eval_plan, margins) {
let score = plan.record(margins, &mut self.history);
scored = true;
if let Some(stopping) = &mut self.stopping {
stop = stopping.observe(iteration, score);
}
}
if let Some(hook) = &mut self.on_round {
let round = match self.history.last() {
Some(round) if scored => round,
_ => RoundEval::unscored(iteration),
};
debug_assert_eq!(round.iteration(), iteration);
stop |= hook(round).is_break();
}
if stop {
ControlFlow::Break(())
} else {
ControlFlow::Continue(())
}
}
pub(super) fn into_result(self, mut model: BoostedModel) -> TrainResult {
let mut best_score = None;
if let Some(stopping) = &self.stopping
&& !self.history.is_empty()
{
let best_iter = stopping.best_round();
best_score = self
.history
.round(best_iter)
.and_then(|round| round.values().last().copied());
model.truncate_shrunk(best_iter + 1);
model.set_best_iteration(Some(best_iter));
}
TrainResult {
model,
history: self.history,
best_score,
}
}
}