mod search;
use super::lightgbm::{SplitOptions, finalize_smoothed_leaves};
use super::partition::{child_histograms, partition_rows};
use super::shared::{
BuilderConfig, InteractionState, LeafRows, finalize_leaf_values, rayon_available, sum_rows,
xgb_calc_weight,
};
use super::{BELOW_ALL_VALUES, BestSplit, SplitLocation, limit_or_unbounded};
use crate::config::{GrowPolicy, TrainingParams};
use crate::data::ghist::GHistIndex;
use crate::data::quantile::HistCuts;
use crate::objective::GradPair;
use crate::tree::constraints::Bounds;
use crate::tree::gain::GradStats;
use crate::tree::hist::quantized::QuantNode;
use crate::tree::hist::{CpuBackend, Histogram, HistogramBackend, zeroed};
use crate::tree::regtree::RegTree;
use crate::tree::reuse::{HistReuse, ReuseSet};
use crate::tree::sampler::{ColumnSampler, FeatureSet};
use rayon::prelude::*;
use std::cmp::Ordering;
use std::collections::{BinaryHeap, HashMap};
const SPECULATE_NODES: usize = 8;
const PARALLEL_EVALUATE_ROWS: usize = 4096;
pub(super) const PARALLEL_FRONTIER_ROWS: usize = 4096;
#[derive(Debug, Clone, Copy)]
pub(super) struct NodeCtx {
pub(super) id: usize,
pub(super) stats: GradStats,
pub(super) bounds: Bounds,
pub(super) rows: usize,
pub(super) output: f64,
pub(super) tree_seed: u64,
}
struct NodeEntry {
nid: usize,
depth: usize,
rows: Vec<u32>,
hist: Histogram,
best: BestSplit,
bounds: Bounds,
allowed: Option<InteractionState>,
tree_seed: u64,
quant: Option<QuantNode>,
}
struct PendingSplit {
entry: NodeEntry,
left_id: usize,
right_id: usize,
left_bounds: Bounds,
right_bounds: Bounds,
left_features: FeatureSet,
right_features: FeatureSet,
terminal: bool,
}
pub(super) static CPU_BACKEND: CpuBackend = CpuBackend;
impl PartialEq for NodeEntry {
fn eq(&self, other: &Self) -> bool {
self.best.loss_chg == other.best.loss_chg
}
}
impl Eq for NodeEntry {}
impl PartialOrd for NodeEntry {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
impl Ord for NodeEntry {
fn cmp(&self, other: &Self) -> Ordering {
self.best.loss_chg.total_cmp(&other.best.loss_chg)
}
}
pub struct HistTreeBuilder<'a> {
config: BuilderConfig<'a>,
backend: &'a dyn HistogramBackend,
options: Option<SplitOptions>,
reuse: Option<HistReuse>,
rounding_seed: u64,
}
impl<'a> HistTreeBuilder<'a> {
pub fn new(params: &'a TrainingParams) -> Self {
HistTreeBuilder {
config: BuilderConfig::new(params),
backend: &CPU_BACKEND,
options: SplitOptions::from_params(params),
reuse: None,
rounding_seed: 0,
}
}
#[must_use]
pub(crate) fn with_reuse(mut self, set: Option<&ReuseSet>, cuts: &HistCuts) -> Self {
self.reuse = set.map(|set| HistReuse::new(set, cuts, BELOW_ALL_VALUES));
self
}
#[must_use]
pub(crate) fn with_rounding_seed(mut self, seed: u64) -> Self {
self.rounding_seed = seed;
self
}
#[must_use]
pub(crate) fn with_backend(mut self, backend: &'a dyn HistogramBackend) -> Self {
self.backend = backend;
self
}
pub fn build(
&self,
ghist: &GHistIndex,
gpair: &[GradPair],
row_subset: &[u32],
sampler: &mut ColumnSampler,
) -> RegTree {
self.build_inner(ghist, gpair, row_subset, sampler, false).0
}
pub(crate) fn build_with_leaf_rows(
&self,
ghist: &GHistIndex,
gpair: &[GradPair],
row_subset: &[u32],
sampler: &mut ColumnSampler,
) -> (RegTree, Vec<LeafRows>) {
self.build_inner(ghist, gpair, row_subset, sampler, true)
}
fn build_inner(
&self,
ghist: &GHistIndex,
gpair: &[GradPair],
row_subset: &[u32],
sampler: &mut ColumnSampler,
capture_rows: bool,
) -> (RegTree, Vec<LeafRows>) {
self.backend.prepare(ghist, gpair);
if self.config.params.grow_policy == GrowPolicy::Symmetric {
return super::oblivious::SymmetricTreeBuilder::new(&self.config, self.backend).build(
ghist,
gpair,
row_subset,
sampler,
capture_rows,
);
}
debug_assert!(
self.reuse
.as_ref()
.is_none_or(|r| r.n_bins() == ghist.total_bins())
);
let renew = self.config.params.quantized.is_some_and(|q| q.renew_leaf());
let (root, root_stats) = self.root(ghist, gpair, row_subset, sampler);
let mut tree = RegTree::with_root(root_stats.hess as f32);
let mut store = NodeStore {
stats: vec![root_stats],
bounds: vec![Bounds::default()],
leaf_rows: (capture_rows || renew).then(Vec::new),
};
match self.config.params.grow_policy {
GrowPolicy::DepthWise => {
self.grow_depthwise(&mut tree, &mut store, ghist, gpair, sampler, root);
}
GrowPolicy::LossGuide => {
self.grow_lossguide(&mut tree, &mut store, ghist, gpair, sampler, root);
}
GrowPolicy::Symmetric => unreachable!("symmetric trees return above"),
}
if renew && let Some(leaves) = &store.leaf_rows {
for leaf in leaves {
store.stats[leaf.node] = sum_rows(gpair, &leaf.rows);
}
}
match &self.options {
Some(options) if options.smoothing() => {
finalize_smoothed_leaves(&mut tree, root_stats, &self.config.reg);
}
_ => finalize_leaf_values(&mut tree, &store.stats, &store.bounds, &self.config.reg),
}
let leaf_rows = if capture_rows {
store.leaf_rows.unwrap_or_default()
} else {
Vec::new()
};
(tree, leaf_rows)
}
fn root(
&self,
ghist: &GHistIndex,
gpair: &[GradPair],
row_subset: &[u32],
sampler: &mut ColumnSampler,
) -> (NodeEntry, GradStats) {
let total_bins = ghist.total_bins();
let (root_stats, root_hist, root_quant) =
if let Some(quantized) = self.config.params.quantized {
let (quant, stats, hist) =
QuantNode::root(ghist, gpair, row_subset, quantized, self.rounding_seed);
(stats, hist, Some(quant))
} else {
let build_hist = || {
let mut root_hist = zeroed(total_bins);
self.backend.build(ghist, row_subset, gpair, &mut root_hist);
root_hist
};
let (root_stats, root_hist) = if rayon_available() {
rayon::join(|| sum_rows(gpair, row_subset), build_hist)
} else {
(sum_rows(gpair, row_subset), build_hist())
};
(root_stats, root_hist, None)
};
let root_feats = sampler.sample(0);
let tree_seed = sampler.seed();
let root_ctx = NodeCtx {
id: 0,
stats: root_stats,
bounds: Bounds::default(),
rows: row_subset.len(),
output: xgb_calc_weight(root_stats, &self.config.reg),
tree_seed,
};
let best = self.evaluate(ghist, &root_hist, &root_feats, None, root_ctx);
let root = NodeEntry {
nid: 0,
depth: 0,
rows: row_subset.to_vec(),
hist: root_hist,
best,
bounds: Bounds::default(),
allowed: None,
tree_seed,
quant: root_quant,
};
(root, root_stats)
}
fn grow_depthwise(
&self,
tree: &mut RegTree,
store: &mut NodeStore,
ghist: &GHistIndex,
gpair: &[GradPair],
sampler: &mut ColumnSampler,
root: NodeEntry,
) {
let limit = limit_or_unbounded(self.config.params.max_depth);
let mut frontier = vec![root];
let mut depth = 0;
while depth < limit && !frontier.is_empty() {
let parallel = frontier.len() > 1
&& frontier.iter().map(|entry| entry.rows.len()).sum::<usize>()
>= PARALLEL_FRONTIER_ROWS
&& rayon_available();
let mut pending = Vec::with_capacity(frontier.len());
for entry in frontier.drain(..) {
if self.valid(&entry.best) {
if let Some(split) =
self.prepare_split(tree, store, ghist.cuts(), sampler, entry)
{
pending.push(split);
}
} else {
store.record_leaf(entry);
}
}
let build = |split| self.build_children(ghist, gpair, split);
let children: Vec<_> = if parallel {
pending.into_par_iter().map(build).collect()
} else {
pending.into_iter().map(build).collect()
};
frontier = children
.into_iter()
.flat_map(|(left, right)| [left, right])
.collect();
depth += 1;
}
for entry in frontier {
store.record_leaf(entry);
}
}
fn grow_lossguide(
&self,
tree: &mut RegTree,
store: &mut NodeStore,
ghist: &GHistIndex,
gpair: &[GradPair],
sampler: &mut ColumnSampler,
root: NodeEntry,
) {
let limit = limit_or_unbounded(self.config.params.max_depth);
let max_leaves = limit_or_unbounded(self.config.params.max_leaves);
let expandable = |entry: &NodeEntry| entry.depth < limit && self.valid(&entry.best);
let speculative = self.speculative_features(sampler);
let mut ready: HashMap<usize, (NodeEntry, NodeEntry)> = HashMap::new();
let mut heap = BinaryHeap::new();
heap.push(root);
let mut n_leaves = 1usize;
while let Some(entry) = heap.pop() {
if n_leaves >= max_leaves {
store.record_leaf(entry);
break;
}
if !expandable(&entry) {
store.record_leaf(entry);
continue; }
let children = if let Some(features) = &speculative {
if !ready.contains_key(&entry.nid) {
let budget = (max_leaves - n_leaves).min(SPECULATE_NODES);
let queued = heap
.iter()
.filter(|e| expandable(e) && !ready.contains_key(&e.nid));
let batch = speculation_batch(&entry, queued, budget);
let built: Vec<_> = batch
.par_iter()
.map(|e| (e.nid, self.speculate_children(ghist, gpair, e, features)))
.collect();
ready.extend(built);
}
let built = ready.remove(&entry.nid);
self.prepare_split(tree, store, ghist.cuts(), sampler, entry)
.map(|split| match built {
Some((mut left, mut right)) => {
left.nid = split.left_id;
right.nid = split.right_id;
(left, right)
}
None => self.build_children(ghist, gpair, split),
})
} else {
self.prepare_split(tree, store, ghist.cuts(), sampler, entry)
.map(|split| self.build_children(ghist, gpair, split))
};
n_leaves += 1; if let Some((l, r)) = children {
heap.push(l);
heap.push(r);
}
}
for entry in heap {
store.record_leaf(entry);
}
}
fn speculative_features(&self, sampler: &ColumnSampler) -> Option<FeatureSet> {
if self.reuse.is_some()
|| self.options.is_some()
|| self.config.params.quantized.is_some()
|| !rayon_available()
{
return None;
}
sampler.fixed_features()
}
fn speculate_children(
&self,
ghist: &GHistIndex,
gpair: &[GradPair],
entry: &NodeEntry,
features: &FeatureSet,
) -> (NodeEntry, NodeEntry) {
let b = &entry.best;
let (left_bounds, right_bounds) =
b.child_bounds(entry.bounds, self.config.cons.dir(b.feature as usize));
let split = PendingSplit {
entry: NodeEntry {
nid: entry.nid,
depth: entry.depth,
rows: entry.rows.clone(),
hist: entry.hist.clone(),
best: entry.best.clone(),
bounds: entry.bounds,
allowed: entry.allowed.clone(),
tree_seed: entry.tree_seed,
quant: None,
},
left_id: 0,
right_id: 0,
left_bounds,
right_bounds,
left_features: features.clone(),
right_features: features.clone(),
terminal: false,
};
self.build_children(ghist, gpair, split)
}
fn valid(&self, best: &BestSplit) -> bool {
best.valid(self.config.params.gamma, self.config.reg.min_child_weight)
}
fn prepare_split(
&self,
tree: &mut RegTree,
store: &mut NodeStore,
cuts: &HistCuts,
sampler: &mut ColumnSampler,
entry: NodeEntry,
) -> Option<PendingSplit> {
let b = &entry.best;
let dir = self.config.cons.dir(b.feature as usize);
let (lb_bounds, rb_bounds) = b.child_bounds(entry.bounds, dir);
let (left_id, right_id) = b.expand(tree, entry.nid, b.route().rule(cuts));
if let Some(reuse) = &self.reuse {
match &b.location {
SplitLocation::Categories(categories) => {
reuse.commit_categorical(b.feature, categories);
}
SplitLocation::Numeric(pos) => reuse.commit_numeric(b.feature, pos.bin()),
}
}
debug_assert_eq!(left_id, store.stats.len());
store.push(b.left, lb_bounds);
store.push(b.right, rb_bounds);
let child_depth = entry.depth + 1;
let terminal = self.config.params.grow_policy == GrowPolicy::DepthWise
&& child_depth >= limit_or_unbounded(self.config.params.max_depth);
if terminal && store.leaf_rows.is_none() {
sampler.sample(child_depth);
sampler.sample(child_depth);
return None;
}
let left_features = sampler.sample(child_depth);
let right_features = sampler.sample(child_depth);
Some(PendingSplit {
entry,
left_id,
right_id,
left_bounds: lb_bounds,
right_bounds: rb_bounds,
left_features,
right_features,
terminal,
})
}
fn build_children(
&self,
ghist: &GHistIndex,
gpair: &[GradPair],
split: PendingSplit,
) -> (NodeEntry, NodeEntry) {
let PendingSplit {
entry,
left_id,
right_id,
left_bounds: lb_bounds,
right_bounds: rb_bounds,
left_features,
right_features,
terminal,
} = split;
let NodeEntry {
depth: parent_depth,
rows: parent_rows,
hist: parent_hist,
best,
allowed: parent_allowed,
tree_seed,
quant: parent_quant,
..
} = entry;
let b = &best;
let (left_rows, right_rows) = partition_rows(ghist, &parent_rows, b.route());
drop(parent_rows);
let (mut left_quant, mut right_quant) = (None, None);
let (left_hist, right_hist) = if terminal {
(Vec::new(), Vec::new())
} else if let Some(quant) = parent_quant {
let ((lq, lh), (rq, rh)) = quant.children(ghist, &left_rows, &right_rows, parent_hist);
(left_quant, right_quant) = (Some(lq), Some(rq));
(lh, rh)
} else {
child_histograms(
self.backend,
ghist,
gpair,
&left_rows,
&right_rows,
parent_hist,
)
};
let child_allowed = self.config.next_allowed(parent_allowed.as_ref(), b.feature);
let (left_best, right_best) = if terminal {
(BestSplit::none(), BestSplit::none())
} else {
let allowed = child_allowed.as_ref();
let left_ctx = NodeCtx {
id: left_id,
stats: b.left,
bounds: lb_bounds,
rows: left_rows.len(),
output: b.w_left,
tree_seed,
};
let right_ctx = NodeCtx {
id: right_id,
stats: b.right,
bounds: rb_bounds,
rows: right_rows.len(),
output: b.w_right,
tree_seed,
};
let left = || self.evaluate(ghist, &left_hist, &left_features, allowed, left_ctx);
let right = || self.evaluate(ghist, &right_hist, &right_features, allowed, right_ctx);
if left_rows.len() + right_rows.len() >= PARALLEL_EVALUATE_ROWS && rayon_available() {
rayon::join(left, right)
} else {
(left(), right())
}
};
let left = NodeEntry {
nid: left_id,
depth: parent_depth + 1,
rows: left_rows,
hist: left_hist,
best: left_best,
bounds: lb_bounds,
allowed: child_allowed.clone(),
tree_seed,
quant: left_quant,
};
let right = NodeEntry {
nid: right_id,
depth: parent_depth + 1,
rows: right_rows,
hist: right_hist,
best: right_best,
bounds: rb_bounds,
allowed: child_allowed,
tree_seed,
quant: right_quant,
};
(left, right)
}
}
fn speculation_batch<'e>(
entry: &'e NodeEntry,
queued: impl Iterator<Item = &'e NodeEntry>,
budget: usize,
) -> Vec<&'e NodeEntry> {
let mut queued: Vec<&NodeEntry> = queued.collect();
queued.sort_by(|a, b| b.cmp(a));
std::iter::once(entry)
.chain(queued.into_iter().take(budget - 1))
.collect()
}
struct NodeStore {
stats: Vec<GradStats>,
bounds: Vec<Bounds>,
leaf_rows: Option<Vec<LeafRows>>,
}
impl NodeStore {
fn record_leaf(&mut self, entry: NodeEntry) {
if let Some(leaves) = &mut self.leaf_rows {
leaves.push(LeafRows {
node: entry.nid,
rows: entry.rows,
});
}
}
#[inline]
fn push(&mut self, stats: GradStats, bounds: Bounds) {
self.stats.push(stats);
self.bounds.push(bounds);
}
}
#[cfg(test)]
mod tests {
use super::super::test_support::{
binned, gp, grow_exact, grow_hist, monotone_v_shape_data, non_decreasing,
};
use super::*;
use crate::config::TrainingParams;
use crate::data::DMatrix;
use crate::tree::builder::all_rows;
#[test]
fn splits_on_separating_feature() {
let x = vec![0.0f32, 0.0, 1.0, 1.0];
let data = DMatrix::from_dense(&x, 4, 1).unwrap();
let ghist = binned(&data, 256);
let gpair = vec![gp(1.0, 1.0), gp(1.0, 1.0), gp(-1.0, 1.0), gp(-1.0, 1.0)];
let params = TrainingParams::builder()
.max_depth(1)
.lambda(0.0)
.min_child_weight(0.0)
.gamma(0.0)
.build()
.unwrap();
let tree = grow_hist(¶ms, &ghist, &gpair);
assert_eq!(tree.num_nodes(), 3);
assert!((tree.predict_row(&data, 0) - (-1.0)).abs() < 1e-6);
assert!((tree.predict_row(&data, 2) - 1.0).abs() < 1e-6);
}
#[test]
fn depth_limit_preserves_reused_column_sampler_state() {
let n = 256;
let features = 8;
let x: Vec<f32> = (0..n * features)
.map(|i| ((i / features * 17 + i % features * 29) % 97) as f32 / 97.0)
.collect();
let gradients: Vec<GradPair> = x
.chunks_exact(features)
.map(|row| gp(row.iter().sum::<f32>() - 4.0, 1.0))
.collect();
let data = DMatrix::from_dense(&x, n, features).unwrap();
let ghist = binned(&data, 64);
let rows = all_rows(n);
for depth in [1, 2, 4] {
let params = TrainingParams::builder().max_depth(depth).build().unwrap();
let builder = HistTreeBuilder::new(¶ms);
let new_sampler = || ColumnSampler::new(features, None, 1.0, 0.75, 0.75, 42);
let mut sampler = new_sampler();
let mut expected = new_sampler();
for _ in 0..3 {
let tree = builder.build(&ghist, &gradients, &rows, &mut sampler);
assert!(tree.num_nodes() > 1);
let mut node_depth = vec![0usize; tree.num_nodes()];
for (nid, node) in tree.nodes().iter().enumerate() {
expected.sample(node_depth[nid]);
if !node.is_leaf() {
node_depth[node.left as usize] = node_depth[nid] + 1;
node_depth[node.right as usize] = node_depth[nid] + 1;
}
}
for depth in 0..4 {
assert_eq!(sampler.sample(depth), expected.sample(depth));
}
}
}
}
#[test]
fn parallel_depthwise_preserves_tree_and_sampler() {
use crate::config::Monotone;
use crate::data::FeatureType;
let n = 8192;
let features = 8;
let serial = rayon::ThreadPoolBuilder::new()
.num_threads(1)
.build()
.unwrap();
let parallel = rayon::ThreadPoolBuilder::new()
.num_threads(4)
.build()
.unwrap();
for mode in ["dense", "missing", "categorical"] {
let mut state = 123u64;
let mut values = Vec::with_capacity(n * features);
let mut gradients = Vec::with_capacity(n);
for row in 0..n {
let mut target = 0.0;
for col in 0..features {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1);
let mut value = (state >> 33) as f32 / (1u32 << 31) as f32;
if mode == "categorical" && col < 2 {
value = (value * 4.0).floor();
}
target += value * (col + 1) as f32;
if mode == "missing" && (row * 13 + col * 7) % 11 < 2 {
value = f32::NAN;
}
values.push(value);
}
gradients.push(gp(18.0 - target, 1.0));
}
let mut data = DMatrix::from_dense(&values, n, features).unwrap();
if mode == "categorical" {
let mut types = vec![FeatureType::Numerical; features];
types[..2].fill(FeatureType::Categorical);
data = data.with_feature_types(&types).unwrap();
}
let ghist = binned(&data, 64);
let rows = all_rows(n);
let params = TrainingParams::builder()
.max_depth(6)
.alpha(0.1)
.monotone_constraints(vec![Monotone::None, Monotone::None, Monotone::Increasing])
.interaction_constraints(vec![vec![0, 1, 2, 3], vec![2, 4, 5, 6, 7]])
.build()
.unwrap();
let builder = HistTreeBuilder::new(¶ms);
let new_sampler = || ColumnSampler::new(features, None, 1.0, 0.75, 0.75, 91);
let mut expected_sampler = new_sampler();
let mut sampler = new_sampler();
let mut captured_sampler = new_sampler();
for _ in 0..3 {
let expected = serial
.install(|| builder.build(&ghist, &gradients, &rows, &mut expected_sampler));
let actual =
parallel.install(|| builder.build(&ghist, &gradients, &rows, &mut sampler));
let (captured, leaves) = parallel.install(|| {
builder.build_with_leaf_rows(&ghist, &gradients, &rows, &mut captured_sampler)
});
assert_eq!(captured, expected, "{mode}");
let mut seen = vec![false; n];
for leaf in leaves {
assert!(captured.node(leaf.node).is_leaf());
for row in leaf.rows {
let row = row as usize;
assert!(!seen[row]);
seen[row] = true;
assert_eq!(
leaf.node,
captured.leaf_id_with(|f| data.get(row, f as usize))
);
}
}
assert!(seen.into_iter().all(|seen| seen));
assert!(actual.num_nodes() > 7, "must exercise multiple depths");
assert_eq!(actual, expected, "{mode}");
let next = sampler.sample(1);
assert_eq!(next, expected_sampler.sample(1), "{mode}");
assert_eq!(next, captured_sampler.sample(1), "{mode}");
}
}
}
#[test]
fn parallel_lossguide_grows_the_serial_tree() {
let half = 6000;
let features = 6;
let mut state = 7u64;
let mut values = Vec::with_capacity(2 * half * features);
let mut gradients = Vec::with_capacity(2 * half);
for side in 0..2 {
for row in 0..half {
let mut target = 0.0;
values.push(side as f32);
for col in 1..features {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1);
let mut value = (state >> 33) as f32 / (1u32 << 31) as f32;
target += value * col as f32;
if (row * 13 + col * 7) % 11 < 2 {
value = f32::NAN;
}
values.push(value);
}
let g = 7.0 - target + 20.0;
gradients.push(gp(if side == 0 { g } else { -g }, 1.0));
}
}
let n = 2 * half;
let data = DMatrix::from_dense(&values, n, features).unwrap();
let ghist = binned(&data, 64);
let rows = all_rows(n);
let params = TrainingParams::builder()
.grow_policy(GrowPolicy::LossGuide)
.unlimited_depth()
.max_leaves(24)
.build()
.unwrap();
let builder = HistTreeBuilder::new(¶ms);
let grow = |threads| {
rayon::ThreadPoolBuilder::new()
.num_threads(threads)
.build()
.unwrap()
.install(|| {
builder.build_with_leaf_rows(
&ghist,
&gradients,
&rows,
&mut ColumnSampler::all(features),
)
})
};
let (expected, _) = grow(1);
assert_eq!(
expected.node(0).split_feature,
0,
"the root separates the halves"
);
assert!(expected.num_nodes() > 20);
for threads in [2, 4, 8] {
let (tree, leaves) = grow(threads);
assert_eq!(tree, expected, "{threads} threads");
let mut seen = vec![false; n];
for leaf in leaves {
for row in leaf.rows {
assert!(!std::mem::replace(&mut seen[row as usize], true));
assert_eq!(
leaf.node,
tree.leaf_id_with(|f| data.get(row as usize, f as usize))
);
}
}
assert!(seen.into_iter().all(|seen| seen));
}
}
#[test]
fn no_split_below_gamma() {
let x = vec![0.0f32, 1.0];
let data = DMatrix::from_dense(&x, 2, 1).unwrap();
let ghist = binned(&data, 256);
let gpair = vec![gp(1.0, 1.0), gp(-1.0, 1.0)];
let params = TrainingParams::builder()
.max_depth(3)
.gamma(1e9)
.build()
.unwrap();
let tree = grow_hist(¶ms, &ghist, &gpair);
assert_eq!(tree.num_nodes(), 1);
let (captured, leaves) = HistTreeBuilder::new(¶ms).build_with_leaf_rows(
&ghist,
&gpair,
&all_rows(2),
&mut ColumnSampler::all(1),
);
assert_eq!(captured, tree);
assert_eq!(leaves.len(), 1);
assert_eq!(leaves[0].node, 0);
assert_eq!(leaves[0].rows, all_rows(2));
}
#[test]
fn lossguide_respects_max_leaves() {
let n = 64;
let x: Vec<f32> = (0..n).map(|i| i as f32).collect();
let mut y = Vec::new();
for i in 0..n {
y.push(gp(if i % 2 == 0 { 1.0 } else { -1.0 }, 1.0));
}
let data = DMatrix::from_dense(&x, n, 1).unwrap();
let ghist = binned(&data, 256);
let params = TrainingParams::builder()
.grow_policy(GrowPolicy::LossGuide)
.max_leaves(4)
.unlimited_depth()
.min_child_weight(0.0)
.gamma(0.0)
.lambda(0.0)
.build()
.unwrap();
let tree = grow_hist(¶ms, &ghist, &y);
assert!(tree.num_leaves() <= 4, "got {} leaves", tree.num_leaves());
}
#[test]
fn monotone_increasing_is_enforced() {
use crate::config::Monotone;
let (data, gpair) = monotone_v_shape_data();
let ghist = binned(&data, 256);
let params = TrainingParams::builder()
.max_depth(4)
.min_child_weight(0.0)
.gamma(0.0)
.lambda(1.0)
.monotone_constraints(vec![Monotone::Increasing])
.build()
.unwrap();
let tree = grow_hist(¶ms, &ghist, &gpair);
assert!(non_decreasing(&tree, &data), "monotonicity violated");
}
fn root_to_leaf_feature_sets(tree: &RegTree) -> Vec<Vec<u32>> {
fn walk(tree: &RegTree, id: usize, path: &mut Vec<u32>, out: &mut Vec<Vec<u32>>) {
let node = tree.node(id);
if node.is_leaf() {
out.push(path.clone());
return;
}
path.push(node.split_feature);
walk(tree, node.left as usize, path, out);
walk(tree, node.right as usize, path, out);
path.pop();
}
let mut out = Vec::new();
walk(tree, 0, &mut Vec::new(), &mut out);
out
}
fn four_feature_data() -> (DMatrix, Vec<GradPair>) {
let n = 32;
let mut x = vec![0.0f32; n * 4];
let mut gpair = Vec::with_capacity(n);
for i in 0..n {
x[i * 4] = (i % 2) as f32;
x[i * 4 + 1] = (i % 4) as f32;
x[i * 4 + 2] = (i % 8) as f32;
x[i * 4 + 3] = (i % 16) as f32;
let g = if i % 2 == 0 { 1.0 } else { -1.0 };
gpair.push(gp(g, 1.0));
}
(DMatrix::from_dense(&x, n, 4).unwrap(), gpair)
}
#[test]
fn interaction_constraints_confine_paths_to_one_group() {
let (data, gpair) = four_feature_data();
let ghist = binned(&data, 256);
let params = TrainingParams::builder()
.max_depth(4)
.min_child_weight(0.0)
.gamma(0.0)
.lambda(0.0)
.interaction_constraints(vec![vec![0, 1], vec![2, 3]])
.build()
.unwrap();
let tree = grow_hist(¶ms, &ghist, &gpair);
for path in root_to_leaf_feature_sets(&tree) {
let has_ab = path.iter().any(|&f| f == 0 || f == 1);
let has_cd = path.iter().any(|&f| f == 2 || f == 3);
assert!(
!(has_ab && has_cd),
"path mixes interaction groups: {path:?}"
);
}
}
#[test]
fn hist_matches_exact_on_small_problem() {
let n = 60;
let mut x = Vec::new();
let mut gpair = Vec::new();
for i in 0..n {
let xi = (i as f32) * 0.1;
x.push(xi);
gpair.push(gp((xi - 3.0).sin(), 1.0));
}
let data = DMatrix::from_dense(&x, n, 1).unwrap();
let params = TrainingParams::builder()
.max_depth(3)
.lambda(1.0)
.min_child_weight(1.0)
.gamma(0.0)
.build()
.unwrap();
let exact = grow_exact(¶ms, &data, &gpair);
let ghist = binned(&data, 256);
let hist = grow_hist(¶ms, &ghist, &gpair);
for r in 0..n {
let pe = exact.predict_row(&data, r);
let ph = hist.predict_row(&data, r);
assert!((pe - ph).abs() < 1e-5, "row {r}: exact {pe} vs hist {ph}");
}
}
#[test]
fn missing_only_left_split_under_monotone_constraint() {
use crate::config::Monotone;
let x = vec![0.0f32, 1.0, 2.0, 3.0, f32::NAN, f32::NAN];
let n = x.len();
let data = DMatrix::from_dense(&x, n, 1).unwrap();
let ghist = binned(&data, 256);
assert!(ghist.dense_stride().is_none());
let mut gpair = vec![gp(-1.0, 1.0); 4];
gpair.extend([gp(1.0, 1.0), gp(1.0, 1.0)]);
let params = TrainingParams::builder()
.max_depth(1)
.lambda(0.0)
.min_child_weight(0.0)
.gamma(0.0)
.monotone_constraints(vec![Monotone::Increasing])
.build()
.unwrap();
let tree = grow_hist(¶ms, &ghist, &gpair);
assert_eq!(tree.num_nodes(), 3);
let root = tree.node(0);
assert!(root.default_left);
assert!(root.split_cond.is_finite());
let left = tree.node(root.left as usize);
let right = tree.node(root.right as usize);
assert!(
(left.sum_hess - 2.0).abs() < 1e-6,
"left cover {}",
left.sum_hess
);
assert!(
(right.sum_hess - 4.0).abs() < 1e-6,
"right cover {}",
right.sum_hess
);
for r in 0..4 {
assert!(
(tree.predict_row(&data, r) - 1.0).abs() < 1e-6,
"present row {r}"
);
}
for r in 4..6 {
assert!(
(tree.predict_row(&data, r) + 1.0).abs() < 1e-6,
"missing row {r}"
);
}
}
#[test]
fn fully_present_feature_never_splits_missing_left() {
use crate::config::Monotone;
let x = vec![
0.0f32,
5.0,
1.0,
5.0,
2.0,
5.0,
3.0,
f32::NAN,
4.0,
5.0,
5.0,
5.0,
];
let n = 6;
let data = DMatrix::from_dense(&x, n, 2).unwrap();
let ghist = binned(&data, 256);
assert!(ghist.dense_stride().is_none());
let gpair = vec![
gp(1.0, 1.0),
gp(1.0, 1.0),
gp(1.0, 1.0),
gp(-1.0, 1.0),
gp(-1.0, 1.0),
gp(-1.0, 1.0),
];
let params = TrainingParams::builder()
.max_depth(1)
.lambda(0.0)
.min_child_weight(0.0)
.gamma(0.0)
.monotone_constraints(vec![Monotone::Increasing, Monotone::None])
.build()
.unwrap();
let tree = grow_hist(¶ms, &ghist, &gpair);
assert_eq!(tree.num_nodes(), 3);
let root = tree.node(0);
assert_eq!(root.split_feature, 0);
assert!(!root.default_left);
assert!(
root.split_cond > 2.0 && root.split_cond <= 3.0,
"{}",
root.split_cond
);
for r in 0..3 {
assert!((tree.predict_row(&data, r) + 1.0).abs() < 1e-6, "row {r}");
}
for r in 3..6 {
assert!((tree.predict_row(&data, r) - 1.0).abs() < 1e-6, "row {r}");
}
}
}