use super::hist::PARALLEL_FRONTIER_ROWS;
use super::partition::{SplitRoute, child_histograms, partition_rows};
use super::shared::{
BuilderConfig, InteractionState, LeafRows, finalize_leaf_values, permits, rayon_available,
sum_rows, xgb_node_gain,
};
use super::split::SplitScorer;
use super::{Score, SplitLocation, SplitPos};
use crate::K_RT_EPS;
use crate::config::TreeMethod;
use crate::data::ghist::GHistIndex;
use crate::data::{DMatrix, FeatureType};
use crate::error::{HessboostError, Result};
use crate::objective::GradPair;
use crate::tree::constraints::{Bounds, child_bounds};
use crate::tree::gain::GradStats;
use crate::tree::hist::{Histogram, HistogramBackend, zeroed};
use crate::tree::sampler::ColumnSampler;
use crate::tree::{ChildLeaf, RegTree};
use rayon::prelude::*;
use std::num::NonZeroUsize;
const PARALLEL_SCORE_BINS: usize = 8192;
pub(crate) fn check_symmetric_input(method: TreeMethod, dtrain: &DMatrix) -> Result<()> {
if method == TreeMethod::Exact {
return Err(HessboostError::invalid_param(
"grow_policy",
"`symmetric` growth requires `tree_method=hist` or `approx`",
));
}
if dtrain.feature_types().contains(&FeatureType::Categorical) {
return Err(HessboostError::invalid_param(
"grow_policy",
"`symmetric` growth does not support categorical features",
));
}
Ok(())
}
struct LevelNode {
nid: usize,
rows: Vec<u32>,
hist: Histogram,
total: GradStats,
bounds: Bounds,
root_gain: f32,
}
#[derive(Debug, Clone, Copy)]
struct LevelSplit {
feature: u32,
offset: usize,
missing_left: bool,
score: f64,
}
impl LevelSplit {
fn pos(&self, fs: usize) -> SplitPos {
if self.missing_left {
SplitPos::backward(fs, self.offset)
} else {
SplitPos::Bin(fs + self.offset)
}
}
}
struct NodeSplit {
left: GradStats,
right: GradStats,
loss_chg: f32,
w_left: f32,
w_right: f32,
}
struct ExpandedNode {
node: LevelNode,
split: NodeSplit,
ids: [usize; 2],
bounds: [Bounds; 2],
}
pub(super) struct SymmetricTreeBuilder<'a> {
config: &'a BuilderConfig<'a>,
backend: &'a dyn HistogramBackend,
}
impl<'a> SymmetricTreeBuilder<'a> {
pub(super) fn new(config: &'a BuilderConfig<'a>, backend: &'a dyn HistogramBackend) -> Self {
SymmetricTreeBuilder { config, backend }
}
pub(super) fn build(
&self,
ghist: &GHistIndex,
gpair: &[GradPair],
row_subset: &[u32],
sampler: &mut ColumnSampler,
capture_rows: bool,
) -> (RegTree, Vec<LeafRows>) {
self.backend.prepare(ghist, gpair);
let root_stats = sum_rows(gpair, row_subset);
let mut root_hist = zeroed(ghist.total_bins());
self.backend.build(ghist, row_subset, gpair, &mut root_hist);
let mut tree = RegTree::with_root(root_stats.hess as f32);
let mut stats = vec![root_stats];
let mut bounds = vec![Bounds::default()];
let mut leaf_rows = Vec::new();
let mut record_leaf = |node: LevelNode| {
if capture_rows {
leaf_rows.push(LeafRows {
node: node.nid,
rows: node.rows,
});
}
};
let mut level = vec![self.level_node(
0,
row_subset.to_vec(),
root_hist,
root_stats,
Bounds::default(),
)];
let mut allowed: Option<InteractionState> = None;
let depth_limit = self.config.params.max_depth.map_or(0, NonZeroUsize::get);
for depth in 0..depth_limit {
let features: Vec<u32> = sampler
.sample(depth)
.iter()
.copied()
.filter(|&f| permits(allowed.as_ref(), f))
.collect();
let Some(split) = self.best_level_split(ghist, &level, &features) else {
break;
};
let cuts = ghist.cuts();
let feature = split.feature as usize;
let (fs, _) = cuts.feature_bins(feature);
let location = SplitLocation::Numeric(split.pos(fs));
let route = SplitRoute {
feature: split.feature,
location: &location,
default_left: split.missing_left,
};
let dir = self.config.cons.dir(feature);
let mut pending = Vec::with_capacity(level.len());
for node in level {
let Some(s) = self.node_split(ghist, &node, &split) else {
record_leaf(node);
continue;
};
let (lb, rb) =
child_bounds(node.bounds, dir, f64::from(s.w_left), f64::from(s.w_right));
let (left_id, right_id) = tree.expand(
node.nid,
route.rule(cuts),
ChildLeaf::new(s.w_left, s.left.hess as f32),
ChildLeaf::new(s.w_right, s.right.hess as f32),
);
tree.set_split_gain(node.nid, s.loss_chg);
debug_assert_eq!(left_id, stats.len());
stats.extend([s.left, s.right]);
bounds.extend([lb, rb]);
pending.push(ExpandedNode {
node,
split: s,
ids: [left_id, right_id],
bounds: [lb, rb],
});
}
allowed = self.config.next_allowed(allowed.as_ref(), split.feature);
let terminal = depth + 1 == depth_limit;
if terminal && !capture_rows {
level = Vec::new();
break;
}
let parallel = pending.len() > 1
&& pending.iter().map(|p| p.node.rows.len()).sum::<usize>()
>= PARALLEL_FRONTIER_ROWS
&& rayon_available();
let build =
|expanded: ExpandedNode| self.children(ghist, gpair, route, expanded, terminal);
let children: Vec<[LevelNode; 2]> = if parallel {
pending.into_par_iter().map(build).collect()
} else {
pending.into_iter().map(build).collect()
};
level = children.into_iter().flatten().collect();
}
for node in level {
record_leaf(node);
}
finalize_leaf_values(&mut tree, &stats, &bounds, &self.config.reg);
(tree, leaf_rows)
}
fn level_node(
&self,
nid: usize,
rows: Vec<u32>,
hist: Histogram,
total: GradStats,
bounds: Bounds,
) -> LevelNode {
LevelNode {
nid,
rows,
hist,
total,
bounds,
root_gain: xgb_node_gain(total, &self.config.reg, bounds),
}
}
fn children(
&self,
ghist: &GHistIndex,
gpair: &[GradPair],
route: SplitRoute,
expanded: ExpandedNode,
terminal: bool,
) -> [LevelNode; 2] {
let ExpandedNode {
node,
split,
ids: [left_id, right_id],
bounds: [lb, rb],
} = expanded;
let (left_rows, right_rows) = partition_rows(ghist, &node.rows, route);
let (left_hist, right_hist) = if terminal {
(Vec::new(), Vec::new())
} else {
child_histograms(
self.backend,
ghist,
gpair,
&left_rows,
&right_rows,
node.hist,
)
};
[
self.level_node(left_id, left_rows, left_hist, split.left, lb),
self.level_node(right_id, right_rows, right_hist, split.right, rb),
]
}
#[inline]
fn node_gain(
&self,
node: &LevelNode,
left: GradStats,
right: GradStats,
dir: i8,
) -> Option<Score<f32>> {
let scorer = SplitScorer {
reg: &self.config.reg,
root_gain: node.root_gain,
bounds: node.bounds,
dir,
};
let score = scorer.loss_chg(left, right)?;
let gain = f64::from(score.loss_chg);
(score.loss_chg.is_finite() && gain > K_RT_EPS && gain >= self.config.params.gamma)
.then_some(score)
}
fn best_level_split(
&self,
ghist: &GHistIndex,
level: &[LevelNode],
features: &[u32],
) -> Option<LevelSplit> {
let score = |&f: &u32| self.score_feature(ghist, level, f);
let per_feature: Vec<Option<LevelSplit>> =
if level.len() * ghist.total_bins() >= PARALLEL_SCORE_BINS && rayon_available() {
features.par_iter().map(score).collect()
} else {
features.iter().map(score).collect()
};
per_feature
.into_iter()
.flatten()
.fold(None, |best: Option<LevelSplit>, cand| match best {
Some(b) if b.score >= cand.score => Some(b),
_ => Some(cand),
})
}
fn score_feature(&self, ghist: &GHistIndex, level: &[LevelNode], f: u32) -> Option<LevelSplit> {
let (fs, fe) = ghist.cuts().feature_bins(f as usize);
if fe <= fs + 1 {
return None; }
let n_bins = fe - fs;
let dense = ghist.dense_stride().is_some();
let dir = self.config.cons.dir(f as usize);
let gamma = self.config.params.gamma;
let mut forward = vec![(0.0f64, 0u32); n_bins];
let mut backward = if dense {
Vec::new()
} else {
vec![(0.0f64, 0u32); n_bins]
};
let mut node_forward: Vec<Option<f64>> = if dense {
Vec::new()
} else {
vec![None; n_bins]
};
let mut any_missing = false;
let add = |slot: &mut (f64, u32), term: f64| {
slot.0 += term;
slot.1 += 1;
};
for node in level {
let hist = &node.hist[fs..fe];
let total = node.total;
let mut acc = GradStats::default();
for (b, &bin) in hist.iter().enumerate() {
acc.add(bin);
let term = self
.node_gain(node, acc, total.sub(acc), dir)
.map(|score| f64::from(score.loss_chg) - gamma);
if let Some(term) = term {
add(&mut forward[b], term);
}
if !dense {
node_forward[b] = term;
}
}
if dense {
continue;
}
if acc == total {
for b in 1..n_bins {
if let Some(term) = node_forward[b - 1] {
add(&mut backward[b], term);
}
}
} else {
any_missing = true;
let mut suffix = GradStats::default();
for b in (0..n_bins).rev() {
suffix.add(hist[b]);
if let Some(Score { loss_chg, .. }) =
self.node_gain(node, total.sub(suffix), suffix, dir)
{
add(&mut backward[b], f64::from(loss_chg) - gamma);
}
}
}
}
let mut best: Option<LevelSplit> = None;
let mut consider = |(score, splits): (f64, u32), offset: usize, missing_left: bool| {
if splits > 0 && best.is_none_or(|b| score > b.score) {
best = Some(LevelSplit {
feature: f,
offset,
missing_left,
score,
});
}
};
for (b, &slot) in forward.iter().enumerate() {
consider(slot, b, false);
}
if any_missing {
for (b, &slot) in backward.iter().enumerate().rev() {
consider(slot, b, true);
}
}
best
}
fn node_split(
&self,
ghist: &GHistIndex,
node: &LevelNode,
split: &LevelSplit,
) -> Option<NodeSplit> {
let feature = split.feature as usize;
let (fs, fe) = ghist.cuts().feature_bins(feature);
let hist = &node.hist[fs..fe];
let total = node.total;
let prefix = |last: usize| {
let mut acc = GradStats::default();
for &bin in &hist[..=last] {
acc.add(bin);
}
acc
};
let (left, right) = if split.missing_left {
let has_missing = ghist.dense_stride().is_none() && prefix(hist.len() - 1) != total;
if has_missing {
let mut suffix = GradStats::default();
for &bin in hist[split.offset..].iter().rev() {
suffix.add(bin);
}
(total.sub(suffix), suffix)
} else if split.offset == 0 {
return None; } else {
let left = prefix(split.offset - 1);
(left, total.sub(left))
}
} else {
let left = prefix(split.offset);
(left, total.sub(left))
};
let Score {
loss_chg,
w_left,
w_right,
} = self.node_gain(node, left, right, self.config.cons.dir(feature))?;
Some(NodeSplit {
left,
right,
loss_chg,
w_left,
w_right,
})
}
}
#[cfg(test)]
mod tests {
use super::super::test_support::{binned, gp, grow_hist};
use super::*;
use crate::config::{GrowPolicy, Monotone, TrainingParams};
use crate::model::Iterations;
use crate::training::train;
use crate::tree::builder::{HistTreeBuilder, all_rows};
fn unit(i: usize) -> f32 {
let h = (i as u64)
.wrapping_mul(0x9E37_79B9_7F4A_7C15)
.rotate_left(17)
^ 0x2545_F491;
(h.wrapping_mul(0xBF58_476D_1CE4_E5B9) >> 40) as f32 / (1u64 << 24) as f32
}
fn synthetic(n: usize, n_features: usize, missing: bool) -> (DMatrix, Vec<f32>) {
let mut x: Vec<f32> = (0..n * n_features).map(unit).collect();
let y: Vec<f32> = x
.chunks_exact(n_features)
.map(|r| (3.0 * r[0]).sin() + 2.0 * r[1] * r[2] + if r[3] > 0.6 { 1.0 } else { 0.0 })
.collect();
if missing {
for (i, v) in x.iter_mut().enumerate() {
if i % n_features < 2 && (i / n_features) % 7 == 3 {
*v = f32::NAN;
}
}
}
let d = DMatrix::from_dense(&x, n, n_features).unwrap();
(d, y)
}
fn residual_gradients(y: &[f32]) -> Vec<GradPair> {
y.iter().map(|&v| gp(-v, 1.0)).collect()
}
fn symmetric() -> crate::config::TrainingParamsBuilder {
TrainingParams::builder()
.tree_method(TreeMethod::Hist)
.grow_policy(GrowPolicy::Symmetric)
}
fn grow(params: &TrainingParams, data: &DMatrix, gpair: &[GradPair]) -> RegTree {
grow_hist(params, &binned(data, 64), gpair)
}
fn levels(tree: &RegTree) -> Vec<(u32, u32, bool)> {
let mut levels = Vec::new();
let mut frontier = vec![0usize];
while !frontier.is_empty() {
let mut next = Vec::new();
let mut level = None;
for id in frontier {
let n = tree.node(id);
if n.is_leaf() {
continue;
}
let split = (n.split_feature, n.split_cond.to_bits(), n.default_left);
assert_eq!(*level.get_or_insert(split), split, "node {id} breaks level");
next.extend([n.left as usize, n.right as usize]);
}
levels.extend(level);
frontier = next;
}
levels
}
#[test]
fn full_tree_has_one_split_per_level() {
let n = 32 * 20;
let mut x = Vec::new();
let mut gpair = Vec::new();
for i in 0..n {
let row: Vec<f32> = (0..5).map(|b| ((i >> b) & 1) as f32).collect();
let y: f32 = row
.iter()
.zip([1.0, 2.0, 4.0, 8.0, 16.0])
.map(|(v, w)| v * w)
.sum();
x.extend(row);
gpair.push(gp(-y, 1.0));
}
let data = DMatrix::from_dense(&x, n, 5).unwrap();
let params = symmetric().max_depth(5).lambda(0.0).build().unwrap();
let tree = grow(¶ms, &data, &gpair);
let lv = levels(&tree);
let mut features: Vec<u32> = lv.iter().map(|l| l.0).collect();
assert_eq!(features, [4, 3, 2, 1, 0]);
features.dedup();
assert_eq!(features.len(), 5);
assert_eq!(tree.num_leaves(), 32);
assert_eq!(tree.num_nodes(), 63);
}
#[test]
fn level_split_maximizes_summed_gain_not_each_node() {
let n = 400;
let mut x = Vec::new();
let mut gpair = Vec::new();
for i in 0..n {
let (x0, x1, x2) = ((i % 2) as f32, ((i / 2) % 2) as f32, ((i / 4) % 2) as f32);
x.extend([x0, x1, x2]);
let y = if x0 == 0.0 { 5.0 * x1 } else { x2 - 2.5 };
gpair.push(gp(-y, 1.0));
}
let data = DMatrix::from_dense(&x, n, 3).unwrap();
let depthwise = TrainingParams::builder()
.tree_method(TreeMethod::Hist)
.max_depth(2)
.build()
.unwrap();
let dw = grow(&depthwise, &data, &gpair);
let (l, r) = (dw.node(1), dw.node(2));
assert_eq!((l.split_feature, r.split_feature), (1, 2));
let params = symmetric().max_depth(2).build().unwrap();
let tree = grow(¶ms, &data, &gpair);
let lv = levels(&tree);
assert_eq!(lv.iter().map(|l| l.0).collect::<Vec<_>>(), [0, 1]);
assert!(!tree.node(1).is_leaf());
assert!(tree.node(2).is_leaf());
}
#[test]
fn nodes_failing_min_child_weight_or_gamma_stay_leaves() {
let (data, y) = synthetic(600, 5, false);
let gpair = residual_gradients(&y);
let full = grow(
&symmetric()
.max_depth(6)
.min_child_weight(0.0)
.build()
.unwrap(),
&data,
&gpair,
);
for (mcw, gamma) in [(40.0, 0.0), (0.0, 2.0)] {
let params = symmetric()
.max_depth(6)
.min_child_weight(mcw)
.gamma(gamma)
.build()
.unwrap();
let tree = grow(¶ms, &data, &gpair);
assert!(!levels(&tree).is_empty());
assert!(tree.num_leaves() > 1 && tree.num_leaves() < full.num_leaves());
for n in tree.nodes() {
if n.is_leaf() {
assert!(f64::from(n.sum_hess) >= mcw);
} else {
assert!(f64::from(n.split_gain) >= gamma);
}
}
}
let stump = grow(&symmetric().gamma(1e9).build().unwrap(), &data, &gpair);
assert_eq!(stump.num_nodes(), 1);
}
#[test]
fn missing_values_follow_one_direction_per_level_and_leaf_rows_match_routing() {
let (data, y) = synthetic(3000, 5, true);
let ghist = binned(&data, 32);
assert!(ghist.dense_stride().is_none());
let params = symmetric().max_depth(4).build().unwrap();
let (tree, leaf_rows) = HistTreeBuilder::new(¶ms).build_with_leaf_rows(
&ghist,
&residual_gradients(&y),
&all_rows(data.n_rows()),
&mut ColumnSampler::all(data.n_cols()),
);
assert_eq!(levels(&tree).len(), 4);
assert_eq!(leaf_rows.iter().map(|l| l.rows.len()).sum::<usize>(), 3000);
assert_eq!(leaf_rows.len(), tree.num_leaves());
for leaf in &leaf_rows {
assert!(leaf.rows.is_sorted());
for &r in &leaf.rows {
let routed = tree.leaf_id_with(|f| data.get(r as usize, f as usize));
assert_eq!(routed, leaf.node);
}
}
}
#[test]
fn monotone_constraint_holds() {
let (data, y) = synthetic(2000, 4, false);
let params = symmetric()
.max_depth(4)
.monotone_constraints(vec![Monotone::Decreasing])
.build()
.unwrap();
let model = train(¶ms, &data.clone().with_labels(&y).unwrap(), 20).unwrap();
for i in 0..50 {
let base: Vec<f32> = (1..4).map(|f| unit(10_000 + 4 * i + f)).collect();
let preds: Vec<f32> = (0..=40)
.map(|s| {
let mut row = vec![s as f32 / 40.0];
row.extend(&base);
let d = DMatrix::from_dense(&row, 1, 4).unwrap();
*model
.predict(&d, Iterations::Best)
.unwrap()
.get(0, 0)
.unwrap()
})
.collect();
assert!(preds.windows(2).all(|w| w[1] <= w[0]), "{preds:?}");
}
}
#[test]
fn interaction_constraints_bound_every_path() {
let (data, y) = synthetic(2000, 6, false);
let params = symmetric()
.max_depth(4)
.interaction_constraints(vec![vec![0, 1, 2], vec![3, 4]])
.build()
.unwrap();
let model = train(¶ms, &data.with_labels(&y).unwrap(), 10).unwrap();
for tree in model.trees() {
let features: Vec<u32> = levels(tree).iter().map(|l| l.0).collect();
let within = |g: &[u32]| features.iter().all(|f| g.contains(f));
assert!(
within(&[0, 1, 2]) || within(&[3, 4]) || features.len() <= 1,
"{features:?}"
);
}
}
#[test]
fn serial_and_parallel_growth_agree() {
let (data, y) = synthetic(20_000, 8, true);
let data = data.with_labels(&y).unwrap();
let params = symmetric()
.max_depth(6)
.colsample_bylevel(0.75)
.build()
.unwrap();
let run = |threads| {
rayon::ThreadPoolBuilder::new()
.num_threads(threads)
.build()
.unwrap()
.install(|| train(¶ms, &data, 5).unwrap())
};
let (serial, parallel) = (run(1), run(8));
assert_eq!(serial.trees(), parallel.trees());
assert_eq!(serial.trees(), run(8).trees());
}
#[test]
fn quality_is_close_to_depthwise() {
let (dtrain, ytrain) = synthetic(8000, 6, false);
let (dtest, ytest) = synthetic(2000, 6, false);
let dtrain = dtrain.with_labels(&ytrain).unwrap();
let rmse = |policy| {
let params = TrainingParams::builder()
.tree_method(TreeMethod::Hist)
.grow_policy(policy)
.max_depth(6)
.eta(0.1)
.build()
.unwrap();
let pred = train(¶ms, &dtrain, 200)
.unwrap()
.predict(&dtest, Iterations::Best)
.unwrap();
let se: f32 = pred
.as_slice()
.iter()
.zip(&ytest)
.map(|(p, y)| (p - y).powi(2))
.sum();
(se / ytest.len() as f32).sqrt()
};
let mean = ytest.iter().sum::<f32>() / ytest.len() as f32;
let std =
(ytest.iter().map(|y| (y - mean).powi(2)).sum::<f32>() / ytest.len() as f32).sqrt();
let (sym, dw) = (rmse(GrowPolicy::Symmetric), rmse(GrowPolicy::DepthWise));
assert!(sym < 0.2 * std, "symmetric rmse {sym} vs target std {std}");
assert!(sym < 1.5 * dw, "symmetric rmse {sym} vs depthwise {dw}");
}
#[test]
fn unsupported_configurations_are_rejected() {
let (data, y) = synthetic(200, 4, false);
let data = data.with_labels(&y).unwrap();
for depth in [0, crate::config::MAX_SYMMETRIC_DEPTH + 1] {
assert!(symmetric().max_depth(depth).build().is_err());
}
assert!(symmetric().max_leaves(8).build().is_err());
let exact = symmetric().tree_method(TreeMethod::Exact).build().unwrap();
assert!(train(&exact, &data, 1).is_err());
let codes: Vec<f32> = (0..200).map(|i| (i % 3) as f32).collect();
let categorical = DMatrix::from_dense(&codes, 200, 1)
.unwrap()
.with_labels(&y)
.unwrap()
.with_feature_types(&[FeatureType::Categorical])
.unwrap();
let err = train(&symmetric().build().unwrap(), &categorical, 1).unwrap_err();
assert!(err.to_string().contains("categorical"), "{err}");
let approx = symmetric().tree_method(TreeMethod::Approx).build().unwrap();
let model = train(&approx, &data, 3).unwrap();
assert!(model.trees().iter().all(|t| !levels(t).is_empty()));
}
}