hessboost/training/online/mod.rs
1//! In-place data addition and deletion for trained GBDT models
2//! (incremental and decremental learning, machine unlearning; opt-in).
3//!
4//! An [`OnlineModel`] keeps a model together with its training data and,
5//! when approximate updates are enabled, the node statistics that let it
6//! [`update`](OnlineModel::update) the model after rows are added to or
7//! deleted from that data without retraining from scratch. It follows Lin,
8//! Chung, Lao and Zhao, *Online Gradient Boosting Decision Tree: In-Place
9//! Updates for Efficient Adding/Deleting Data* (arXiv 2502.01634; reference
10//! code <https://github.com/huawei-lin/InplaceOnlineGBDT>), which extends
11//! their decremental method (*Machine Unlearning in Gradient Boosting
12//! Decision Trees*, KDD 2023) to additions.
13//!
14//! # Method
15//!
16//! The iterations are walked in order. In each tree the nodes are visited
17//! from the root down; the rows that changed (added, deleted, or whose
18//! gradient is refreshed, below) update the node's cached per-bin gradient
19//! histogram, and the node's current split is ranked among every candidate
20//! of the updated histogram with the histogram builder's own scoring. With
21//! [`OnlineParams::approximate`]'s tolerance `σ`, a split ranked within the top
22//! `max(1, ⌊σ · candidates⌋)` is kept (the paper's split robustness
23//! tolerance); otherwise the subtree under the node is regrown with the
24//! histogram builder on the rows now reaching it. Leaves whose statistics
25//! changed get their weights recomputed.
26//!
27//! Gradients are refreshed lazily (the paper's adaptive lazy update): rows
28//! that were added, or that reached a regrown subtree in an earlier tree,
29//! get their gradients recomputed from the updated model; every other row
30//! keeps the gradient cached for it, since its margin only drifts by the
31//! (small) changes of the leaf weights it reaches. The histogram's bin
32//! boundaries and the model's intercept stay those of the original
33//! training, so an added value at or above a feature's top bin boundary
34//! (beyond the training data's range) is refused in this mode.
35//!
36//! # Exactness
37//!
38//! [`OnlineParams::exact`] is the exact mode: every node's split must still be the
39//! best one on recomputed statistics, bin boundaries and intercept, which is
40//! by definition what retraining produces, so the exact mode defers every
41//! tree to the builder: the result equals [`train`](super::train) on
42//! [`OnlineModel::data`] bit for bit, at about the cost of retraining (a
43//! reference to measure the approximate mode against, and the answer when
44//! unlearning must be exact). [`OnlineParams::approximate`] is approximate: the kept
45//! splits, the fixed bins and intercept, and the lazily refreshed gradients
46//! make the model differ from a retrain, by a gap that grows with the
47//! fraction of rows changed and with `σ`, in exchange for touching only the
48//! changed rows and the regrown subtrees.
49//!
50//! # Supported configurations
51//!
52//! Updating is sound only where retraining the same parameters on the
53//! updated data is deterministic in the data alone and every node's
54//! subtree depends on the node's rows only: `gbtree` with `tree_method =
55//! hist` (or `auto`), depth-wise growth with a positive `max_depth` and no
56//! `max_leaves`, one output, `num_parallel_tree = 1`, no row or column
57//! sampling (a retrain draws its samples sequentially over the rows, so
58//! they would change with any added or deleted row), no monotone or
59//! interaction constraints, none of the opt-in split options
60//! (`extra_trees`, `path_smooth`, `linear_tree`, quantized gradients, reuse
61//! penalties), CPU, and a built-in objective whose gradients are per row
62//! and whose leaves are plain Newton steps (not ranking, `survival:cox`,
63//! `reg:absoluteerror`, `reg:quantileerror`, or a custom loss). The data
64//! may not carry weights, base margins, groups, label bounds, or feature
65//! weights, and the approximate mode needs numerical features. Everything
66//! else is refused.
67//!
68//! # Memory
69//!
70//! The approximate mode caches, per tree, one gradient pair per row and one
71//! histogram (`total_bins` pairs of `f64`) per internal node: `trees × (8 ·
72//! rows + 16 · total_bins · internal nodes)` bytes. An update works on a
73//! copy of it, swapped in on success, so an abandoned update restores the
74//! exact state it started from; the copy briefly doubles that memory.
75//!
76//! # Accuracy and speed
77//!
78//! On 20,000-row synthetic tasks (Friedman #1 regression, its thresholded
79//! classification, and a 30-feature variant; 100 trees of depth 6, `σ =
80//! 0.1`), updates of 0.1% and 1% of the rows ran 1.3–4.8x faster than
81//! retraining, with test losses within about 0.5% of the retrained
82//! model's; at 5% they were slower than retraining (regrown subtrees
83//! refresh most rows). Deleting rows raises the model's loss on them, but
84//! by a fraction of what retraining does (e.g. regression RMSE on 200
85//! deleted rows 0.877 → 0.903, retrained 1.080): the approximate mode
86//! forgets only partially, so unlearning that must be complete needs the
87//! exact mode.
88//!
89//! # Interruption
90//!
91//! [`OnlineModel::train_with`] and [`OnlineModel::update_with`] call a
92//! per-iteration hook as [`Trainer::on_round`] does. Breaking stops
93//! training after the iteration (as it stops [`Trainer`]) and abandons an
94//! update, leaving the model, data and state unchanged.
95//! [`OnlineModel::update_with_commit`] also asks for a last confirmation
96//! once the update is computed, before it is applied.
97//!
98//! # Example
99//!
100//! ```
101//! use hessboost::prelude::*;
102//! use hessboost::training::online::{OnlineModel, OnlineParams};
103//!
104//! # fn main() -> Result<()> {
105//! let x: Vec<f32> = (0..200).map(|i| (i % 50) as f32 / 50.0).collect();
106//! let y: Vec<f32> = x.iter().map(|v| 3.0 * v).collect();
107//! let data = DMatrix::from_dense(&x, 200, 1)?.with_labels(&y)?;
108//! let params = TrainingParams::builder()
109//! .tree_method(TreeMethod::Hist)
110//! .max_depth(3)
111//! .build()?;
112//!
113//! let mut online = OnlineModel::train(¶ms, &data, 20, OnlineParams::default())?;
114//! // Forget the first ten rows, learn two new ones.
115//! let new = DMatrix::from_dense(&[0.5, 0.25], 2, 1)?.with_labels(&[1.5, 0.75])?;
116//! let report = online.update(Some(&new), &(0..10).collect::<Vec<_>>())?;
117//! assert_eq!(online.data().n_rows(), 192);
118//! assert!(report.nodes_kept > 0);
119//! # Ok(())
120//! # }
121//! ```
122
123mod cache;
124mod update;
125
126use std::num::NonZeroUsize;
127use std::ops::ControlFlow;
128
129use self::cache::Cache;
130use self::update::Incremental;
131use super::api::{RoundEval, Trainer};
132use super::eval::configured_metrics;
133use super::train::{initial_intercepts, with_thread_pool};
134use super::validate::{validate_trained_model, validate_training_data};
135use crate::config::{
136 BoosterKind, Device, GrowPolicy, ProcessType, SamplingMethod, TrainingParams, TreeMethod,
137};
138use crate::data::quantile::HistCuts;
139use crate::data::{DMatrix, FeatureType};
140use crate::error::{HessboostError, Result};
141use crate::model::BoostedModel;
142use crate::objective::Objective;
143use crate::tree::RegTree;
144
145/// Settings of [`OnlineModel`]: its [`OnlineMode`], the exact mode
146/// ([`Self::exact`], see the [module docs](self#exactness)) or the
147/// approximate one with a split robustness tolerance
148/// ([`Self::approximate`]). The default is approximate at `0.1`, the
149/// paper's recommendation.
150///
151/// ```
152/// use hessboost::training::online::{OnlineMode, OnlineParams};
153///
154/// # fn main() -> hessboost::error::Result<()> {
155/// let OnlineMode::Approximate { tolerance, .. } = OnlineParams::default().mode() else {
156/// unreachable!("the default is approximate");
157/// };
158/// assert_eq!(tolerance, 0.1);
159/// assert_eq!(OnlineParams::exact().mode(), OnlineMode::Exact);
160/// // The exact mode is `exact()`, not a tolerance of 0.
161/// assert!(OnlineParams::approximate(0.0).is_err());
162/// # Ok(())
163/// # }
164/// ```
165#[derive(Debug, Clone, Copy, PartialEq)]
166pub struct OnlineParams {
167 mode: OnlineMode,
168}
169
170/// How an [`OnlineModel`] updates (see the [module docs](self#exactness)),
171/// read from [`OnlineParams::mode`]; built by [`OnlineParams::exact`] and
172/// [`OnlineParams::approximate`].
173#[derive(Debug, Clone, Copy, PartialEq)]
174#[non_exhaustive]
175pub enum OnlineMode {
176 /// Every update reproduces retraining bit for bit, at about its cost.
177 Exact,
178 /// Splits ranked within the tolerance are kept, the rest regrown, on the
179 /// original training's bins and intercept, with lazily refreshed
180 /// gradients.
181 #[non_exhaustive]
182 Approximate {
183 /// The split robustness tolerance `σ`, in `(0, 1]`.
184 tolerance: f64,
185 },
186}
187
188impl Default for OnlineParams {
189 fn default() -> Self {
190 OnlineParams {
191 mode: OnlineMode::Approximate { tolerance: 0.1 },
192 }
193 }
194}
195
196impl OnlineParams {
197 /// The exact mode: every update reproduces retraining bit for bit.
198 pub fn exact() -> Self {
199 OnlineParams {
200 mode: OnlineMode::Exact,
201 }
202 }
203
204 /// The approximate mode with split robustness tolerance `σ` in
205 /// `(0, 1]`: a node keeps its split while it ranks within the top
206 /// `max(1, ⌊σ · candidates⌋)` candidates; `1` regrows only nodes whose
207 /// split stopped being a valid candidate (a child below
208 /// `min_child_weight`, a gain below `gamma`).
209 ///
210 /// # Errors
211 ///
212 /// `tolerance` outside `(0, 1]` (for the exact mode use
213 /// [`Self::exact`]), named `tolerance`.
214 pub fn approximate(tolerance: f64) -> Result<Self> {
215 crate::check::fraction("tolerance", tolerance)?;
216 Ok(OnlineParams {
217 mode: OnlineMode::Approximate { tolerance },
218 })
219 }
220
221 /// The update mode.
222 pub fn mode(&self) -> OnlineMode {
223 self.mode
224 }
225
226 /// Whether updates are approximate (and keep the approximate mode's
227 /// state).
228 fn is_approximate(self) -> bool {
229 matches!(self.mode, OnlineMode::Approximate { .. })
230 }
231}
232
233/// What an [`OnlineModel::update`] did.
234#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
235#[non_exhaustive]
236pub struct UpdateReport {
237 /// Nodes whose split (or leaf) was kept, over all trees.
238 pub nodes_kept: usize,
239 /// Subtrees regrown by the builder (the exact mode regrows every tree).
240 pub subtrees_regrown: usize,
241 /// Rows whose gradients were recomputed in at least one tree.
242 pub rows_refreshed: usize,
243}
244
245/// A trained model with its training data, updatable in place. See the
246/// [module docs](self).
247#[derive(Debug, Clone)]
248pub struct OnlineModel {
249 params: TrainingParams,
250 online: OnlineParams,
251 data: DMatrix,
252 model: BoostedModel,
253 cache: Option<Cache>,
254}
255
256impl OnlineModel {
257 /// Train `num_boost_round` iterations on `data` (as
258 /// [`train`](super::train) does) and keep what updates need.
259 ///
260 /// # Errors
261 ///
262 /// The refusals of the [module docs](self#supported-configurations) and
263 /// the errors of training.
264 pub fn train(
265 params: &TrainingParams,
266 data: &DMatrix,
267 num_boost_round: usize,
268 online: OnlineParams,
269 ) -> Result<Self> {
270 Self::train_with(params, data, num_boost_round, online, |_| {
271 ControlFlow::Continue(())
272 })
273 }
274
275 /// [`Self::train`] calling `on_round` after every iteration, as
276 /// [`Trainer::on_round`] does: [`ControlFlow::Break`] stops training
277 /// after the iteration, and the online model keeps the iterations so
278 /// far (updates then retrain or update that many).
279 ///
280 /// # Errors
281 ///
282 /// Those of [`Self::train`].
283 pub fn train_with(
284 params: &TrainingParams,
285 data: &DMatrix,
286 num_boost_round: usize,
287 online: OnlineParams,
288 on_round: impl FnMut(RoundEval<'_>) -> ControlFlow<()> + Send,
289 ) -> Result<Self> {
290 check_supported(params, data, online)?;
291 let model = Trainer::new(params, data, num_boost_round)
292 .on_round(on_round)
293 .train()?
294 .model;
295 Self::from_model(model, params, data, online)
296 }
297
298 /// Resume from a `model` trained with `params` on `data` (for example one
299 /// loaded from a file), rebuilding the update state by replaying its
300 /// trees over `data`.
301 ///
302 /// # Errors
303 ///
304 /// The refusals of [`Self::train`],
305 /// [`HessboostError::IncompatibleModel`] (`model`) for a model that
306 /// `params` could not have trained (another objective or
307 /// `max_delta_step`, several outputs, weighted trees, categorical trees
308 /// in the approximate mode, linear leaves, a `gblinear`, `boulevard`, or
309 /// `ebm` booster, for example from an imported LightGBM `linear_tree`
310 /// model, model shrinkage, a different feature count, or, in the
311 /// approximate mode, a first tree whose leaves are not these parameters'
312 /// Newton steps on `data`: another `eta` or `lambda`, other data, or an
313 /// imported LightGBM model, whose first tree carries the label average)
314 /// and for an early-stopped model (`best_iteration` set: slice it to its
315 /// best iterations first), and [`HessboostError::InvalidParameter`] for
316 /// `eval_metric`s training would refuse.
317 pub fn from_model(
318 model: BoostedModel,
319 params: &TrainingParams,
320 data: &DMatrix,
321 online: OnlineParams,
322 ) -> Result<Self> {
323 check_supported(params, data, online)?;
324 // Training refuses metrics the objective cannot score before any
325 // round; a resumed model gets the same check.
326 configured_metrics(params, params.loss(1)?.as_ref())?;
327 let categorical = model
328 .trees()
329 .iter()
330 .any(|t| t.nodes().iter().any(|n| n.is_categorical));
331 if model.objective().built_in() != Some(¶ms.objective)
332 || model.n_outputs() != 1
333 || model.num_parallel_tree() != 1
334 || model.n_features() != data.n_cols()
335 // The approximate mode replays numeric splits only; the exact
336 // mode retrains, so categorical trees are fine there.
337 || (categorical && online.is_approximate())
338 // The objective's `max_delta_step` shapes every leaf: a model
339 // trained with another one is not a model of these parameters.
340 || model.max_delta_step() != params.effective_max_delta_step()
341 || model.linear().is_some()
342 || model.boulevard().is_some()
343 || model.ebm().is_some()
344 || model.trees().iter().any(|t| t.linear_leaves().is_some())
345 || model
346 .trees()
347 .iter()
348 .any(|t| splits_below(t, params.max_depth))
349 || (0..model.num_trees()).any(|t| model.tree_weight(t) != 1.0)
350 // Model shrinkage rescales every earlier tree each round, so
351 // an update of one node's subtree would not reproduce a retrain.
352 || model.shrinkage().is_some()
353 {
354 return Err(HessboostError::incompatible_model(
355 "model",
356 "not a single-output, unweighted, unshrunk, numeric gbtree model (not \
357 Boulevard or EBM) with constant leaves of these parameters and data",
358 ));
359 }
360 if let Some(best) = model.best_iteration() {
361 // Early stopping keeps every trained iteration but predicts with
362 // the first `best + 1`; an update updates (and predicts with)
363 // them all, and a model without `best_iteration`.
364 return Err(HessboostError::incompatible_model(
365 "model",
366 format!(
367 "an early-stopped model (best_iteration {best} of {} iterations) is not \
368 updatable: slice it to its best iterations first (`slice(..{}, 1)`)",
369 model.num_boost_rounds(),
370 best + 1
371 ),
372 ));
373 }
374 let cache = match online.mode {
375 OnlineMode::Exact => None,
376 OnlineMode::Approximate { tolerance } => Some(with_thread_pool(params, || {
377 Cache::build(&model, params, data, tolerance)
378 })?),
379 };
380 Ok(OnlineModel {
381 params: params.clone(),
382 online,
383 data: data.clone(),
384 model,
385 cache,
386 })
387 }
388
389 /// Add the rows of `additions` (features and labels, like the training
390 /// data) and delete the training rows `deletions` (indices into
391 /// [`Self::data`]), then update the model. The new data is the kept rows
392 /// in their order followed by the added ones.
393 ///
394 /// # Errors
395 ///
396 /// [`HessboostError::InvalidParameter`] (`deletions`) for out-of-range
397 /// or repeated deletions or deleting every row;
398 /// [`HessboostError::InvalidData`] (`additions`) for additions without
399 /// labels or with metadata, or of another shape, and (approximate mode)
400 /// added values beyond the training data's bins; whatever training refuses on the
401 /// updated data (such as labels outside the objective's domain), checked
402 /// before anything changes; [`HessboostError::ModelFormat`] for an
403 /// update whose arithmetic overflows `f32`, as training refuses such a
404 /// model; the errors of training.
405 pub fn update(
406 &mut self,
407 additions: Option<&DMatrix>,
408 deletions: &[usize],
409 ) -> Result<UpdateReport> {
410 self.update_with(additions, deletions, |_| ControlFlow::Continue(()))
411 }
412
413 /// [`Self::update`] calling `on_round` after every updated iteration, as
414 /// [`Trainer::on_round`] does during training. [`ControlFlow::Break`]
415 /// abandons the update: the model and data stay as they were and
416 /// [`HessboostError::InvalidParameter`] (`on_round`) is returned.
417 ///
418 /// # Errors
419 ///
420 /// Those of [`Self::update`], and the interruption.
421 pub fn update_with(
422 &mut self,
423 additions: Option<&DMatrix>,
424 deletions: &[usize],
425 on_round: impl FnMut(RoundEval<'_>) -> ControlFlow<()> + Send,
426 ) -> Result<UpdateReport> {
427 self.update_with_commit(additions, deletions, on_round, || ControlFlow::Continue(()))
428 }
429
430 /// [`Self::update_with`] asking `commit` once the update is computed,
431 /// just before it is applied: [`ControlFlow::Break`] abandons it as a
432 /// break from `on_round` does (for a caller whose interruption can
433 /// arrive after the last iteration's hook, such as a signal).
434 ///
435 /// # Errors
436 ///
437 /// Those of [`Self::update_with`].
438 pub fn update_with_commit(
439 &mut self,
440 additions: Option<&DMatrix>,
441 deletions: &[usize],
442 mut on_round: impl FnMut(RoundEval<'_>) -> ControlFlow<()> + Send,
443 commit: impl FnOnce() -> ControlFlow<()>,
444 ) -> Result<UpdateReport> {
445 let deleted = self.check_change(additions, deletions)?;
446 let updated = compose(&self.data, &deleted, additions)?;
447 // Refuse what retraining on `updated` refuses (e.g. labels outside
448 // the loss's domain) before any state changes: the approximate mode
449 // computes gradients without going through the trainer's checks.
450 validate_training_data(&self.params, &updated)?;
451 // Retraining estimates the intercept from the labels (unless
452 // `base_score` is set) and refuses a non-finite one; the approximate
453 // mode keeps the original intercept, so check it here.
454 initial_intercepts(
455 &self.params,
456 self.params.loss(1)?.as_ref(),
457 &updated.info(),
458 1,
459 )?;
460 let rounds = self.model.num_boost_rounds();
461 // The approximate mode updates a copy of the cache, swapped in on
462 // success, so an abandoned update leaves exactly the state it started
463 // from (a rebuild would recompute the bins and gradients earlier
464 // updates keep fixed).
465 let Some(cache) = self.cache.as_ref() else {
466 let mut stopped = false;
467 let model = Trainer::new(&self.params, &updated, rounds)
468 .on_round(|round| {
469 let flow = on_round(round);
470 stopped |= flow.is_break();
471 flow
472 })
473 .train()?
474 .model;
475 if stopped || model.num_boost_rounds() != rounds || commit().is_break() {
476 return Err(interrupted());
477 }
478 let report = UpdateReport {
479 nodes_kept: kept_nodes(&self.model, &model),
480 subtrees_regrown: rounds,
481 rows_refreshed: updated.n_rows(),
482 };
483 self.model = model;
484 self.data = updated;
485 return Ok(report);
486 };
487 let mut cache = cache.clone();
488 let run = Incremental {
489 params: &self.params,
490 tolerance: cache.tolerance,
491 old: &self.data,
492 new: &updated,
493 deleted: &deleted,
494 model: &self.model,
495 };
496 // On `nthread` threads, as training runs.
497 let outcome = with_thread_pool(&self.params, || run.run(&mut cache, &mut on_round))
498 .and_then(|done| match commit() {
499 ControlFlow::Continue(()) => Ok(done),
500 ControlFlow::Break(()) => Err(interrupted()),
501 });
502 let (trees, report) = outcome?;
503 // Arithmetic that overflows `f32` (extreme labels or margins) leaves
504 // non-finite leaves the model formats refuse: refuse the update, as
505 // training refuses such a model, before anything changes.
506 let model = self.model.with_trees(trees);
507 validate_trained_model(&model)?;
508 self.model = model;
509 self.data = updated;
510 self.cache = Some(cache);
511 Ok(report)
512 }
513
514 /// The current model.
515 pub fn model(&self) -> &BoostedModel {
516 &self.model
517 }
518
519 /// The current training data: the original rows minus deletions plus
520 /// additions, in update order.
521 pub fn data(&self) -> &DMatrix {
522 &self.data
523 }
524
525 /// The training parameters.
526 pub fn params(&self) -> &TrainingParams {
527 &self.params
528 }
529
530 /// The update settings.
531 pub fn online_params(&self) -> OnlineParams {
532 self.online
533 }
534
535 /// The model, dropping the data and update state.
536 pub fn into_model(self) -> BoostedModel {
537 self.model
538 }
539
540 /// Validate a change; returns the deletion mask over [`Self::data`].
541 fn check_change(&self, additions: Option<&DMatrix>, deletions: &[usize]) -> Result<Vec<bool>> {
542 let n = self.data.n_rows();
543 let mut deleted = vec![false; n];
544 for &row in deletions {
545 if row >= n {
546 return Err(HessboostError::invalid_param(
547 "deletions",
548 format!("row {row} is out of range for {n} rows"),
549 ));
550 }
551 if std::mem::replace(&mut deleted[row], true) {
552 return Err(HessboostError::invalid_param(
553 "deletions",
554 format!("row {row} is deleted twice"),
555 ));
556 }
557 }
558 let added = additions.map_or(0, DMatrix::n_rows);
559 if deletions.len() == n && added == 0 {
560 return Err(HessboostError::invalid_param(
561 "deletions",
562 "an update must leave at least one row",
563 ));
564 }
565 if let Some(a) = additions {
566 check_data(a, "additions")?;
567 if a.n_cols() != self.data.n_cols() || a.feature_types() != self.data.feature_types() {
568 return Err(HessboostError::invalid_data(
569 "additions",
570 "added rows need the training data's columns and feature types",
571 ));
572 }
573 if let Some(cache) = &self.cache {
574 check_within_cuts(&cache.cuts, a)?;
575 }
576 }
577 Ok(deleted)
578 }
579}
580
581/// Refuse an added value at or above its feature's top cut. The approximate
582/// mode keeps the original training's bins, whose last bin ends at that cut:
583/// a value past it would count in the last bin of the histograms that rank
584/// splits while prediction routes it right of a split at the top cut, so
585/// the ranked and the actual children would differ.
586fn check_within_cuts(cuts: &HistCuts, additions: &DMatrix) -> Result<()> {
587 let mut outside = None;
588 for row in 0..additions.n_rows() {
589 additions.for_row_entry(row, |c, v| {
590 let (start, end) = cuts.feature_bins(c as usize);
591 if outside.is_none() && (end == start || v >= cuts.cut_value(end - 1)) {
592 outside = Some((row, c, v));
593 }
594 });
595 if let Some((row, c, v)) = outside {
596 return Err(HessboostError::invalid_data(
597 "additions",
598 format!(
599 "added row {row} has feature {c} = {v}, beyond the training data's bins, \
600 which the approximate mode keeps fixed; use the exact mode (`OnlineParams::exact`, \
601 Python `mode=Exact()`), \
602 or rebuild the state on data covering it with `OnlineModel::from_model`"
603 ),
604 ));
605 }
606 }
607 Ok(())
608}
609
610fn interrupted() -> HessboostError {
611 HessboostError::invalid_param("on_round", "the update was interrupted; nothing changed")
612}
613
614/// Nodes of `after` equal (split or leaf, same position) to `before`'s.
615fn kept_nodes(before: &BoostedModel, after: &BoostedModel) -> usize {
616 before
617 .trees()
618 .iter()
619 .zip(after.trees())
620 .map(|(a, b)| {
621 a.nodes()
622 .iter()
623 .zip(b.nodes())
624 .filter(|(x, y)| {
625 x.left == y.left
626 && x.split_feature == y.split_feature
627 && x.split_cond == y.split_cond
628 && x.default_left == y.default_left
629 })
630 .count()
631 })
632 .sum()
633}
634
635/// The refusals of the module docs.
636fn check_supported(params: &TrainingParams, data: &DMatrix, online: OnlineParams) -> Result<()> {
637 params.validate()?;
638 let refuse = |name: &'static str, why: &str| {
639 Err(HessboostError::invalid_param(
640 name,
641 format!("in-place updates need {why}"),
642 ))
643 };
644 if params.booster != BoosterKind::GbTree {
645 return refuse("booster", "booster = gbtree");
646 }
647 if !matches!(params.tree_method, TreeMethod::Hist | TreeMethod::Auto) {
648 return refuse("tree_method", "tree_method = hist");
649 }
650 if params.grow_policy != GrowPolicy::DepthWise
651 || params.max_leaves.is_some()
652 || params.max_depth.is_none()
653 {
654 return refuse(
655 "grow_policy",
656 "depth-wise growth with a max_depth and no max_leaves (a node's subtree must \
657 depend on its rows only)",
658 );
659 }
660 if params.subsample != 1.0
661 || params.sampling_method != SamplingMethod::Uniform
662 || params.colsample_bytree != 1.0
663 || params.colsample_bylevel != 1.0
664 || params.colsample_bynode != 1.0
665 {
666 return refuse(
667 "subsample",
668 "no row or column sampling (a retrain would draw different samples)",
669 );
670 }
671 if params.num_parallel_tree != 1 {
672 return refuse("num_parallel_tree", "num_parallel_tree = 1");
673 }
674 if params.process_type != ProcessType::Default {
675 return refuse("process_type", "process_type = default");
676 }
677 if !params.monotone_constraints.is_empty() || !params.interaction_constraints.is_empty() {
678 return refuse(
679 "monotone_constraints",
680 "no monotone or interaction constraints",
681 );
682 }
683 if params.extra_trees.is_some()
684 || params.path_smooth != 0.0
685 || params.linear_tree.is_some()
686 || params.quantized.is_some()
687 || params.toad_penalty_feature != 0.0
688 || params.toad_penalty_threshold != 0.0
689 {
690 return refuse(
691 "extra_trees",
692 "no extra_trees, path_smooth, linear_tree, quantized gradients, or reuse penalties",
693 );
694 }
695 if params.device != Device::Cpu {
696 return refuse("device", "device = cpu");
697 }
698 // Every other setting, including any added later, stays at its default:
699 // a new training option is refused until it is shown sound here.
700 let reference = TrainingParams {
701 booster: params.booster,
702 nthread: params.nthread,
703 seed: params.seed,
704 device: params.device,
705 objective: params.objective.clone(),
706 base_score: params.base_score,
707 eval_metric: params.eval_metric.clone(),
708 eta: params.eta,
709 gamma: params.gamma,
710 max_depth: params.max_depth,
711 min_child_weight: params.min_child_weight,
712 max_delta_step: params.max_delta_step,
713 lambda: params.lambda,
714 alpha: params.alpha,
715 tree_method: params.tree_method,
716 max_bin: params.max_bin,
717 multi_strategy: params.multi_strategy,
718 ..TrainingParams::default()
719 };
720 params.refuse_changes_from(
721 &reference,
722 "params",
723 "in-place updates support only the settings they are proven sound for",
724 )?;
725 // Exhaustive, so a new objective has to be classified here.
726 let per_row_newton = match ¶ms.objective {
727 Objective::SquaredError(_)
728 | Objective::SquaredLogError
729 | Objective::PseudoHuber(_)
730 | Objective::Expectile(_)
731 | Objective::RegLogistic(_)
732 | Objective::BinaryLogistic(_)
733 | Objective::BinaryLogitRaw(_)
734 | Objective::BinaryHinge
735 | Objective::Softmax(_)
736 | Objective::Softprob(_)
737 | Objective::Poisson
738 | Objective::Gamma(_)
739 | Objective::Tweedie(_)
740 | Objective::Aft(_)
741 | Objective::Dist(_) => true,
742 Objective::AbsoluteError
743 | Objective::Quantile(_)
744 | Objective::RankPairwise(_)
745 | Objective::RankNdcg(_)
746 | Objective::RankMap(_)
747 | Objective::RankXendcg
748 | Objective::Cox
749 | Objective::Custom(_) => false,
750 };
751 if !per_row_newton {
752 return refuse(
753 "objective",
754 "a built-in objective with per-row gradients and Newton-step leaves (not \
755 ranking, survival:cox, reg:absoluteerror, reg:quantileerror, or a custom loss)",
756 );
757 }
758 let objective = params.loss(1)?;
759 if objective.n_outputs() != 1 {
760 return refuse("objective", "a single-output objective");
761 }
762 check_data(data, "data")?;
763 if online.is_approximate() && data.feature_types().contains(&FeatureType::Categorical) {
764 return refuse(
765 "data",
766 "numerical features in the approximate mode (the exact mode accepts categorical ones)",
767 );
768 }
769 validate_training_data(params, data)
770}
771
772/// Labels and none of the metadata updates cannot honor.
773fn check_data(data: &DMatrix, name: &'static str) -> Result<()> {
774 if data.labels().is_none() || data.n_targets() != 1 {
775 return Err(HessboostError::invalid_data(
776 name,
777 "in-place updates need one label per row",
778 ));
779 }
780 if data.weights().is_some()
781 || data.base_margin().is_some()
782 || data.group().is_some()
783 || data.label_lower_bound().is_some()
784 || data.feature_weights().is_some()
785 {
786 return Err(HessboostError::invalid_data(
787 name,
788 "in-place updates do not support weights, base margins, groups, label bounds, or \
789 feature weights",
790 ));
791 }
792 Ok(())
793}
794
795/// The rows of `data` not in `deleted`, then those of `additions`: dense
796/// (missing as NaN) when `data` is dense, else compressed sparse rows.
797fn compose(data: &DMatrix, deleted: &[bool], additions: Option<&DMatrix>) -> Result<DMatrix> {
798 let p = data.n_cols();
799 let rows = deleted
800 .iter()
801 .enumerate()
802 .filter(|(_, d)| !**d)
803 .map(|(r, _)| (data, r))
804 .chain(
805 additions
806 .into_iter()
807 .flat_map(|a| (0..a.n_rows()).map(move |r| (a, r))),
808 );
809 let mut labels = Vec::new();
810 let mut out = if data.dense_values().is_some() {
811 let mut values = Vec::new();
812 for (m, row) in rows {
813 let start = values.len();
814 values.resize(start + p, f32::NAN);
815 m.for_row_entry(row, |c, v| values[start + c as usize] = v);
816 labels.push(m.labels().map_or(0.0, |l| l[row]));
817 }
818 let n = labels.len();
819 DMatrix::from_dense(&values, n, p)?
820 } else {
821 let mut indptr = vec![0usize];
822 let (mut indices, mut values) = (Vec::new(), Vec::new());
823 for (m, row) in rows {
824 m.for_row_entry(row, |c, v| {
825 indices.push(c);
826 values.push(v);
827 });
828 indptr.push(values.len());
829 labels.push(m.labels().map_or(0.0, |l| l[row]));
830 }
831 DMatrix::from_csr(indptr, indices, values, p)?
832 };
833 if data.feature_types().contains(&FeatureType::Categorical) {
834 out = out.with_feature_types(data.feature_types())?;
835 }
836 out.with_labels(&labels)
837}
838
839/// Whether `tree` splits a node at depth `max_depth` or deeper (`None`: no
840/// limit), which depth-wise training to that depth never does.
841fn splits_below(tree: &RegTree, max_depth: Option<NonZeroUsize>) -> bool {
842 let Some(limit) = max_depth else {
843 return false;
844 };
845 let mut stack = vec![(0usize, 0usize)];
846 while let Some((id, depth)) = stack.pop() {
847 let node = tree.node(id);
848 if node.is_leaf() {
849 continue;
850 }
851 if depth >= limit.get() {
852 return true;
853 }
854 stack.push((node.left as usize, depth + 1));
855 stack.push((node.right as usize, depth + 1));
856 }
857 false
858}