mod categories;
pub mod compact;
pub(crate) mod container;
mod embed;
mod io;
mod lightgbm;
pub(crate) mod native;
mod objective;
mod predict;
mod predictions;
pub(crate) mod sections;
mod serde;
mod shap;
mod shrinkage;
mod slice;
mod ubjson;
pub mod uncertainty;
mod validate;
mod xgboost;
pub(crate) use shrinkage::{Shrinkage, shrink_margins};
pub use embed::EmbeddedModel;
pub use io::ModelFormat;
pub use objective::ModelObjective;
pub use predict::Iterations;
use predict::RowBlock;
pub(crate) use predict::{initial_margins, transform_margins_in_place, transform_model_margins};
pub use predictions::{Contributions, Interactions, Predictions};
pub(crate) use validate::{check_objective_width, validate_prediction_data};
use self::serde::UncheckedBoostedModel;
use crate::data::DMatrix;
use crate::ebm::EbmInfo;
use crate::error::{HessboostError, Result};
use crate::inference::BoulevardInfo;
use crate::objective::{Loss, LossContext};
use crate::tree::compact::CompactForest;
use crate::tree::{RegTree, scalar_tree_output};
use ::serde::{Deserialize, Serialize};
use std::collections::BTreeMap;
use std::ops::{Bound, Range, RangeBounds};
use std::sync::Arc;
use std::sync::OnceLock;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum ImportanceType {
Weight,
TotalGain,
Gain,
TotalCover,
Cover,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(try_from = "UncheckedBoostedModel")]
#[allow(
clippy::unsafe_derive_deserialize,
reason = "the type's only unsafe code is the optional Metal backend's \
buffer handling; every serialized field is plain data"
)]
pub struct BoostedModel {
trees: Vec<RegTree>,
base_score: Vec<f32>,
objective: ModelObjective,
max_delta_step: f64,
num_class: usize,
n_outputs: usize,
n_targets: usize,
n_features: usize,
best_iteration: Option<usize>,
tree_weights: TreeWeights,
num_parallel_tree: usize,
linear: Option<LinearModel>,
shrinkage: Option<Shrinkage>,
boulevard: Option<BoulevardInfo>,
ebm: Option<EbmInfo>,
compact: OnceLock<CompactForest>,
}
#[derive(Debug, Clone, PartialEq)]
pub(crate) enum TreeWeights {
Unit,
Explicit(Vec<f32>),
}
impl TreeWeights {
pub(crate) fn from_vec(weights: Vec<f32>) -> Self {
if weights.is_empty() {
Self::Unit
} else {
Self::Explicit(weights)
}
}
pub(crate) fn as_slice(&self) -> &[f32] {
match self {
Self::Unit => &[],
Self::Explicit(weights) => weights,
}
}
pub(crate) fn iter(&self) -> std::slice::Iter<'_, f32> {
self.as_slice().iter()
}
#[inline]
fn get(&self, i: usize) -> f32 {
match self {
Self::Unit => 1.0,
Self::Explicit(weights) => weights[i],
}
}
fn push(&mut self, n_trees: usize, weight: f32) {
self.materialize(n_trees);
match self {
Self::Unit => *self = Self::Explicit(vec![weight]),
Self::Explicit(weights) => weights.push(weight),
}
}
fn materialize(&mut self, n_trees: usize) {
match self {
Self::Unit if n_trees > 0 => *self = Self::Explicit(vec![1.0; n_trees]),
Self::Unit => {}
Self::Explicit(weights) => weights.resize(n_trees, 1.0),
}
}
fn scale(&mut self, i: usize, factor: f32) {
if let Self::Explicit(weights) = self
&& let Some(weight) = weights.get_mut(i)
{
*weight *= factor;
}
}
fn select(&self, layers: impl Iterator<Item = Range<usize>>) -> Self {
match self {
Self::Unit => Self::Unit,
Self::Explicit(weights) => Self::from_vec(
layers
.flat_map(|layer| weights[layer].iter().copied())
.collect(),
),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub(crate) struct LinearModel {
weights: Vec<f32>,
bias: Vec<f32>,
}
impl LinearModel {
pub(crate) fn new(weights: Vec<f32>, bias: Vec<f32>) -> Self {
LinearModel { weights, bias }
}
pub(crate) fn bias(&self) -> &[f32] {
&self.bias
}
pub(crate) fn weights(&self) -> &[f32] {
&self.weights
}
}
pub(crate) fn for_each_present_value(data: &DMatrix, row: usize, mut f: impl FnMut(usize, f32)) {
for feat in 0..data.n_cols() {
if let Some(x) = data.get(row, feat) {
f(feat, x);
}
}
}
struct AttributionPrologue<'a> {
n: usize,
k: usize,
nf: usize,
width: usize,
trees: &'a [RegTree],
initial: Vec<f32>,
}
pub(crate) struct ModelSpec {
pub(crate) objective: ModelObjective,
pub(crate) max_delta_step: f64,
pub(crate) num_class: usize,
pub(crate) n_outputs: usize,
pub(crate) n_targets: usize,
pub(crate) n_features: usize,
}
impl BoostedModel {
pub(crate) fn new(base_score: Vec<f32>, spec: ModelSpec) -> Self {
Self::from_parts(Vec::new(), Vec::new(), base_score, spec)
}
pub(crate) fn set_linear(&mut self, linear: LinearModel) {
self.linear = Some(linear);
}
pub(crate) fn with_trees(&self, trees: Vec<RegTree>) -> BoostedModel {
BoostedModel {
trees,
base_score: self.base_score.clone(),
objective: self.objective.clone(),
max_delta_step: self.max_delta_step,
num_class: self.num_class,
n_outputs: self.n_outputs,
n_targets: self.n_targets,
n_features: self.n_features,
best_iteration: None,
tree_weights: TreeWeights::Unit,
num_parallel_tree: self.num_parallel_tree,
linear: None,
shrinkage: None,
boulevard: None,
ebm: None,
compact: OnceLock::new(),
}
}
pub(crate) fn push_tree_weighted(&mut self, tree: RegTree, weight: f32) {
self.tree_weights.push(self.trees.len(), weight);
self.trees.push(tree);
self.compact = OnceLock::new();
}
pub(crate) fn compact_forest(&self) -> &CompactForest {
self.compact
.get_or_init(|| CompactForest::from_trees(&self.trees))
}
#[inline]
pub(crate) fn tree_weight(&self, i: usize) -> f32 {
self.tree_weights.get(i)
}
#[cfg(all(target_os = "macos", feature = "metal"))]
pub(crate) fn is_gblinear(&self) -> bool {
self.linear.is_some()
}
#[cfg(all(target_os = "macos", feature = "metal"))]
pub(crate) fn has_linear_leaves(&self) -> bool {
self.trees.iter().any(|tree| tree.linear_leaves().is_some())
}
#[cfg(all(target_os = "macos", feature = "metal"))]
pub(crate) fn tree_is_vector_leaf(&self, t: usize) -> bool {
self.trees[t].is_vector_leaf()
}
pub(crate) fn for_each_linear_contribution(
&self,
data: &DMatrix,
row: usize,
mut f: impl FnMut(usize, usize, f64),
) {
let Some(lm) = &self.linear else {
return;
};
let k = self.n_outputs();
for_each_present_value(data, row, |feat, x| {
for c in 0..k {
f(feat, c, f64::from(lm.weights[feat * k + c]) * f64::from(x));
}
});
}
pub(crate) fn scale_tree_weight(&mut self, i: usize, factor: f32) {
self.tree_weights.scale(i, factor);
}
pub(crate) fn predict_margin_dropout(&self, data: &DMatrix, dropped: &[bool]) -> Vec<f32> {
let mut out = self.initial_margins(data);
let weight = |ti: usize| {
if dropped.get(ti).copied().unwrap_or(false) {
0.0
} else {
self.tree_weight(ti)
}
};
self.accumulate_forest(data, &mut out, 0..self.trees.len(), weight);
out
}
pub(crate) fn set_best_iteration(&mut self, it: Option<usize>) {
self.best_iteration = it;
}
pub(crate) fn set_boulevard(&mut self, info: Option<BoulevardInfo>) {
self.boulevard = info;
}
pub fn boulevard(&self) -> Option<&BoulevardInfo> {
self.boulevard.as_ref()
}
pub(crate) fn set_ebm(&mut self, info: Option<EbmInfo>) {
self.ebm = info;
}
pub fn ebm(&self) -> Option<&EbmInfo> {
self.ebm.as_ref()
}
pub(crate) fn from_parts(
trees: Vec<RegTree>,
tree_weights: Vec<f32>,
base_score: Vec<f32>,
spec: ModelSpec,
) -> Self {
BoostedModel {
trees,
base_score,
objective: spec.objective,
max_delta_step: spec.max_delta_step,
num_class: spec.num_class,
n_outputs: spec.n_outputs,
n_targets: spec.n_targets,
n_features: spec.n_features,
best_iteration: None,
tree_weights: TreeWeights::from_vec(tree_weights),
num_parallel_tree: 1,
linear: None,
shrinkage: None,
boulevard: None,
ebm: None,
compact: OnceLock::new(),
}
}
pub(crate) fn num_class(&self) -> usize {
self.num_class
}
pub(crate) fn max_delta_step(&self) -> f64 {
self.max_delta_step
}
#[inline]
pub fn n_outputs(&self) -> usize {
self.n_outputs
}
#[inline]
pub fn n_targets(&self) -> usize {
self.n_targets
}
pub fn has_vector_leaves(&self) -> bool {
self.trees.first().is_some_and(RegTree::is_vector_leaf)
}
pub fn base_score(&self) -> f32 {
self.base_score[0]
}
pub fn base_scores(&self) -> &[f32] {
&self.base_score
}
pub fn objective(&self) -> &ModelObjective {
&self.objective
}
pub fn best_iteration(&self) -> Option<usize> {
self.best_iteration
}
pub fn feature_importance(&self, kind: ImportanceType) -> BTreeMap<usize, f64> {
let value = |node: &crate::tree::Node| match kind {
ImportanceType::Weight => 1.0,
ImportanceType::Cover | ImportanceType::TotalCover => f64::from(node.sum_hess),
ImportanceType::Gain | ImportanceType::TotalGain => f64::from(node.split_gain),
};
let mut totals: BTreeMap<usize, (f64, f64)> = BTreeMap::new();
for tree in &self.trees {
for node in tree.nodes() {
if node.is_leaf() {
continue;
}
let (total, count) = totals.entry(node.split_feature as usize).or_default();
*total += value(node);
*count += 1.0;
}
}
let average = matches!(kind, ImportanceType::Cover | ImportanceType::Gain);
totals
.into_iter()
.map(|(f, (total, count))| (f, if average { total / count } else { total }))
.collect()
}
#[cfg(not(all(target_os = "macos", feature = "metal")))]
pub fn to_gpu(&self) -> Result<crate::backend::metal::GpuModel> {
Err(HessboostError::gpu(
"GPU prediction requires the `metal` feature on macOS",
))
}
pub fn trees(&self) -> &[RegTree] {
&self.trees
}
pub fn n_features(&self) -> usize {
self.n_features
}
pub(crate) fn linear(&self) -> Option<&LinearModel> {
self.linear.as_ref()
}
pub(crate) fn has_non_unit_tree_weights(&self) -> bool {
(0..self.trees.len()).any(|i| self.tree_weight(i) != 1.0)
}
pub(crate) fn initial_margins(&self, data: &DMatrix) -> Vec<f32> {
initial_margins(&self.base_score, data)
}
fn attribution_prologue(
&self,
data: &DMatrix,
iterations: Iterations,
what: &str,
) -> Result<AttributionPrologue<'_>> {
self.validate_prediction_data(data)?;
let end = self.prefix_trees(iterations, what)?;
self.refuse_partial_shrunk_range(&(0..end), what)?;
let trees = &self.trees[..end];
if trees.iter().any(|tree| tree.linear_leaves().is_some()) {
return Err(HessboostError::incompatible_model(
"linear_tree",
"SHAP contributions and interactions are not defined for models with linear leaves",
));
}
let nf = self.n_features;
Ok(AttributionPrologue {
n: data.n_rows(),
k: self.n_outputs(),
nf,
width: nf + 1,
trees,
initial: self.initial_margins(data),
})
}
pub(crate) fn set_num_parallel_tree(&mut self, num_parallel_tree: usize) {
self.num_parallel_tree = num_parallel_tree;
}
pub(crate) fn set_base_scores(&mut self, base_score: Vec<f32>) {
self.base_score = base_score;
}
pub(crate) fn set_objective(&mut self, objective: ModelObjective, max_delta_step: f64) {
self.objective = objective;
self.max_delta_step = max_delta_step;
}
pub(crate) fn materialize_tree_weights(&mut self) {
self.tree_weights.materialize(self.trees.len());
}
pub(crate) fn take_trees(&mut self) -> Vec<RegTree> {
self.tree_weights = TreeWeights::Unit;
self.compact = OnceLock::new();
std::mem::take(&mut self.trees)
}
pub(crate) fn scale_all_leaves(&mut self, factor: f32) {
for tree in self.trees_mut() {
tree.scale_leaves(factor);
}
}
pub(crate) fn trees_mut(&mut self) -> &mut [RegTree] {
self.compact = OnceLock::new();
&mut self.trees
}
pub fn num_trees(&self) -> usize {
self.trees.len()
}
#[inline]
pub fn num_parallel_tree(&self) -> usize {
self.num_parallel_tree
}
#[inline]
pub fn trees_per_iteration(&self) -> usize {
if self.has_vector_leaves() {
self.num_parallel_tree
} else {
self.n_outputs * self.num_parallel_tree
}
}
pub(crate) fn check_iteration_size(n_outputs: usize, num_parallel_tree: usize) -> Result<()> {
if n_outputs.checked_mul(num_parallel_tree).is_none() {
return Err(HessboostError::invalid_param(
"num_parallel_tree",
format!(
"{num_parallel_tree} parallel trees for {n_outputs} outputs overflow the \
trees per iteration"
),
));
}
Ok(())
}
#[inline]
pub(crate) fn tree_output(&self, t: usize) -> usize {
if self.has_vector_leaves() {
0
} else {
scalar_tree_output(t, self.num_parallel_tree, self.n_outputs)
}
}
pub fn num_boost_rounds(&self) -> usize {
self.trees.len() / self.trees_per_iteration()
}
pub(crate) fn effective_num_trees(&self) -> usize {
self.best_iteration.map_or(self.trees.len(), |it| {
((it + 1) * self.trees_per_iteration()).min(self.trees.len())
})
}
fn iteration_bounds(&self, iterations: Iterations) -> (Bound<usize>, Bound<usize>) {
match iterations {
Iterations::Best => {
let end = self
.best_iteration
.map_or(Bound::Unbounded, |it| Bound::Excluded(it + 1));
(Bound::Unbounded, end)
}
Iterations::Range { start, end } => (start, end),
}
}
pub(crate) fn resolve_iterations(
&self,
iterations: Iterations,
param: &'static str,
) -> Result<Range<usize>> {
let iterations = self.iteration_bounds(iterations);
if self.linear.is_some() {
let whole = matches!(
iterations.start_bound(),
Bound::Unbounded | Bound::Included(0)
) && iterations.end_bound() == Bound::Unbounded;
if !whole {
return Err(HessboostError::incompatible_model(
param,
"gblinear models have no boosting iterations to select; pass `..`",
));
}
return Ok(0..0);
}
let rounds = self.num_boost_rounds();
let overflow = || HessboostError::invalid_param(param, "range bound overflows");
let begin = match iterations.start_bound() {
Bound::Included(&b) => b,
Bound::Excluded(&b) => b.checked_add(1).ok_or_else(overflow)?,
Bound::Unbounded => 0,
};
let end = match iterations.end_bound() {
Bound::Included(&e) => e.checked_add(1).ok_or_else(overflow)?,
Bound::Excluded(&e) => e,
Bound::Unbounded => rounds,
};
if end > rounds {
return Err(HessboostError::incompatible_model(
param,
format!("{begin}..{end} is out of range for a model with {rounds} iterations"),
));
}
if begin > end {
return Err(HessboostError::invalid_param(
param,
format!("{begin}..{end} is an inverted range"),
));
}
Ok(begin..end)
}
pub(crate) fn iteration_trees(&self, iterations: Range<usize>) -> Range<usize> {
let per = self.trees_per_iteration();
iterations.start * per..iterations.end * per
}
fn prefix_trees(&self, iterations: Iterations, what: &str) -> Result<usize> {
let trees = self.iteration_trees(self.resolve_iterations(iterations, "iterations")?);
if trees.start != 0 {
return Err(HessboostError::invalid_param(
"iterations",
format!(
"{what} supports only ranges starting at iteration 0; slice the model instead"
),
));
}
Ok(trees.end)
}
pub(crate) fn set_shrinkage(&mut self, shrinkage: Shrinkage) {
let (tree_weights, base_score) =
shrinkage.scaling(self.num_boost_rounds(), self.trees_per_iteration());
self.tree_weights = TreeWeights::from_vec(tree_weights);
self.base_score = base_score;
self.shrinkage = Some(shrinkage);
}
pub(crate) fn truncate_shrunk(&mut self, k: usize) {
if let Some(shrinkage) = &self.shrinkage
&& k < self.num_boost_rounds()
{
*self = self.shrunk_prefix(shrinkage, k);
}
}
pub(crate) fn shrinkage(&self) -> Option<&Shrinkage> {
self.shrinkage.as_ref()
}
pub(crate) fn validate_prediction_data(&self, data: &DMatrix) -> Result<()> {
validate_prediction_data(self.n_features, self.n_outputs(), data)
}
pub(crate) fn rebuild_objective(&self) -> Option<Result<Arc<dyn Loss>>> {
rebuild_objective(&self.objective, self.max_delta_step, self.n_targets)
}
}
fn rebuild_objective(
objective: &ModelObjective,
max_delta_step: f64,
n_targets: usize,
) -> Option<Result<Arc<dyn Loss>>> {
objective.built_in().map(|objective| {
objective.build_loss(&LossContext {
n_targets,
max_delta_step,
shared_tree_seed: None,
seed: 0,
})
})
}
#[cfg(test)]
mod tests {
use super::BoostedModel;
use crate::config::TrainingParams;
use crate::data::DMatrix;
use crate::error::HessboostError;
use crate::model::Iterations;
use crate::model::ModelFormat;
use crate::objective::{Objective, RegLoss};
use crate::test_support::labeled_dense;
use crate::training::train;
#[test]
fn training_rejects_overflowing_iteration_size() {
let x = [0.0f32, 1.0, 2.0, 3.0];
let y = [0.0f32, 1.0, 1.0, 0.0, 2.0, 3.0, 3.0, 2.0];
let d = DMatrix::from_dense(&x, 4, 1)
.unwrap()
.with_label_matrix(&y, 2)
.unwrap();
let params = TrainingParams {
num_parallel_tree: 1usize << (usize::BITS - 1),
..TrainingParams::default()
};
for rounds in [0, 1] {
assert!(matches!(
train(¶ms, &d, rounds),
Err(HessboostError::InvalidParameter { name, .. }) if name == "num_parallel_tree"
));
}
}
#[test]
fn loading_propagates_invalid_builtin_objective() {
let d = labeled_dense(&[0.0, 1.0], 2, 1, &[0.0, 1.0]);
let model = train(&TrainingParams::default(), &d, 1).unwrap();
let mut value: serde_json::Value =
serde_json::from_slice(&model.encode(ModelFormat::Json).unwrap()).unwrap();
value["n_targets"] = 2.into();
value["objective"] = "count:poisson".into();
assert!(matches!(
BoostedModel::decode(value.to_string(), ModelFormat::Json),
Err(HessboostError::ModelFormat(_))
));
value["objective"] = "my:custom".into();
let custom = BoostedModel::decode(value.to_string(), ModelFormat::Json).unwrap();
assert_eq!(custom.objective().name(), "my:custom");
assert_eq!(custom.objective().built_in(), None);
}
#[test]
fn loading_refuses_non_finite_linear_parameters() {
let x: Vec<f32> = (0..20).map(|i| i as f32).collect();
let d = labeled_dense(&x, 10, 2, &[1.0; 10]);
let params = TrainingParams::builder()
.booster(crate::config::BoosterKind::GbLinear)
.build()
.unwrap();
let model = train(¶ms, &d, 2).unwrap();
assert!(
BoostedModel::decode(
model.encode(ModelFormat::Binary).unwrap(),
ModelFormat::Binary
)
.is_ok()
);
for (weight, bias) in [(f32::INFINITY, 0.0), (0.0, f32::NAN)] {
let mut corrupt = model.clone();
let linear = corrupt.linear.as_mut().unwrap();
linear.weights[0] = weight;
linear.bias[0] = bias;
assert!(matches!(
BoostedModel::decode(
corrupt.encode(ModelFormat::Binary).unwrap(),
ModelFormat::Binary
),
Err(HessboostError::ModelFormat(_))
));
}
}
fn gblinear_doc() -> (BoostedModel, serde_json::Value) {
let x: Vec<f32> = (0..20).map(|i| i as f32).collect();
let d = labeled_dense(&x, 10, 2, &[1.0; 10]);
let params = TrainingParams::builder()
.booster(crate::config::BoosterKind::GbLinear)
.build()
.unwrap();
let model = train(¶ms, &d, 2).unwrap();
let doc = serde_json::from_slice(&model.encode(ModelFormat::Json).unwrap()).unwrap();
(model, doc)
}
#[test]
fn serde_deserialization_validates_the_model() {
let (model, mut doc) = gblinear_doc();
let valid: BoostedModel = serde_json::from_value(doc.clone()).unwrap();
let d = DMatrix::from_dense(&[1.0, 2.0], 1, 2).unwrap();
assert_eq!(
valid.predict(&d, Iterations::Best).unwrap(),
model.predict(&d, Iterations::Best).unwrap()
);
doc["linear"]["bias"] = serde_json::json!([]);
assert!(serde_json::from_value::<BoostedModel>(doc.clone()).is_err());
assert!(matches!(
BoostedModel::decode(doc.to_string(), ModelFormat::Json),
Err(HessboostError::ModelFormat(_))
));
let d = labeled_dense(&[0.0, 1.0, 2.0, 3.0], 4, 1, &[0.0, 0.0, 1.0, 1.0]);
let model = train(&TrainingParams::default(), &d, 1).unwrap();
let mut doc: serde_json::Value =
serde_json::from_slice(&model.encode(ModelFormat::Json).unwrap()).unwrap();
assert!(doc["trees"][0]["nodes"].as_array().unwrap().len() > 1);
doc["trees"][0]["nodes"][0]["left"] = 0.into();
assert!(serde_json::from_value::<BoostedModel>(doc.clone()).is_err());
assert!(matches!(
BoostedModel::decode(doc.to_string(), ModelFormat::Json),
Err(HessboostError::ModelFormat(_))
));
}
#[test]
fn gblinear_refuses_best_iteration() {
let (model, mut doc) = gblinear_doc();
doc["best_iteration"] = 0.into();
assert!(matches!(
BoostedModel::decode(doc.to_string(), ModelFormat::Json),
Err(HessboostError::ModelFormat(_))
));
let mut stopped = model.clone();
stopped.set_best_iteration(Some(0));
assert!(matches!(
BoostedModel::decode(
stopped.encode(ModelFormat::Binary).unwrap(),
ModelFormat::Binary
),
Err(HessboostError::ModelFormat(_))
));
}
#[test]
fn num_class_applies_only_to_multiclass_objectives() {
let x = [0.0f32, 1.0, 2.0, 3.0];
let y = [0.0f32, 1.0, 0.0, 1.0, 1.0, 0.0, 1.0, 0.0];
let d = DMatrix::from_dense(&x, 4, 1)
.unwrap()
.with_label_matrix(&y, 2)
.unwrap();
assert!(matches!(
TrainingParams::from_xgboost([
("objective", serde_json::json!("binary:logistic")),
("num_class", serde_json::json!(2)),
]),
Err(HessboostError::InvalidParameter { name, .. }) if name == "num_class"
));
let params = TrainingParams::builder()
.objective(Objective::BinaryLogistic(RegLoss::default()))
.build()
.unwrap();
let model = train(¶ms, &d, 1).unwrap();
let mut doc: serde_json::Value =
serde_json::from_slice(&model.encode(ModelFormat::Json).unwrap()).unwrap();
assert!(BoostedModel::decode(doc.to_string(), ModelFormat::Json).is_ok());
doc["num_class"] = 2.into();
assert!(matches!(
BoostedModel::decode(doc.to_string(), ModelFormat::Json),
Err(HessboostError::ModelFormat(_))
));
}
}