hessboost 0.2.0

Fast, deterministic gradient boosting (GBDT) in Rust: conformal intervals, explainable boosting machines, distributional boosting, tree-based diffusion, and XGBoost model interchange
Documentation
//! The refresh updater behind `process_type=update` (XGBoost's
//! `updater=refresh`, `src/tree/updater_refresh.cc`).
//!
//! Refreshing keeps a tree's split structure and recomputes its node
//! statistics from new gradients: every row is routed from the root to its
//! leaf, adding its gradient pair to each node on the path (`f64`
//! accumulation, like XGBoost's `GradStats`). Each node then gets
//!
//! * `sum_hess` = the accumulated Hessian (its cover),
//! * `split_gain` = `gain(left) + gain(right) - gain(node)` for internal nodes,
//! * with `refresh_leaf`, a leaf value `calc_weight(stats) * learning_rate`
//!   (the node's base weight, shrunk by the per-tree learning rate
//!   `eta / num_parallel_tree`); without it the leaf values stay as they are.
//!
//! No rows are subsampled. Unlike split finding, XGBoost's refresh applies
//! no `min_child_weight` floor: a node with any positive Hessian gets its
//! regularized weight (`CalcWeight` / `CalcGain` in `src/tree/param.h`).

use crate::config::{Refresh, TrainingParams};
use crate::data::DMatrix;
use crate::objective::GradPair;
use crate::tree::RegTree;
use crate::tree::gain::{GradStats, RegParams, calc_gain, calc_weight};
use rayon::prelude::*;

/// Rows per statistics block. Blocks are reduced in index order, so the sums
/// do not depend on the thread count.
const REFRESH_BLOCK_ROWS: usize = 4096;

/// Blocks accumulated concurrently per worker thread before they are reduced
/// into the running total: bounds the buffered per-block node statistics to
/// `threads × REFRESH_BATCH_PER_THREAD × n_nodes` regardless of the row count.
const REFRESH_BATCH_PER_THREAD: usize = 4;

/// Refresh `tree` in place from the per-row gradients `gpair` (one pair per
/// row of `data`) with `params`' regularization and `refresh_leaf`, shrinking
/// refreshed leaves by `learning_rate`. See the module docs for the
/// recomputed quantities.
pub(super) fn refresh_tree(
    tree: &mut RegTree,
    data: &DMatrix,
    gpair: &[GradPair],
    params: &TrainingParams,
    refresh: Refresh,
    learning_rate: f32,
) {
    let reg = &RegParams {
        min_child_weight: 0.0,
        ..RegParams::from_params(params)
    };
    let stats = node_stats(tree, data, gpair);
    for nid in 0..tree.num_nodes() {
        let node = *tree.node(nid);
        tree.set_sum_hess(nid, stats[nid].hess as f32);
        if node.is_leaf() {
            if refresh.refresh_leaf() {
                let base_weight = calc_weight(stats[nid], reg) as f32;
                tree.set_leaf_value(nid, base_weight * learning_rate);
            }
        } else {
            let gain = calc_gain(stats[node.left as usize], reg)
                + calc_gain(stats[node.right as usize], reg)
                - calc_gain(stats[nid], reg);
            tree.set_split_gain(nid, gain as f32);
        }
    }
}

/// Per-node gradient statistics of `tree` over every row of `data`.
fn node_stats(tree: &RegTree, data: &DMatrix, gpair: &[GradPair]) -> Vec<GradStats> {
    let batch_blocks = rayon::current_num_threads().max(1) * REFRESH_BATCH_PER_THREAD;
    node_stats_batched(tree, data, gpair, batch_blocks)
}

/// [`node_stats`] holding at most `batch_blocks` block accumulators at once.
/// Blocks are still reduced one by one in index order, so the result is the
/// same for every batch size (and thread count).
fn node_stats_batched(
    tree: &RegTree,
    data: &DMatrix,
    gpair: &[GradPair],
    batch_blocks: usize,
) -> Vec<GradStats> {
    let n_nodes = tree.num_nodes();
    let block_stats = |rows: std::ops::Range<usize>| {
        let mut stats = vec![GradStats::default(); n_nodes];
        for row in rows {
            let gp = GradStats::from_pair(gpair[row]);
            let mut nid = 0;
            stats[nid].add(gp);
            while !tree.node(nid).is_leaf() {
                let feature = tree.node(nid).split_feature as usize;
                nid = tree.child(nid, data.get(row, feature));
                stats[nid].add(gp);
            }
        }
        stats
    };
    let n = data.n_rows();
    let n_blocks = n.div_ceil(REFRESH_BLOCK_ROWS);
    let batch = batch_blocks.max(1);
    let mut total = vec![GradStats::default(); n_nodes];
    for first in (0..n_blocks).step_by(batch) {
        let blocks: Vec<Vec<GradStats>> = (first..(first + batch).min(n_blocks))
            .into_par_iter()
            .map(|b| block_stats(b * REFRESH_BLOCK_ROWS..((b + 1) * REFRESH_BLOCK_ROWS).min(n)))
            .collect();
        for block in blocks {
            for (acc, s) in total.iter_mut().zip(block) {
                acc.add(s);
            }
        }
    }
    total
}

