hessboost 0.2.2

Fast, deterministic gradient boosting (GBDT) in Rust: conformal intervals, explainable boosting machines, distributional boosting, tree-based diffusion, and XGBoost model interchange
Documentation
//! The honest ("integrity") refit of a Boulevard model's leaves on an
//! independent sample.

use super::check_data;
use crate::data::DMatrix;
use crate::ebm::{EbmBoulevard, EbmInfo};
use crate::error::{HessboostError, Result};
use crate::model::BoostedModel;
use crate::rng::Rng;
use crate::training::boulevard::{Recursion, RoundRequest, Schedule};
use crate::tree::RegTree;

/// The refit's RNG salt, separating its draws from the training run's.
const REFIT_SALT: u64 = 0x1_0E57;

/// Every node's count of the rows whose leaves are `leaf_counts` (per node
/// id; internal nodes `0`), summed up the tree.
fn node_counts(tree: &RegTree, leaf_counts: &[usize]) -> Vec<usize> {
    let nodes = tree.nodes();
    let mut counts = leaf_counts.to_vec();
    // Post-order without assuming children follow their parents.
    let mut stack = vec![(0usize, false)];
    while let Some((id, expanded)) = stack.pop() {
        let node = &nodes[id];
        if node.is_leaf() {
            continue;
        }
        let (l, r) = (node.left as usize, node.right as usize);
        if expanded {
            counts[id] = counts[l] + counts[r];
        } else {
            stack.push((id, true));
            stack.push((l, false));
            stack.push((r, false));
        }
    }
    counts
}

/// The leaf node id of every row of the refit sample in every tree.
struct LeafIds {
    /// `[row][tree]`.
    ids: Vec<u32>,
    rows: usize,
    trees: usize,
}

impl LeafIds {
    fn new(model: &BoostedModel, values: &DMatrix) -> Result<Self> {
        Ok(LeafIds {
            ids: model.predict_leaf(values, ..)?.into_vec(),
            rows: values.n_rows(),
            trees: model.num_trees(),
        })
    }

    /// The leaf of row `row` in tree `tree`.
    fn leaf(&self, row: usize, tree: usize) -> usize {
        self.ids[row * self.trees + tree] as usize
    }
}

/// Per node of one tree, the sum and count of the residuals of the sampled
/// rows reaching it (leaves only; internal nodes stay `0`).
struct LeafSums {
    sums: Vec<f64>,
    counts: Vec<usize>,
}

impl LeafSums {
    /// Add `residual(row)` of every row `in_bag` keeps to its leaf of tree
    /// `tree` (of `nodes` nodes), in row order.
    fn accumulate(
        leaves: &LeafIds,
        tree: usize,
        nodes: usize,
        in_bag: &[bool],
        residual: impl Fn(usize) -> f64,
    ) -> Self {
        let mut sums = vec![0.0f64; nodes];
        let mut counts = vec![0usize; nodes];
        for (row, &keep) in in_bag.iter().enumerate() {
            if keep {
                let leaf = leaves.leaf(row, tree);
                sums[leaf] += residual(row);
                counts[leaf] += 1;
            }
        }
        LeafSums { sums, counts }
    }

    /// Every leaf's `Σ z / (m + lambda)` (`0` when that denominator is not
    /// positive, and at internal nodes), per node id of `tree`.
    fn leaf_values(&self, tree: &RegTree, reg_lambda: f64) -> Vec<f64> {
        let mut values = vec![0.0f64; tree.num_nodes()];
        for (id, node) in tree.nodes().iter().enumerate() {
            let denom = self.counts[id] as f64 + reg_lambda;
            if node.is_leaf() && denom > 0.0 {
                values[id] = self.sums[id] / denom;
            }
        }
        values
    }
}

/// Set every node's cover to the number of rows of `leaves` reaching it,
/// tree by tree.
fn update_covers(trees: &mut [RegTree], leaves: &LeafIds) {
    for (t, tree) in trees.iter_mut().enumerate() {
        let mut leaf_counts = vec![0usize; tree.num_nodes()];
        for row in 0..leaves.rows {
            leaf_counts[leaves.leaf(row, t)] += 1;
        }
        for (id, &c) in node_counts(tree, &leaf_counts).iter().enumerate() {
            tree.set_sum_hess(id, c as f32);
        }
    }
}

