use hessboost::config::{
BalancedBagging, BoosterKind, Dart, ProcessType, QueryBagging, Refresh, SamplingMethod,
TrainingParamsBuilder,
};
use hessboost::objective::{LambdaRank, RegLoss};
use hessboost::prelude::*;
use hessboost::tree::RegTree;
mod common;
use common::{invalid_data, invalid_param, labeled_dense, lcg};
fn dataset(n: usize, f: usize) -> DMatrix {
let mut next = lcg(0x1234_5678);
let mut x = Vec::with_capacity(n * f);
let mut y = Vec::with_capacity(n);
for row in 0..n {
let mut target = 0.0;
for col in 0..f {
let v = next();
target += v * (col as f32 + 1.0);
x.push(v);
}
if row % 37 == 0 {
target += 10.0;
}
y.push(target);
}
labeled_dense(&x, f, &y)
}
fn mvs_params(seed: u64) -> TrainingParams {
TrainingParams::builder()
.tree_method(TreeMethod::Hist)
.sampling_method(SamplingMethod::GradientBased)
.subsample(0.3)
.max_depth(4)
.seed(seed)
.build()
.unwrap()
}
fn predictions(params: &TrainingParams, data: &DMatrix, rounds: usize) -> Vec<f32> {
train(params, data, rounds)
.unwrap()
.predict(data, Iterations::Best)
.unwrap()
.into_vec()
}
fn path_features(tree: &RegTree) -> Vec<Vec<u32>> {
let mut out = Vec::new();
let mut stack = vec![(0usize, Vec::new())];
while let Some((nid, mut path)) = stack.pop() {
let node = tree.node(nid);
if node.is_leaf() {
out.push(path);
continue;
}
path.push(node.split_feature);
stack.push((node.left as usize, path.clone()));
stack.push((node.right as usize, path));
}
out
}
fn split_features(model: &BoostedModel) -> Vec<u32> {
let mut all: Vec<u32> = model
.trees()
.iter()
.flat_map(|t| path_features(t).into_iter().flatten())
.collect();
all.sort_unstable();
all.dedup();
all
}
#[test]
fn gradient_based_sampling_is_seeded_and_thread_count_independent() {
let data = dataset(6000, 5);
let run = |threads: usize, seed: u64| {
common::with_threads(threads, || predictions(&mvs_params(seed), &data, 8))
};
let serial = run(1, 3);
assert_eq!(serial, run(4, 3), "serial and parallel training differ");
assert_ne!(serial, run(1, 4), "the seed drives the sample");
let full = predictions(
&TrainingParams::builder().max_depth(4).build().unwrap(),
&data,
8,
);
assert_ne!(serial, full, "gradient-based sampling had no effect");
}
#[test]
fn gradient_based_at_full_subsample_is_the_unsampled_model() {
let data = dataset(500, 3);
let base = || TrainingParams::builder().max_depth(3).colsample_bynode(0.7);
let uniform = predictions(&base().build().unwrap(), &data, 5);
let mvs = predictions(
&base()
.sampling_method(SamplingMethod::GradientBased)
.build()
.unwrap(),
&data,
5,
);
assert_eq!(uniform, mvs);
}
#[test]
fn gradient_based_sampling_tree_method_support() {
let data = dataset(300, 3);
let with = |method: TreeMethod, subsample: f64| {
TrainingParams::builder()
.tree_method(method)
.sampling_method(SamplingMethod::GradientBased)
.subsample(subsample)
.build()
.unwrap()
};
assert_eq!(
invalid_param(train(&with(TreeMethod::Exact, 0.5), &data, 2)),
"sampling_method"
);
assert!(train(&with(TreeMethod::Exact, 1.0), &data, 2).is_ok());
for method in [TreeMethod::Hist, TreeMethod::Approx, TreeMethod::Auto] {
let model = train(&with(method, 0.4), &data, 3).unwrap();
assert_eq!(model.num_trees(), 3, "{method:?}");
}
let dart = TrainingParams::builder()
.booster(BoosterKind::Dart(Dart::default()))
.sampling_method(SamplingMethod::GradientBased)
.subsample(0.4)
.build()
.unwrap();
assert!(train(&dart, &data, 3).is_ok());
}
#[test]
fn bagging_by_query_is_seeded_and_changes_the_trees() {
let n_groups = 30;
let group_size = 4;
let n = n_groups * group_size;
let mut x = Vec::with_capacity(n);
let mut y = Vec::with_capacity(n);
for group in 0..n_groups {
for doc in 0..group_size {
x.push(doc as f32 + group as f32 * 0.001);
y.push(doc as f32);
}
}
let data = labeled_dense(&x, 1, &y)
.with_group_sizes(&vec![group_size; n_groups])
.unwrap();
let params = TrainingParams::builder()
.objective(Objective::RankNdcg(LambdaRank::default()))
.tree_method(TreeMethod::Hist)
.bagging_by_query(QueryBagging::new(0.5).unwrap())
.seed(22)
.max_depth(2)
.build()
.unwrap();
let first = train(¶ms, &data, 4).unwrap().trees().to_vec();
let second = train(¶ms, &data, 4).unwrap().trees().to_vec();
assert_eq!(first, second);
for tree in &first {
assert!(tree.nodes().iter().any(|node| !node.is_leaf()));
}
let mut unbagged = params;
unbagged.bagging_by_query = None;
assert_ne!(first, train(&unbagged, &data, 4).unwrap().trees().to_vec());
}
#[test]
fn bagging_by_query_refuses_non_ranking_or_incompatible_sampling() {
let bagging = QueryBagging::new(0.5).unwrap();
let ranking = || {
TrainingParams::builder()
.objective(Objective::RankNdcg(LambdaRank::default()))
.bagging_by_query(bagging)
};
for (params, name) in [
(ranking().subsample(0.8), "subsample"),
(
ranking().sampling_method(SamplingMethod::GradientBased),
"sampling_method",
),
(ranking().booster(BoosterKind::GbLinear), "bagging_by_query"),
(
TrainingParams::builder().bagging_by_query(bagging),
"bagging_by_query",
),
(
ranking().balanced_bagging(BalancedBagging::new(0.5, 1.0).unwrap()),
"bagging_by_query",
),
] {
assert_eq!(invalid_param(params.build()), name);
}
let ungrouped = labeled_dense(&[0.0, 1.0], 1, &[0.0, 1.0]);
assert_eq!(
invalid_data(train(&ranking().build().unwrap(), &ungrouped, 1)),
("group_sizes", None)
);
}
#[test]
fn balanced_bagging_changes_binary_training_deterministically() {
let (n_pos, n_neg) = (200usize, 800usize);
let n = n_pos + n_neg;
let mut x = Vec::with_capacity(n);
let mut y = Vec::with_capacity(n);
for i in 0..n {
x.push((i % 100) as f32 / 100.0);
y.push(if i < n_pos { 1.0 } else { 0.0 });
}
let data = labeled_dense(&x, 1, &y);
let base = || {
TrainingParams::builder()
.objective(binary())
.tree_method(TreeMethod::Hist)
.max_depth(2)
.seed(53)
};
let all = train(&base().build().unwrap(), &data, 3)
.unwrap()
.trees()
.to_vec();
let balanced = base()
.balanced_bagging(BalancedBagging::new(0.6, 0.1).unwrap())
.build()
.unwrap();
let sampled = train(&balanced, &data, 3).unwrap().trees().to_vec();
assert_ne!(all, sampled);
assert_eq!(sampled, train(&balanced, &data, 3).unwrap().trees());
}
#[test]
fn balanced_bagging_refuses_unsupported_parameters_and_labels() {
let bagging = BalancedBagging::new(0.5, 1.0).unwrap();
let builder = || TrainingParams::builder().balanced_bagging(bagging);
for (params, name) in [
(builder(), "pos_bagging_fraction"),
(
builder()
.objective(binary())
.sampling_method(SamplingMethod::GradientBased),
"sampling_method",
),
(builder().objective(binary()).subsample(0.8), "subsample"),
(
builder().objective(binary()).booster(BoosterKind::GbLinear),
"pos_bagging_fraction",
),
] {
assert_eq!(invalid_param(params.build()), name);
}
let binary = builder().objective(binary()).build().unwrap();
let multi = DMatrix::from_dense(&[0.0, 1.0, 1.0, 0.0], 2, 2)
.unwrap()
.with_label_matrix(&[0.0, 1.0, 1.0, 0.0], 2)
.unwrap();
assert_eq!(invalid_data(train(&binary, &multi, 1)), ("labels", None));
let graded = labeled_dense(&[0.0, 1.0], 1, &[0.0, 0.5]);
assert_eq!(invalid_data(train(&binary, &graded, 1)), ("labels", None));
}
fn binary() -> Objective {
Objective::BinaryLogistic(RegLoss::default())
}
#[test]
fn zero_weight_features_are_practically_never_split_on() {
let data = dataset(800, 4)
.with_feature_weights(&[0.0, 3.0, 0.0, 1.0])
.unwrap();
let base = |method: TreeMethod| {
TrainingParams::builder()
.tree_method(method)
.max_depth(4)
.seed(11)
};
for method in [TreeMethod::Hist, TreeMethod::Approx, TreeMethod::Exact] {
for params in [
base(method).colsample_bytree(0.5).build().unwrap(),
base(method).colsample_bylevel(0.5).build().unwrap(),
base(method).colsample_bynode(0.5).build().unwrap(),
] {
let model = train(¶ms, &data, 10).unwrap();
assert_eq!(split_features(&model), vec![1, 3], "{method:?}");
}
}
let unweighted = dataset(800, 4);
let model = train(
&base(TreeMethod::Hist)
.colsample_bynode(0.5)
.build()
.unwrap(),
&unweighted,
10,
)
.unwrap();
assert_eq!(split_features(&model), vec![0, 1, 2, 3]);
}
#[test]
fn tree_feature_shares_follow_weights() {
let weights = [1.0f32, 2.0, 5.0];
let data = dataset(400, 3).with_feature_weights(&weights).unwrap();
let params = TrainingParams::builder()
.colsample_bytree(0.34) .max_depth(2)
.eta(0.01)
.seed(5)
.build()
.unwrap();
let model = train(¶ms, &data, 800).unwrap();
let mut counts = [0usize; 3];
for tree in model.trees() {
let used: Vec<u32> = path_features(tree).into_iter().flatten().collect();
assert!(
used.windows(2).all(|w| w[0] == w[1]),
"one feature per tree"
);
if let Some(&f) = used.first() {
counts[f as usize] += 1;
}
}
let total: usize = counts.iter().sum();
assert!(total > 700, "trees must split");
for (f, &count) in counts.iter().enumerate() {
let share = count as f64 / total as f64;
let expected = f64::from(weights[f]) / 8.0;
assert!(
(share - expected).abs() < 0.06,
"feature {f}: {share} vs {expected}"
);
}
}
#[test]
fn interaction_constraints_filter_the_weighted_sample() {
let data = dataset(1500, 4)
.with_feature_weights(&[1.0, 1.0, 0.0, 0.0])
.unwrap();
for method in [TreeMethod::Hist, TreeMethod::Exact] {
let params = TrainingParams::builder()
.tree_method(method)
.colsample_bynode(0.5)
.interaction_constraints(vec![vec![0, 2], vec![1, 3]])
.max_depth(4)
.seed(2)
.build()
.unwrap();
let model = train(¶ms, &data, 10).unwrap();
let mut used_any = [false; 4];
for tree in model.trees() {
for path in path_features(tree) {
assert!(
path.windows(2).all(|w| w[0] == w[1]),
"{method:?}: {path:?}"
);
for f in path {
used_any[f as usize] = true;
}
}
}
assert_eq!(used_any, [true, true, false, false], "{method:?}");
}
}
#[test]
fn weighted_column_sampling_is_seeded() {
let data = dataset(600, 6)
.with_feature_weights(&[0.5, 1.0, 1.5, 2.0, 2.5, 3.0])
.unwrap();
let params = |seed: u64| {
TrainingParams::builder()
.colsample_bytree(0.8)
.colsample_bylevel(0.8)
.colsample_bynode(0.6)
.max_depth(4)
.seed(seed)
.build()
.unwrap()
};
let a = predictions(¶ms(1), &data, 6);
assert_eq!(a, predictions(¶ms(1), &data, 6));
assert_ne!(a, predictions(¶ms(2), &data, 6));
}
fn step_rows(feature: impl Fn(usize) -> f32) -> DMatrix {
let x: Vec<f32> = (0..8).map(feature).collect();
let y: Vec<f32> = (0..8).map(|i| if i < 4 { 0.0 } else { 1.0 }).collect();
labeled_dense(&x, 1, &y)
}
fn plain(method: TreeMethod, booster: BoosterKind) -> TrainingParamsBuilder {
TrainingParams::builder()
.tree_method(method)
.booster(booster)
.base_score(0.0)
.eta(1.0)
.lambda(0.0)
.min_child_weight(0.0)
}
#[test]
fn approx_parallel_trees_share_one_row_sample() {
let data = step_rows(|_| 0.0);
for booster in [BoosterKind::GbTree, BoosterKind::Dart(Dart::default())] {
for (sampling, subsample) in [
(SamplingMethod::GradientBased, 0.25),
(SamplingMethod::Uniform, 0.5),
] {
let leaves = |method: TreeMethod| {
let params = plain(method, booster)
.sampling_method(sampling)
.subsample(subsample)
.num_parallel_tree(4)
.seed(13)
.build()
.unwrap();
let model = train(¶ms, &data, 3).unwrap();
let leaves: Vec<f32> = model.trees().iter().map(|t| t.node(0).leaf_value).collect();
assert_eq!(leaves.len(), 12);
leaves
};
let differs = |leaves: &[f32]| {
leaves
.chunks(4)
.any(|forest| forest.iter().any(|&v| v != forest[0]))
};
let approx = leaves(TreeMethod::Approx);
assert!(!differs(&approx), "{booster:?} {sampling:?}: {approx:?}");
let hist = leaves(TreeMethod::Hist);
assert!(differs(&hist), "{booster:?} {sampling:?}: {hist:?}");
}
}
}
#[test]
fn approx_gradient_sampling_continuation_matches_uninterrupted_training() {
let step = step_rows(|i| i as f32);
let wide = dataset(400, 3);
for booster in [BoosterKind::GbTree, BoosterKind::Dart(Dart::default())] {
for (data, depth, split) in [(&step, 1, [1, 1]), (&wide, 3, [2, 3])] {
let params = plain(TreeMethod::Approx, booster)
.sampling_method(SamplingMethod::GradientBased)
.subsample(0.25)
.max_depth(depth)
.seed(13)
.build()
.unwrap();
let full = train(¶ms, data, split[0] + split[1]).unwrap();
let head = train(¶ms, data, split[0]).unwrap();
let resumed = Trainer::new(¶ms, data, split[1])
.init_model(&head)
.train()
.unwrap()
.model;
assert_eq!(
full.predict(data, Iterations::Best).unwrap(),
resumed.predict(data, Iterations::Best).unwrap(),
"{booster:?} depth {depth}"
);
}
}
}
#[test]
fn gradient_based_sampling_survives_overflowing_squares() {
let data = labeled_dense(&[0.0; 8], 1, &[1e20; 8]);
let params = plain(TreeMethod::Hist, BoosterKind::GbTree)
.sampling_method(SamplingMethod::GradientBased)
.subsample(0.5)
.seed(3)
.build()
.unwrap();
let preds = predictions(¶ms, &data, 1);
assert!(
preds.iter().all(|&p| p > 1e19 && p.is_finite()),
"{preds:?}"
);
}
#[test]
fn feature_weights_are_refused_where_columns_are_not_sampled() {
let data = dataset(200, 3);
let weighted = data.clone().with_feature_weights(&[1.0, 2.0, 3.0]).unwrap();
let linear = TrainingParams::builder()
.booster(BoosterKind::GbLinear)
.build()
.unwrap();
assert_eq!(
invalid_data(train(&linear, &weighted, 2)),
("feature_weights", None)
);
let model = train(&TrainingParams::default(), &data, 2).unwrap();
let update = TrainingParams::builder()
.process_type(ProcessType::Update(Refresh::default()))
.build()
.unwrap();
let refresh = |data: &DMatrix| Trainer::new(&update, data, 2).init_model(&model).train();
assert_eq!(invalid_data(refresh(&weighted)), ("feature_weights", None));
assert!(refresh(&data).is_ok());
}