mod exact;
mod hist;
pub use exact::{ExactTreeBuilder, SortedColumns, all_features, all_rows};
pub use hist::HistTreeBuilder;
pub(crate) use hist::LeafRows;
use std::collections::BTreeSet;
use crate::objective::GradPair;
use crate::tree::constraints::{Bounds, calc_weight_bounded, gain_at_weight, satisfies};
use crate::tree::gain::{GradStats, RegParams, calc_gain, threshold_l1};
use crate::tree::regtree::RegTree;
pub(super) const K_RT_EPS: f64 = 1e-6;
#[derive(Debug, Clone)]
pub(super) struct BestSplit {
pub(super) loss_chg: f64,
pub(super) feature: u32,
pub(super) threshold: f32,
pub(super) split_bin: Option<usize>,
pub(super) default_left: bool,
pub(super) left: GradStats,
pub(super) right: GradStats,
pub(super) w_left: f64,
pub(super) w_right: f64,
pub(super) is_categorical: bool,
pub(super) cat_left: Vec<u32>,
}
impl BestSplit {
pub(super) fn none() -> Self {
BestSplit {
loss_chg: 0.0,
feature: 0,
threshold: 0.0,
split_bin: None,
default_left: true,
left: GradStats::default(),
right: GradStats::default(),
w_left: 0.0,
w_right: 0.0,
is_categorical: false,
cat_left: Vec::new(),
}
}
#[allow(clippy::too_many_arguments)]
pub(super) fn numeric(
loss_chg: f64,
feature: u32,
pos: SplitPos,
default_left: bool,
left: GradStats,
right: GradStats,
w_left: f64,
w_right: f64,
) -> Self {
let (threshold, split_bin) = match pos {
SplitPos::Value(t) => (t, None),
SplitPos::Bin(b) => (0.0, Some(b)),
SplitPos::BelowBins => (0.0, None),
};
BestSplit {
loss_chg,
feature,
threshold,
split_bin,
default_left,
left,
right,
w_left,
w_right,
is_categorical: false,
cat_left: Vec::new(),
}
}
#[allow(clippy::too_many_arguments)]
pub(super) fn categorical(
loss_chg: f64,
feature: u32,
left: GradStats,
right: GradStats,
w_left: f64,
w_right: f64,
cat_left: Vec<u32>,
) -> Self {
BestSplit {
loss_chg,
feature,
threshold: 0.0,
split_bin: None,
default_left: false,
left,
right,
w_left,
w_right,
is_categorical: true,
cat_left,
}
}
#[inline]
pub(super) fn found(&self) -> bool {
self.loss_chg > K_RT_EPS
}
pub(super) fn valid(&self, gamma: f64, min_child_weight: f64) -> bool {
self.found()
&& self.loss_chg >= gamma
&& self.left.hess > 0.0
&& self.right.hess > 0.0
&& self.left.hess >= min_child_weight
&& self.right.hess >= min_child_weight
}
}
#[inline]
pub(super) fn candidate_gain(
left: GradStats,
right: GradStats,
parent: f64,
bounds: Bounds,
dir: i8,
constrained: bool,
reg: &RegParams,
) -> Option<(f64, f64, f64)> {
if constrained {
let wl = calc_weight_bounded(left, reg, bounds);
let wr = calc_weight_bounded(right, reg, bounds);
if !satisfies(dir, wl, wr) {
return None;
}
let g = gain_at_weight(left, reg, wl) + gain_at_weight(right, reg, wr) - parent;
Some((g, wl, wr))
} else {
let g = calc_gain(left, reg) + calc_gain(right, reg) - parent;
Some((g, 0.0, 0.0))
}
}
#[derive(Debug, Clone, Copy)]
pub(super) enum SplitPos {
Value(f32),
Bin(usize),
BelowBins,
}
pub(super) const BELOW_ALL_VALUES: f32 = f32::MIN;
#[inline]
pub(super) fn xgb_weight(stats: GradStats, reg: &RegParams, bounds: Bounds) -> f32 {
let w = if stats.hess <= 0.0 {
0.0
} else {
let mut w = -threshold_l1(stats.grad, reg.alpha) / (stats.hess + reg.lambda);
if reg.max_delta_step != 0.0 && w.abs() > reg.max_delta_step {
w = reg.max_delta_step.copysign(w);
}
w
} as f32;
let (lower, upper) = (bounds.lower as f32, bounds.upper as f32);
if w < lower {
lower
} else if w > upper {
upper
} else {
w
}
}
#[inline]
fn xgb_gain_given_weight(stats: GradStats, reg: &RegParams, w: f32) -> f64 {
-(2.0 * stats.grad * f64::from(w)
+ (stats.hess + reg.lambda) * f64::from(w * w)
+ 2.0 * reg.alpha * f64::from(w.abs()))
}
pub(super) fn xgb_node_gain(stats: GradStats, reg: &RegParams, bounds: Bounds) -> f32 {
if stats.hess <= 0.0 {
return 0.0;
}
xgb_gain_given_weight(stats, reg, xgb_weight(stats, reg, bounds)) as f32
}
#[inline]
pub(super) fn xgb_loss_chg(
left: GradStats,
right: GradStats,
root_gain: f32,
reg: &RegParams,
bounds: Bounds,
dir: i8,
) -> Option<(f32, f32, f32)> {
let mcw = reg.min_child_weight;
if !(left.hess > 0.0 && right.hess > 0.0 && left.hess >= mcw && right.hess >= mcw) {
return None;
}
let wl = xgb_weight(left, reg, bounds);
let wr = xgb_weight(right, reg, bounds);
if !satisfies(dir, f64::from(wl), f64::from(wr)) {
return None;
}
let gain =
xgb_gain_given_weight(left, reg, wl) as f32 + xgb_gain_given_weight(right, reg, wr) as f32;
Some((gain - root_gain, wl, wr))
}
#[allow(clippy::too_many_arguments)]
pub(super) fn xgb_update(
best: &mut BestSplit,
loss_chg: f32,
feature: u32,
pos: SplitPos,
default_left: bool,
left: GradStats,
right: GradStats,
w_left: f32,
w_right: f32,
) -> bool {
if loss_chg.is_infinite() {
return false;
}
let incumbent = best.loss_chg as f32;
let replace = if best.feature <= feature {
loss_chg > incumbent
} else {
incumbent.partial_cmp(&loss_chg) != Some(std::cmp::Ordering::Greater)
};
if replace {
*best = BestSplit::numeric(
f64::from(loss_chg),
feature,
pos,
default_left,
left,
right,
f64::from(w_left),
f64::from(w_right),
);
}
replace
}
#[allow(clippy::too_many_arguments)]
pub(super) fn sweep_categorical(
best: &mut BestSplit,
cats: &mut [(u32, GradStats)],
total: GradStats,
parent_gain: f64,
bounds: Bounds,
dir: i8,
constrained: bool,
reg: &RegParams,
feature: u32,
) {
if cats.len() < 2 {
return; }
let ratio = |s: GradStats| s.grad / (s.hess + reg.lambda);
cats.sort_by(|a, b| ratio(a.1).total_cmp(&ratio(b.1)));
let mcw = reg.min_child_weight;
let mut left = GradStats::default();
let mut cats_left: Vec<u32> = Vec::new();
for &(cat, s) in &cats[..cats.len() - 1] {
left.add(s);
cats_left.push(cat);
let right = total.sub(left);
if left.hess < mcw || right.hess < mcw {
continue;
}
let Some((g, wl, wr)) =
candidate_gain(left, right, parent_gain, bounds, dir, constrained, reg)
else {
continue;
};
if g > best.loss_chg + K_RT_EPS {
*best = BestSplit::categorical(g, feature, left, right, wl, wr, cats_left.clone());
}
}
}
#[derive(Clone)]
pub(super) struct InteractionState {
path: Vec<u32>,
allowed: Vec<u32>,
}
pub(super) fn build_interaction_sets(groups: &[Vec<u32>]) -> Option<Vec<Vec<u32>>> {
if groups.is_empty() {
return None;
}
Some(
groups
.iter()
.map(|group| {
let mut group = group.clone();
group.sort_unstable();
group.dedup();
group
})
.collect(),
)
}
pub(super) fn next_allowed(
parent: Option<&InteractionState>,
feature: u32,
groups: Option<&[Vec<u32>]>,
) -> Option<InteractionState> {
let groups = groups?;
let mut path = parent.map_or_else(Vec::new, |state| state.path.clone());
if let Err(pos) = path.binary_search(&feature) {
path.insert(pos, feature);
}
let mut allowed: BTreeSet<u32> = path.iter().copied().collect();
for group in groups {
if path
.iter()
.all(|feature| group.binary_search(feature).is_ok())
{
allowed.extend(group.iter().copied());
}
}
Some(InteractionState {
path,
allowed: allowed.into_iter().collect(),
})
}
pub(super) fn permits(state: Option<&InteractionState>, feature: u32) -> bool {
state.is_none_or(|state| state.allowed.binary_search(&feature).is_ok())
}
pub(super) fn sum_rows(gpair: &[GradPair], rows: &[u32]) -> GradStats {
let mut total = GradStats::default();
for &r in rows {
total.add(GradStats::from_pair(gpair[r as usize]));
}
total
}
pub(super) fn finalize_leaf_values(
tree: &mut RegTree,
stats: &[GradStats],
bounds: &[Bounds],
reg: &RegParams,
) {
#[allow(clippy::needless_range_loop)]
for id in 0..tree.num_nodes() {
if tree.node(id).is_leaf() {
let w = calc_weight_bounded(stats[id], reg, bounds[id]);
tree.set_leaf_value(id, w as f32);
}
}
}
#[cfg(test)]
mod test_support {
use crate::data::DMatrix;
use crate::objective::GradPair;
pub(super) fn gp(g: f32, h: f32) -> GradPair {
GradPair::new(g, h)
}
pub(super) fn monotone_v_shape_data() -> (DMatrix, Vec<GradPair>) {
let n = 60;
let mut x = Vec::new();
let mut gpair = Vec::new();
for i in 0..n {
let xi = i as f32 / n as f32;
x.push(xi);
let target = (xi - 0.5).abs(); gpair.push(gp(-(target - 0.25), 1.0)); }
(DMatrix::from_dense(&x, n, 1).unwrap(), gpair)
}
}