/// Refit every leaf of the Boulevard model `model` on `values`, labelled
/// rows independent of its training data, keeping every tree's structure:
/// the Boulevard recursion the model was trained with (BRAT-D with its
/// dropout, or BRAT-P, with the same learning rate, row subsample ratio, L2
/// penalty, and truncation) is rerun on `values`, each tree's leaves set to
/// `Σ z / (m + lambda)` over the `m` rows of a fresh row sample reaching
/// them (`0` for a leaf none reaches), and the result scaled as training
/// scales it. A label-mean intercept is re-estimated on `values`.
///
/// A Boulevard EBM (`booster = ebm` with `ebm_boulevard`) is refitted the
/// same way through its own recursion (Algorithm 1 of Fang, Tan, Pipping &
/// Hooker): stage by stage, every round's trees on the same residuals,
/// each centered on `values`; its term means are re-estimated on `values`.
///
/// The structures then depend on the training labels only, and the leaf
/// values on `values`' labels only: Fang, Tan & Hooker's *integrity*
/// (Zhou & Hooker's structure–value isolation), under which
/// [`BoulevardInference`](super::BoulevardInference) (or
/// [`EbmInference`](super::EbmInference)) is fitted on `values` (not the
/// training rows). Node covers become `values`' row counts, so SHAP values
/// of the refitted model are relative to `values`.
///
/// # Errors
///
/// [`HessboostError::IncompatibleModel`] (`model`) when `model` is not a
/// Boulevard fit or Boulevard EBM; [`HessboostError::InvalidData`]
/// (`values`) when `values` lacks labels, has non-finite labels, row
/// weights or base margins, or refits leaves that overflow `f32`;
/// [`HessboostError::DimensionMismatch`] for a different feature count.
pub fn honest_refit(model: &BoostedModel, values: &DMatrix) -> Result<BoostedModel> {
    if let Some(info) = model.ebm()
        && let Some(settings) = info.boulevard
    {
        return ebm_refit(model, info, settings, values);
    }
    let info = *model.boulevard().ok_or_else(|| {
        HessboostError::incompatible_model(
            "model",
            "not a Boulevard fit: train it with `booster = boulevard` (or `booster = ebm` with \
             `ebm_boulevard`)",
        )
    })?;
    let labels = finite_labels(model, values)?;
    let n = values.n_rows();
    let mu = if info.intercept_from_labels {
        (labels.iter().map(|&y| f64::from(y)).sum::<f64>() / n as f64) as f32
    } else {
        model.base_score()
    };
    let leaves = LeafIds::new(model, values)?;
    let parallel = model.num_parallel_tree();
    let schedule = Schedule::from_info(&info, parallel, REFIT_SALT);
    let mut refit = model.clone();
    let mut recursion = Recursion::new(schedule, n);
    let trees = refit.trees_mut();
    for _ in 0..model.num_boost_rounds() {
        recursion.step(|request| {
            let RoundRequest {
                index,
                first_slot,
                offsets,
                rng,
            } = request;
            let mut preds = Vec::with_capacity(offsets.len());
            for (j, offset) in offsets.iter().enumerate() {
                let t = index * parallel + first_slot + j;
                let in_bag = row_sample(n, info.subsample, rng);
                let tree = &mut trees[t];
                let sums = LeafSums::accumulate(&leaves, t, tree.num_nodes(), &in_bag, |row| {
                    f64::from(labels[row]) - f64::from(mu) - offset[row]
                });
                let values: Vec<f32> = sums
                    .leaf_values(tree, info.reg_lambda)
                    .into_iter()
                    .map(|v| v as f32)
                    .collect();
                for (id, &v) in values.iter().enumerate() {
                    if tree.nodes()[id].is_leaf() {
                        tree.set_leaf_value(id, v);
                    }
                }
                preds.push((0..n).map(|row| values[leaves.leaf(row, t)]).collect());
            }
            Ok(preds)
        })?;
    }
    let scale = recursion.scale() as f32;
    for tree in trees.iter_mut() {
        tree.scale_leaves(scale);
    }
    update_covers(trees, &leaves);
    refit.set_base_scores(vec![mu]);
    refit.set_best_iteration(None);
    if refit
        .trees()
        .iter()
        .any(|t| t.nodes().iter().any(|n| !n.leaf_value.is_finite()))
    {
        return Err(HessboostError::invalid_data(
            "values",
            "the refitted leaves overflow f32",
        ));
    }
    Ok(refit)
}

