Skip to main content

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(&params, &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(&params.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 &params.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}