mod cache;
mod update;
use std::num::NonZeroUsize;
use std::ops::ControlFlow;
use self::cache::Cache;
use self::update::Incremental;
use super::api::{RoundEval, Trainer};
use super::eval::configured_metrics;
use super::train::{initial_intercepts, with_thread_pool};
use super::validate::{validate_trained_model, validate_training_data};
use crate::config::{
BoosterKind, Device, GrowPolicy, ProcessType, SamplingMethod, TrainingParams, TreeMethod,
};
use crate::data::quantile::HistCuts;
use crate::data::{DMatrix, FeatureType};
use crate::error::{HessboostError, Result};
use crate::model::BoostedModel;
use crate::objective::Objective;
use crate::tree::RegTree;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct OnlineParams {
mode: OnlineMode,
}
#[derive(Debug, Clone, Copy, PartialEq)]
#[non_exhaustive]
pub enum OnlineMode {
Exact,
#[non_exhaustive]
Approximate {
tolerance: f64,
},
}
impl Default for OnlineParams {
fn default() -> Self {
OnlineParams {
mode: OnlineMode::Approximate { tolerance: 0.1 },
}
}
}
impl OnlineParams {
pub fn exact() -> Self {
OnlineParams {
mode: OnlineMode::Exact,
}
}
pub fn approximate(tolerance: f64) -> Result<Self> {
crate::check::fraction("tolerance", tolerance)?;
Ok(OnlineParams {
mode: OnlineMode::Approximate { tolerance },
})
}
pub fn mode(&self) -> OnlineMode {
self.mode
}
fn is_approximate(self) -> bool {
matches!(self.mode, OnlineMode::Approximate { .. })
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub struct UpdateReport {
pub nodes_kept: usize,
pub subtrees_regrown: usize,
pub rows_refreshed: usize,
}
#[derive(Debug, Clone)]
pub struct OnlineModel {
params: TrainingParams,
online: OnlineParams,
data: DMatrix,
model: BoostedModel,
cache: Option<Cache>,
}
impl OnlineModel {
pub fn train(
params: &TrainingParams,
data: &DMatrix,
num_boost_round: usize,
online: OnlineParams,
) -> Result<Self> {
Self::train_with(params, data, num_boost_round, online, |_| {
ControlFlow::Continue(())
})
}
pub fn train_with(
params: &TrainingParams,
data: &DMatrix,
num_boost_round: usize,
online: OnlineParams,
on_round: impl FnMut(RoundEval<'_>) -> ControlFlow<()> + Send,
) -> Result<Self> {
check_supported(params, data, online)?;
let model = Trainer::new(params, data, num_boost_round)
.on_round(on_round)
.train()?
.model;
Self::from_model(model, params, data, online)
}
pub fn from_model(
model: BoostedModel,
params: &TrainingParams,
data: &DMatrix,
online: OnlineParams,
) -> Result<Self> {
check_supported(params, data, online)?;
configured_metrics(params, params.loss(1)?.as_ref())?;
let categorical = model
.trees()
.iter()
.any(|t| t.nodes().iter().any(|n| n.is_categorical));
if model.objective().built_in() != Some(¶ms.objective)
|| model.n_outputs() != 1
|| model.num_parallel_tree() != 1
|| model.n_features() != data.n_cols()
|| (categorical && online.is_approximate())
|| model.max_delta_step() != params.effective_max_delta_step()
|| model.linear().is_some()
|| model.boulevard().is_some()
|| model.ebm().is_some()
|| model.trees().iter().any(|t| t.linear_leaves().is_some())
|| model
.trees()
.iter()
.any(|t| splits_below(t, params.max_depth))
|| (0..model.num_trees()).any(|t| model.tree_weight(t) != 1.0)
|| model.shrinkage().is_some()
{
return Err(HessboostError::incompatible_model(
"model",
"not a single-output, unweighted, unshrunk, numeric gbtree model (not \
Boulevard or EBM) with constant leaves of these parameters and data",
));
}
if let Some(best) = model.best_iteration() {
return Err(HessboostError::incompatible_model(
"model",
format!(
"an early-stopped model (best_iteration {best} of {} iterations) is not \
updatable: slice it to its best iterations first (`slice(..{}, 1)`)",
model.num_boost_rounds(),
best + 1
),
));
}
let cache = match online.mode {
OnlineMode::Exact => None,
OnlineMode::Approximate { tolerance } => Some(with_thread_pool(params, || {
Cache::build(&model, params, data, tolerance)
})?),
};
Ok(OnlineModel {
params: params.clone(),
online,
data: data.clone(),
model,
cache,
})
}
pub fn update(
&mut self,
additions: Option<&DMatrix>,
deletions: &[usize],
) -> Result<UpdateReport> {
self.update_with(additions, deletions, |_| ControlFlow::Continue(()))
}
pub fn update_with(
&mut self,
additions: Option<&DMatrix>,
deletions: &[usize],
on_round: impl FnMut(RoundEval<'_>) -> ControlFlow<()> + Send,
) -> Result<UpdateReport> {
self.update_with_commit(additions, deletions, on_round, || ControlFlow::Continue(()))
}
pub fn update_with_commit(
&mut self,
additions: Option<&DMatrix>,
deletions: &[usize],
mut on_round: impl FnMut(RoundEval<'_>) -> ControlFlow<()> + Send,
commit: impl FnOnce() -> ControlFlow<()>,
) -> Result<UpdateReport> {
let deleted = self.check_change(additions, deletions)?;
let updated = compose(&self.data, &deleted, additions)?;
validate_training_data(&self.params, &updated)?;
initial_intercepts(
&self.params,
self.params.loss(1)?.as_ref(),
&updated.info(),
1,
)?;
let rounds = self.model.num_boost_rounds();
let Some(cache) = self.cache.as_ref() else {
let mut stopped = false;
let model = Trainer::new(&self.params, &updated, rounds)
.on_round(|round| {
let flow = on_round(round);
stopped |= flow.is_break();
flow
})
.train()?
.model;
if stopped || model.num_boost_rounds() != rounds || commit().is_break() {
return Err(interrupted());
}
let report = UpdateReport {
nodes_kept: kept_nodes(&self.model, &model),
subtrees_regrown: rounds,
rows_refreshed: updated.n_rows(),
};
self.model = model;
self.data = updated;
return Ok(report);
};
let mut cache = cache.clone();
let run = Incremental {
params: &self.params,
tolerance: cache.tolerance,
old: &self.data,
new: &updated,
deleted: &deleted,
model: &self.model,
};
let outcome = with_thread_pool(&self.params, || run.run(&mut cache, &mut on_round))
.and_then(|done| match commit() {
ControlFlow::Continue(()) => Ok(done),
ControlFlow::Break(()) => Err(interrupted()),
});
let (trees, report) = outcome?;
let model = self.model.with_trees(trees);
validate_trained_model(&model)?;
self.model = model;
self.data = updated;
self.cache = Some(cache);
Ok(report)
}
pub fn model(&self) -> &BoostedModel {
&self.model
}
pub fn data(&self) -> &DMatrix {
&self.data
}
pub fn params(&self) -> &TrainingParams {
&self.params
}
pub fn online_params(&self) -> OnlineParams {
self.online
}
pub fn into_model(self) -> BoostedModel {
self.model
}
fn check_change(&self, additions: Option<&DMatrix>, deletions: &[usize]) -> Result<Vec<bool>> {
let n = self.data.n_rows();
let mut deleted = vec![false; n];
for &row in deletions {
if row >= n {
return Err(HessboostError::invalid_param(
"deletions",
format!("row {row} is out of range for {n} rows"),
));
}
if std::mem::replace(&mut deleted[row], true) {
return Err(HessboostError::invalid_param(
"deletions",
format!("row {row} is deleted twice"),
));
}
}
let added = additions.map_or(0, DMatrix::n_rows);
if deletions.len() == n && added == 0 {
return Err(HessboostError::invalid_param(
"deletions",
"an update must leave at least one row",
));
}
if let Some(a) = additions {
check_data(a, "additions")?;
if a.n_cols() != self.data.n_cols() || a.feature_types() != self.data.feature_types() {
return Err(HessboostError::invalid_data(
"additions",
"added rows need the training data's columns and feature types",
));
}
if let Some(cache) = &self.cache {
check_within_cuts(&cache.cuts, a)?;
}
}
Ok(deleted)
}
}
fn check_within_cuts(cuts: &HistCuts, additions: &DMatrix) -> Result<()> {
let mut outside = None;
for row in 0..additions.n_rows() {
additions.for_row_entry(row, |c, v| {
let (start, end) = cuts.feature_bins(c as usize);
if outside.is_none() && (end == start || v >= cuts.cut_value(end - 1)) {
outside = Some((row, c, v));
}
});
if let Some((row, c, v)) = outside {
return Err(HessboostError::invalid_data(
"additions",
format!(
"added row {row} has feature {c} = {v}, beyond the training data's bins, \
which the approximate mode keeps fixed; use the exact mode (`OnlineParams::exact`, \
Python `mode=Exact()`), \
or rebuild the state on data covering it with `OnlineModel::from_model`"
),
));
}
}
Ok(())
}
fn interrupted() -> HessboostError {
HessboostError::invalid_param("on_round", "the update was interrupted; nothing changed")
}
fn kept_nodes(before: &BoostedModel, after: &BoostedModel) -> usize {
before
.trees()
.iter()
.zip(after.trees())
.map(|(a, b)| {
a.nodes()
.iter()
.zip(b.nodes())
.filter(|(x, y)| {
x.left == y.left
&& x.split_feature == y.split_feature
&& x.split_cond == y.split_cond
&& x.default_left == y.default_left
})
.count()
})
.sum()
}
fn check_supported(params: &TrainingParams, data: &DMatrix, online: OnlineParams) -> Result<()> {
params.validate()?;
let refuse = |name: &'static str, why: &str| {
Err(HessboostError::invalid_param(
name,
format!("in-place updates need {why}"),
))
};
if params.booster != BoosterKind::GbTree {
return refuse("booster", "booster = gbtree");
}
if !matches!(params.tree_method, TreeMethod::Hist | TreeMethod::Auto) {
return refuse("tree_method", "tree_method = hist");
}
if params.grow_policy != GrowPolicy::DepthWise
|| params.max_leaves.is_some()
|| params.max_depth.is_none()
{
return refuse(
"grow_policy",
"depth-wise growth with a max_depth and no max_leaves (a node's subtree must \
depend on its rows only)",
);
}
if params.subsample != 1.0
|| params.sampling_method != SamplingMethod::Uniform
|| params.colsample_bytree != 1.0
|| params.colsample_bylevel != 1.0
|| params.colsample_bynode != 1.0
{
return refuse(
"subsample",
"no row or column sampling (a retrain would draw different samples)",
);
}
if params.num_parallel_tree != 1 {
return refuse("num_parallel_tree", "num_parallel_tree = 1");
}
if params.process_type != ProcessType::Default {
return refuse("process_type", "process_type = default");
}
if !params.monotone_constraints.is_empty() || !params.interaction_constraints.is_empty() {
return refuse(
"monotone_constraints",
"no monotone or interaction constraints",
);
}
if params.extra_trees.is_some()
|| params.path_smooth != 0.0
|| params.linear_tree.is_some()
|| params.quantized.is_some()
|| params.toad_penalty_feature != 0.0
|| params.toad_penalty_threshold != 0.0
{
return refuse(
"extra_trees",
"no extra_trees, path_smooth, linear_tree, quantized gradients, or reuse penalties",
);
}
if params.device != Device::Cpu {
return refuse("device", "device = cpu");
}
let reference = TrainingParams {
booster: params.booster,
nthread: params.nthread,
seed: params.seed,
device: params.device,
objective: params.objective.clone(),
base_score: params.base_score,
eval_metric: params.eval_metric.clone(),
eta: params.eta,
gamma: params.gamma,
max_depth: params.max_depth,
min_child_weight: params.min_child_weight,
max_delta_step: params.max_delta_step,
lambda: params.lambda,
alpha: params.alpha,
tree_method: params.tree_method,
max_bin: params.max_bin,
multi_strategy: params.multi_strategy,
..TrainingParams::default()
};
params.refuse_changes_from(
&reference,
"params",
"in-place updates support only the settings they are proven sound for",
)?;
let per_row_newton = match ¶ms.objective {
Objective::SquaredError(_)
| Objective::SquaredLogError
| Objective::PseudoHuber(_)
| Objective::Expectile(_)
| Objective::RegLogistic(_)
| Objective::BinaryLogistic(_)
| Objective::BinaryLogitRaw(_)
| Objective::BinaryHinge
| Objective::Softmax(_)
| Objective::Softprob(_)
| Objective::Poisson
| Objective::Gamma(_)
| Objective::Tweedie(_)
| Objective::Aft(_)
| Objective::Dist(_) => true,
Objective::AbsoluteError
| Objective::Quantile(_)
| Objective::RankPairwise(_)
| Objective::RankNdcg(_)
| Objective::RankMap(_)
| Objective::RankXendcg
| Objective::Cox
| Objective::Custom(_) => false,
};
if !per_row_newton {
return refuse(
"objective",
"a built-in objective with per-row gradients and Newton-step leaves (not \
ranking, survival:cox, reg:absoluteerror, reg:quantileerror, or a custom loss)",
);
}
let objective = params.loss(1)?;
if objective.n_outputs() != 1 {
return refuse("objective", "a single-output objective");
}
check_data(data, "data")?;
if online.is_approximate() && data.feature_types().contains(&FeatureType::Categorical) {
return refuse(
"data",
"numerical features in the approximate mode (the exact mode accepts categorical ones)",
);
}
validate_training_data(params, data)
}
fn check_data(data: &DMatrix, name: &'static str) -> Result<()> {
if data.labels().is_none() || data.n_targets() != 1 {
return Err(HessboostError::invalid_data(
name,
"in-place updates need one label per row",
));
}
if data.weights().is_some()
|| data.base_margin().is_some()
|| data.group().is_some()
|| data.label_lower_bound().is_some()
|| data.feature_weights().is_some()
{
return Err(HessboostError::invalid_data(
name,
"in-place updates do not support weights, base margins, groups, label bounds, or \
feature weights",
));
}
Ok(())
}
fn compose(data: &DMatrix, deleted: &[bool], additions: Option<&DMatrix>) -> Result<DMatrix> {
let p = data.n_cols();
let rows = deleted
.iter()
.enumerate()
.filter(|(_, d)| !**d)
.map(|(r, _)| (data, r))
.chain(
additions
.into_iter()
.flat_map(|a| (0..a.n_rows()).map(move |r| (a, r))),
);
let mut labels = Vec::new();
let mut out = if data.dense_values().is_some() {
let mut values = Vec::new();
for (m, row) in rows {
let start = values.len();
values.resize(start + p, f32::NAN);
m.for_row_entry(row, |c, v| values[start + c as usize] = v);
labels.push(m.labels().map_or(0.0, |l| l[row]));
}
let n = labels.len();
DMatrix::from_dense(&values, n, p)?
} else {
let mut indptr = vec![0usize];
let (mut indices, mut values) = (Vec::new(), Vec::new());
for (m, row) in rows {
m.for_row_entry(row, |c, v| {
indices.push(c);
values.push(v);
});
indptr.push(values.len());
labels.push(m.labels().map_or(0.0, |l| l[row]));
}
DMatrix::from_csr(indptr, indices, values, p)?
};
if data.feature_types().contains(&FeatureType::Categorical) {
out = out.with_feature_types(data.feature_types())?;
}
out.with_labels(&labels)
}
fn splits_below(tree: &RegTree, max_depth: Option<NonZeroUsize>) -> bool {
let Some(limit) = max_depth else {
return false;
};
let mut stack = vec![(0usize, 0usize)];
while let Some((id, depth)) = stack.pop() {
let node = tree.node(id);
if node.is_leaf() {
continue;
}
if depth >= limit.get() {
return true;
}
stack.push((node.left as usize, depth + 1));
stack.push((node.right as usize, depth + 1));
}
false
}