/// The labels of `values` (checked finite) after [`check_data`].
fn finite_labels<'v>(model: &BoostedModel, values: &'v DMatrix) -> Result<&'v [f32]> {
    check_data(model, values, "values", true)?;
    let labels = values.labels().unwrap_or_default();
    if labels.iter().any(|y| !y.is_finite()) {
        return Err(HessboostError::invalid_data(
            "values",
            "labels must be finite",
        ));
    }
    Ok(labels)
}

/// [`honest_refit`] of a Boulevard EBM: its main-effect stage from the
/// label mean, then its pair stage from the refitted main effects, each
/// round refitting one tree per term of the stage (trees are stored stage
/// by stage, round by round, in term order).
fn ebm_refit(
    model: &BoostedModel,
    info: &EbmInfo,
    settings: EbmBoulevard,
    values: &DMatrix,
) -> Result<BoostedModel> {
    let labels = finite_labels(model, values)?;
    let n = values.n_rows();
    let mu = labels.iter().map(|&y| f64::from(y)).sum::<f64>() / n as f64;
    let leaves = LeafIds::new(model, values)?;
    let mut refit = model.clone();
    let mut base = vec![mu; n];
    for stage in info.stages()? {
        let terms = stage.terms.len();
        let schedule = Schedule {
            dropout: 0.0,
            learning_rate: settings.learning_rate,
            truncation: None,
            parallel: 1,
            seed: 0,
            salt: REFIT_SALT ^ stage.index as u64,
        };
        let mut recursion = Recursion::new(schedule, n);
        let mut total = vec![0.0f64; n];
        let trees = refit.trees_mut();
        for round in 0..stage.rounds {
            recursion.step(|request| {
                let RoundRequest { offsets, rng, .. } = request;
                let mut round_sum = vec![0.0f64; n];
                for k in 0..terms {
                    let t = stage.trees.start + round * terms + k;
                    let in_bag = row_sample(n, settings.subsample, rng);
                    let tree = &mut trees[t];
                    let sums = LeafSums::accumulate(&leaves, t, tree.num_nodes(), &in_bag, |row| {
                        f64::from(labels[row]) - base[row] - offsets[0][row]
                    });
                    let leaf_values = sums.leaf_values(tree, settings.reg_lambda);
                    let mean = (0..n)
                        .map(|row| leaf_values[leaves.leaf(row, t)])
                        .sum::<f64>()
                        / n as f64;
                    for (id, v) in leaf_values.iter().enumerate() {
                        if tree.nodes()[id].is_leaf() {
                            tree.set_leaf_value(id, (v - mean) as f32);
                        }
                    }
                    for (row, s) in round_sum.iter_mut().enumerate() {
                        *s += f64::from(tree.nodes()[leaves.leaf(row, t)].leaf_value);
                    }
                }
                for (t, &s) in total.iter_mut().zip(&round_sum) {
                    *t += s;
                }
                Ok(vec![round_sum.iter().map(|&s| s as f32).collect()])
            })?;
        }
        let scale = recursion.scale();
        for tree in &mut trees[stage.trees] {
            tree.scale_leaves(scale as f32);
        }
        for (b, &t) in base.iter_mut().zip(&total) {
            *b += scale * t;
        }
    }
    update_covers(refit.trees_mut(), &leaves);
    refit.set_base_scores(vec![mu as f32]);
    if refit
        .trees()
        .iter()
        .any(|t| t.nodes().iter().any(|n| !n.leaf_value.is_finite()))
    {
        return Err(HessboostError::invalid_data(
            "values",
            "the refitted leaves overflow f32",
        ));
    }
    let mut refit_info = info.clone();
    refit_info.term_means = refit_info.term_means_on(&refit, values);
    refit.set_ebm(Some(refit_info));
    Ok(refit)
}

/// A Bernoulli row sample: row `i` is kept with probability `ratio` (every
/// row when `ratio >= 1`), at least one row, as training samples them.
fn row_sample(n: usize, ratio: f64, rng: &mut Rng) -> Vec<bool> {
    if ratio >= 1.0 {
        return vec![true; n];
    }
    let mut keep: Vec<bool> = (0..n).map(|_| rng.f64() < ratio).collect();
    if !keep.iter().any(|&k| k) {
        keep[rng.range(0..n)] = true;
    }
    keep
}