#[cfg(test)]
mod tests {
    use super::*;
    use crate::tree::{ChildLeaf, SplitRule};

    /// Stump on feature 0 at 0.5 (missing left) with placeholder statistics.
    fn stump() -> RegTree {
        let mut t = RegTree::with_root(99.0);
        t.expand(
            0,
            SplitRule::numeric(0, 0.5, true),
            ChildLeaf::new(7.0, 99.0),
            ChildLeaf::new(-7.0, 99.0),
        );
        t
    }

    #[test]
    fn refresh_recomputes_cover_gain_and_leaves_from_the_routed_rows() {
        // Rows 0, 1 and the missing row 3 go left; row 2 goes right.
        let data = DMatrix::from_dense(&[0.1, 0.2, 0.9, f32::NAN], 4, 1).unwrap();
        let gpair = [
            GradPair::new(1.0, 1.0),
            GradPair::new(2.0, 1.0),
            GradPair::new(-3.0, 2.0),
            GradPair::new(0.5, 0.5),
        ];
        // min_child_weight = 3 exceeds the right leaf's Hessian (2): refresh
        // still gives it a weight, as XGBoost's refresh does.
        let params = TrainingParams::builder()
            .lambda(1.0)
            .min_child_weight(3.0)
            .build()
            .unwrap();
        let reg = RegParams {
            min_child_weight: 0.0,
            ..RegParams::from_params(&params)
        };
        let mut tree = stump();
        refresh_tree(&mut tree, &data, &gpair, &params, Refresh::default(), 0.25);

        let (left, right, root) = (
            GradStats::new(3.5, 2.5),
            GradStats::new(-3.0, 2.0),
            GradStats::new(0.5, 4.5),
        );
        assert_eq!(tree.node(0).sum_hess, 4.5);
        assert_eq!(tree.node(1).sum_hess, 2.5);
        assert_eq!(tree.node(2).sum_hess, 2.0);
        let gain = calc_gain(left, &reg) + calc_gain(right, &reg) - calc_gain(root, &reg);
        assert_eq!(tree.node(0).split_gain, gain as f32);
        // Leaf = base weight -G / (H + lambda), shrunk by the learning rate.
        assert_eq!(tree.node(1).leaf_value, (-3.5f64 / 3.5) as f32 * 0.25);
        assert_eq!(tree.node(2).leaf_value, (3.0f64 / 3.0) as f32 * 0.25);

        // Without refresh_leaf the statistics change but the leaves do not.
        let mut kept = stump();
        refresh_tree(
            &mut kept,
            &data,
            &gpair,
            &params,
            Refresh::stats_only(),
            0.25,
        );
        assert_eq!(kept.node(1).leaf_value, 7.0);
        assert_eq!(kept.node(2).leaf_value, -7.0);
        assert_eq!(kept.node(0).sum_hess, 4.5);
        assert_eq!(kept.node(0).split_gain, gain as f32);
    }

    #[test]
    fn a_leaf_no_row_reaches_gets_zero_weight() {
        let data = DMatrix::from_dense(&[0.1, 0.2], 2, 1).unwrap();
        let gpair = [GradPair::new(1.0, 1.0), GradPair::new(1.0, 1.0)];
        let mut tree = stump();
        refresh_tree(
            &mut tree,
            &data,
            &gpair,
            &TrainingParams::default(),
            Refresh::default(),
            0.3,
        );
        assert_eq!(tree.node(2).sum_hess, 0.0);
        assert_eq!(tree.node(2).leaf_value, 0.0);
    }

    /// Bounding the buffered blocks does not change the ordered reduction:
    /// any batch size gives bit-identical statistics.
    #[test]
    fn batched_block_reduction_matches_for_every_batch_size() {
        let n = 5 * REFRESH_BLOCK_ROWS + 123;
        let x: Vec<f32> = (0..n).map(|i| ((i * 37) % 101) as f32 / 101.0).collect();
        let data = DMatrix::from_dense(&x, n, 1).unwrap();
        let gpair: Vec<GradPair> = (0..n)
            .map(|i| {
                let g = ((i * 7919) % 1000) as f32 * 1e-3 - 0.37;
                GradPair::new(g, 0.1 + (i % 13) as f32)
            })
            .collect();
        let tree = stump();
        let unbatched = node_stats_batched(&tree, &data, &gpair, usize::MAX);
        assert!(unbatched.iter().all(|s| s.hess > 0.0));
        for batch in [0, 1, 2, 4, 6] {
            let stats = node_stats_batched(&tree, &data, &gpair, batch);
            for (a, b) in stats.iter().zip(&unbatched) {
                assert_eq!(
                    (a.grad.to_bits(), a.hess.to_bits()),
                    (b.grad.to_bits(), b.hess.to_bits()),
                    "batch {batch}"
                );
            }
        }
    }
}