use hessboost::config::{
BoosterKind, Dart, GrowPolicy, Monotone, MultiStrategy, ProcessType, Refresh, TreeMethod,
};
use hessboost::data::FeatureType;
use hessboost::model::{Contributions, Predictions};
use hessboost::objective::{CustomLoss, GradPair, SplitGradient};
use hessboost::prelude::*;
use std::num::NonZeroUsize;
mod common;
use common::{four_features, incompatible_model, invalid_param, labeled_dense, rmse};
const N: usize = 300;
const COLS: usize = 4;
const K: usize = 3;
fn data() -> (Vec<f32>, Vec<f32>) {
let mut x = Vec::with_capacity(N * COLS);
let mut y = Vec::with_capacity(N * K);
for i in 0..N {
let [a, b, c, d] = four_features(i);
x.extend([a, b, c, d]);
y.extend([
2.0 * a - b,
(6.0 * c).sin(),
if d.is_nan() { 1.0 } else { d * a },
]);
}
(x, y)
}
fn dtrain() -> DMatrix {
let (x, y) = data();
DMatrix::from_dense(&x, N, COLS)
.unwrap()
.with_label_matrix(&y, K)
.unwrap()
}
fn vector_params() -> hessboost::config::TrainingParamsBuilder {
TrainingParams::builder()
.multi_strategy(MultiStrategy::MultiOutputTree)
.max_depth(4)
.eta(0.3)
}
fn model() -> BoostedModel {
train(&vector_params().build().unwrap(), &dtrain(), 8).unwrap()
}
fn reference_margins(model: &BoostedModel, x: &[f32], n: usize) -> Vec<f32> {
let mut out = Vec::with_capacity(n * K);
for r in 0..n {
let row = &x[r * COLS..(r + 1) * COLS];
let mut m = model.base_scores().to_vec();
for tree in model.trees() {
let leaf = tree.leaf_id_dense(row, f32::NAN);
for (o, v) in m.iter_mut().zip(tree.leaf_vector(leaf)) {
*o += v;
}
}
out.extend(m);
}
out
}
fn assert_contribs_sum_to(contribs: &Contributions, margins: &Predictions) {
assert_eq!(
(
contribs.n_rows(),
contribs.n_outputs(),
contribs.n_features()
),
(margins.n_rows(), margins.width(), COLS)
);
for (r, row) in margins.rows().enumerate() {
for (o, m) in row.iter().enumerate() {
let sum: f32 = contribs.get(r, o).unwrap().iter().sum();
assert!((sum - m).abs() < 1e-4, "{sum} vs {m}");
}
}
}
#[test]
fn one_vector_tree_per_round_predicts_every_output() {
let model = model();
assert!(model.has_vector_leaves());
assert_eq!(model.num_trees(), 8);
assert_eq!(model.num_boost_rounds(), 8);
assert!(model.trees().iter().all(|t| t.size_leaf_vector() == K));
let (x, _) = data();
let want = reference_margins(&model, &x, N);
let d = DMatrix::from_dense(&x, N, COLS).unwrap();
assert_eq!(
model
.predict_margin(&d, Iterations::Best)
.unwrap()
.as_slice(),
want
);
let few = DMatrix::from_dense(&x[..5 * COLS], 5, COLS).unwrap();
assert_eq!(
model
.predict_margin(&few, Iterations::Best)
.unwrap()
.as_slice(),
&want[..5 * K]
);
let (mut indptr, mut indices, mut values) = (vec![0], Vec::new(), Vec::new());
for r in 0..N {
for c in 0..COLS {
let v = x[r * COLS + c];
if !v.is_nan() {
indices.push(c as u32);
values.push(v);
}
}
indptr.push(indices.len());
}
let csr = DMatrix::from_csr(indptr, indices, values, COLS).unwrap();
assert_eq!(
model
.predict_margin(&csr, Iterations::Best)
.unwrap()
.as_slice(),
want
);
assert_eq!(
model.predict(&d, Iterations::Best).unwrap().as_slice(),
want
);
}
#[test]
fn training_margins_match_the_final_model() {
let dtrain = dtrain();
let result = Trainer::new(
&vector_params().subsample(0.7).seed(3).build().unwrap(),
&dtrain,
6,
)
.eval(&dtrain, "train")
.train()
.unwrap();
let rmse = rmse(&result.model, &dtrain);
let last = result.history.last().unwrap().values()[0];
assert!((last - rmse).abs() < 1e-6, "history {last} vs model {rmse}");
}
#[test]
fn shap_is_additive_per_output() {
let model = model();
let (x, _) = data();
let n = 40;
let d = DMatrix::from_dense(&x[..n * COLS], n, COLS).unwrap();
let margin = model.predict_margin(&d, Iterations::Best).unwrap();
let width = COLS + 1;
let contribs = model.predict_contribs(&d, Iterations::Best).unwrap();
assert_contribs_sum_to(&contribs, &margin);
let inter = model.predict_interactions(&d, Iterations::Best).unwrap();
assert_eq!(
(inter.n_rows(), inter.n_outputs(), inter.n_features()),
(n, K, COLS)
);
for r in 0..n {
for o in 0..K {
for (i, &p) in contribs.get(r, o).unwrap().iter().enumerate() {
let sum: f32 = (0..width).map(|j| inter.at(r, o, i, j).unwrap()).sum();
assert!((sum - p).abs() < 1e-4, "{sum} vs {p}");
}
}
}
}
#[test]
fn formats_round_trip_vector_leaves() {
let model = model();
let d = dtrain();
let want = model.predict(&d, Iterations::Best).unwrap();
let reloaded = [
BoostedModel::decode(
model.encode(ModelFormat::Binary).unwrap(),
ModelFormat::Binary,
)
.unwrap(),
BoostedModel::decode(model.encode(ModelFormat::Json).unwrap(), ModelFormat::Json).unwrap(),
BoostedModel::decode(
model.encode(ModelFormat::XgboostJson).unwrap(),
ModelFormat::XgboostJson,
)
.unwrap(),
BoostedModel::decode(
model.encode(ModelFormat::XgboostUbjson).unwrap(),
ModelFormat::XgboostUbjson,
)
.unwrap(),
];
for m in reloaded {
assert!(m.has_vector_leaves());
assert_eq!(m.predict(&d, Iterations::Best).unwrap(), want);
assert_eq!(m.num_boost_rounds(), 8);
}
}
#[test]
fn early_stopping_keeps_whole_vector_rounds() {
let dtrain = dtrain();
let (x, y) = data();
let shuffled: Vec<f32> = y
.as_chunks::<K>()
.0
.iter()
.rev()
.flatten()
.copied()
.collect();
let holdout = DMatrix::from_dense(&x, N, COLS)
.unwrap()
.with_label_matrix(&shuffled, K)
.unwrap();
let result = Trainer::new(
&vector_params().eta(1.0).max_depth(6).build().unwrap(),
&dtrain,
40,
)
.eval(&holdout, "holdout")
.early_stopping_rounds(NonZeroUsize::new(1).unwrap())
.train()
.unwrap();
let model = result.model;
let best = model.best_iteration().expect("early stopping triggers");
assert_eq!(model.num_trees(), best + 2);
let margin = model.predict_margin(&dtrain, Iterations::Best).unwrap();
assert_eq!(margin, model.predict_margin(&dtrain, ..=best).unwrap());
assert_ne!(margin, model.predict_margin(&dtrain, ..).unwrap());
}
#[test]
fn single_output_builds_scalar_trees() {
let (x, y) = data();
let y0: Vec<f32> = y.iter().step_by(K).copied().collect();
let d = labeled_dense(&x, COLS, &y0);
let vector = train(&vector_params().build().unwrap(), &d, 5).unwrap();
let scalar = train(
&TrainingParams::builder()
.max_depth(4)
.eta(0.3)
.build()
.unwrap(),
&d,
5,
)
.unwrap();
assert!(!vector.has_vector_leaves());
assert_eq!(
vector.predict(&d, Iterations::Best).unwrap(),
scalar.predict(&d, Iterations::Best).unwrap()
);
}
#[test]
fn dart_rounds_train_vector_trees() {
let params = vector_params()
.booster(BoosterKind::Dart(
Dart::builder()
.rate_drop(0.5)
.skip_drop(0.0)
.build()
.unwrap(),
))
.seed(7)
.build()
.unwrap();
let d = dtrain();
let model = train(¶ms, &d, 6).unwrap();
assert!(model.has_vector_leaves());
assert_eq!(model.num_trees(), 6);
let (x, _) = data();
let margins = model.predict_margin(&d, Iterations::Best).unwrap();
let unweighted = reference_margins(&model, &x, N);
assert_ne!(margins.as_slice(), unweighted);
assert_contribs_sum_to(
&model.predict_contribs(&d, Iterations::Best).unwrap(),
&margins,
);
}
fn squared_error(k: usize) -> CustomLoss {
CustomLoss::new("custom:sqerr", k, |p, y, _w, out| {
for (o, (p, y)) in out.iter_mut().zip(p.iter().zip(y)) {
*o = GradPair::new(p - y, 1.0);
}
})
.with_base_margin(0.5)
}
fn mean_sketch(g: &[GradPair]) -> SplitGradient {
let gpair = g
.as_chunks::<K>()
.0
.iter()
.map(|r| {
let (g, h) = r
.iter()
.fold((0.0, 0.0), |(g, h), p| (g + p.grad, h + p.hess));
GradPair::new(g / K as f32, h / K as f32)
})
.collect();
SplitGradient::new(gpair, 1)
}
#[test]
fn reduced_gradients_grow_structure_from_the_sketch() {
let dtrain = dtrain();
let sketched = vector_params()
.lambda(0.0)
.objective(Objective::custom(
squared_error(K).with_split_gradient(|_, g| Some(mean_sketch(g))),
))
.build()
.unwrap();
let model = train(&sketched, &dtrain, 1).unwrap();
let tree = &model.trees()[0];
let (x, y) = data();
let mut sums = vec![[0.0f64; K]; tree.num_nodes()];
let mut counts = vec![0usize; tree.num_nodes()];
for r in 0..N {
let leaf = tree.leaf_id_dense(&x[r * COLS..(r + 1) * COLS], f32::NAN);
counts[leaf] += 1;
for (sum, &label) in sums[leaf].iter_mut().zip(&y[r * K..(r + 1) * K]) {
*sum += f64::from(label - 0.5);
}
}
for (leaf, &count) in counts.iter().enumerate().filter(|(_, c)| **c > 0) {
for (t, sum) in sums[leaf].iter().enumerate() {
let want = 0.3 * sum / count as f64;
let got = f64::from(tree.leaf_vector(leaf)[t]);
assert!(
(got - want).abs() < 1e-5,
"leaf {leaf} target {t}: {got} vs {want}"
);
}
}
let full_params = vector_params()
.lambda(0.0)
.objective(Objective::custom(squared_error(K)))
.build()
.unwrap();
let full = train(&full_params, &dtrain, 1).unwrap();
assert_ne!(full.trees()[0].nodes(), tree.nodes());
}
#[test]
fn unsupported_combinations_are_rejected() {
let dtrain = dtrain();
let exact = vector_params()
.tree_method(TreeMethod::Exact)
.build()
.unwrap();
assert_eq!(invalid_param(train(&exact, &dtrain, 1)), "multi_strategy");
let sketch =
Objective::custom(squared_error(K).with_split_gradient(|_, g| Some(mean_sketch(g))));
let per_output = TrainingParams::builder()
.objective(sketch.clone())
.build()
.unwrap();
assert_eq!(invalid_param(train(&per_output, &dtrain, 1)), "objective");
for strategy in [
MultiStrategy::OneOutputPerTree,
MultiStrategy::MultiOutputTree,
] {
let linear = TrainingParams::builder()
.booster(BoosterKind::GbLinear)
.multi_strategy(strategy)
.objective(sketch.clone())
.build()
.unwrap();
assert_eq!(invalid_param(train(&linear, &dtrain, 1)), "objective");
}
let monotone = vector_params()
.monotone_constraints(vec![Monotone::Increasing])
.objective(sketch)
.build()
.unwrap();
assert_eq!(
invalid_param(train(&monotone, &dtrain, 1)),
"monotone_constraints"
);
let wrong = vector_params()
.objective(Objective::custom(squared_error(K).with_split_gradient(
|_, g| Some(SplitGradient::new(g[1..].to_vec(), 1)),
)))
.build()
.unwrap();
assert!(matches!(
train(&wrong, &dtrain, 1),
Err(HessboostError::DimensionMismatch { .. })
));
}
#[test]
fn vector_forests_hold_num_parallel_tree_trees_per_iteration() {
let params = vector_params().num_parallel_tree(3).build().unwrap();
let d = dtrain();
let model = train(¶ms, &d, 4).unwrap();
assert_eq!(model.num_trees(), 12);
assert_eq!(model.trees_per_iteration(), 3);
assert_eq!(model.num_boost_rounds(), 4);
let (x, _) = data();
let all = model.predict_margin(&d, Iterations::Best).unwrap();
assert_eq!(all.as_slice(), reference_margins(&model, &x, N));
let first_two = model.predict_margin(&d, ..2).unwrap();
assert_eq!(
first_two,
model
.slice(..2, 1)
.unwrap()
.predict_margin(&d, Iterations::Best)
.unwrap()
);
assert_ne!(first_two, all);
assert_contribs_sum_to(&model.predict_contribs(&d, ..2).unwrap(), &first_two);
}
#[test]
fn continued_vector_training_matches_one_run() {
let dtrain = dtrain();
let params = vector_params().subsample(0.8).seed(11).build().unwrap();
let full = train(¶ms, &dtrain, 8).unwrap();
let first = train(¶ms, &dtrain, 5).unwrap();
let continued = Trainer::new(¶ms, &dtrain, 3)
.init_model(&first)
.train()
.unwrap()
.model;
assert_eq!(continued.num_trees(), 8);
assert_eq!(
continued.predict_margin(&dtrain, Iterations::Best).unwrap(),
full.predict_margin(&dtrain, Iterations::Best).unwrap()
);
}
#[test]
fn unsupported_vector_layouts_are_rejected() {
let dtrain = dtrain();
let vector = model();
let refresh = vector_params()
.process_type(ProcessType::Update(Refresh::default()))
.build()
.unwrap();
assert_eq!(
incompatible_model(
Trainer::new(&refresh, &dtrain, 2)
.init_model(&vector)
.train()
),
"process_type"
);
let scalar = TrainingParams::builder().max_depth(4).build().unwrap();
assert_eq!(
incompatible_model(
Trainer::new(&scalar, &dtrain, 2)
.init_model(&vector)
.train()
),
"multi_strategy"
);
let unchecked = |edit: fn(&mut TrainingParams)| {
let mut params = vector_params().build().unwrap();
edit(&mut params);
params
};
for (params, name) in [
(
unchecked(|p| p.grow_policy = GrowPolicy::Symmetric),
"grow_policy",
),
(
unchecked(|p| p.toad_penalty_feature = 0.1),
"toad_penalty_feature",
),
] {
assert_eq!(invalid_param(train(¶ms, &dtrain, 1)), name);
}
assert!(matches!(
vector.to_compact_bytes(),
Err(HessboostError::ModelFormat(_))
));
}
fn categorical_two_targets(x: &[f32], cols: usize, y: &[f32]) -> DMatrix {
let mut types = vec![FeatureType::Numerical; cols];
types[0] = FeatureType::Categorical;
let labels: Vec<f32> = y.iter().flat_map(|&v| [v, v]).collect();
DMatrix::from_dense(x, y.len(), cols)
.unwrap()
.with_label_matrix(&labels, 2)
.unwrap()
.with_feature_types(&types)
.unwrap()
}
fn one_round_margins(params: hessboost::config::TrainingParamsBuilder, d: &DMatrix) -> Predictions {
let params = params
.multi_strategy(MultiStrategy::MultiOutputTree)
.lambda(0.0)
.eta(1.0)
.base_score(0.0)
.build()
.unwrap();
train(¶ms, d, 1)
.unwrap()
.predict_margin(d, Iterations::Best)
.unwrap()
}
fn assert_margins(got: &Predictions, want: &[f32]) {
assert_eq!((got.n_rows(), got.width()), (want.len(), 2));
for (r, (row, &w)) in got.rows().zip(want).enumerate() {
for g in row {
assert!((g - w).abs() < 1e-4, "row {r}: {got:?} vs {want:?}");
}
}
}
#[test]
fn low_cardinality_categories_split_one_hot() {
let nan = f32::NAN;
for (x, y) in [
(vec![0.0, 0.0, nan, nan], vec![0.0, 0.0, 2.0, 2.0]),
(
vec![0.0, 0.0, 1.0, 1.0, 2.0, 2.0],
vec![-5.0, -5.0, 10.0, 10.0, -5.0, -5.0],
),
(
vec![0.0, 0.0, 1.0, 1.0, 2.0, 2.0, nan, nan],
vec![-10.0, -10.0, 10.0, 10.0, -10.0, -10.0, 10.0, 10.0],
),
] {
let d = categorical_two_targets(&x, 1, &y);
let margins = one_round_margins(TrainingParams::builder().max_depth(1), &d);
assert_margins(&margins, &y);
}
}
#[test]
fn categorical_children_keep_xgboost_priority() {
let (mut x, mut y, mut want) = (Vec::new(), Vec::new(), Vec::new());
for c in 0..4 {
for b in 0..2 {
let v = if c < 2 { -10.0 } else { 10.0 } + 2.0 * b as f32 - 1.0;
x.extend([c as f32, b as f32]);
y.push(v);
want.push(if c < 2 { -10.0 } else { v });
}
}
let d = categorical_two_targets(&x, 2, &y);
for policy in [GrowPolicy::DepthWise, GrowPolicy::LossGuide] {
let params = TrainingParams::builder()
.grow_policy(policy)
.max_depth(2)
.max_leaves(3);
assert_margins(&one_round_margins(params, &d), &want);
}
}