use super::categorical::{MAX_CAT_THRESHOLD, MAX_CAT_TO_ONEHOT};
use super::partition::{SplitRoute, partition_rows, with_sibling};
use super::shared::{
BuilderConfig, InteractionState, LeafRows, apply_bounds, permits, rayon_available,
xgb_calc_weight, xgb_gain_given_weight,
};
use super::{SplitLocation, SplitPos, limit_or_unbounded, need_replace};
use crate::K_RT_EPS_F32;
use crate::config::{GrowPolicy, TrainingParams};
use crate::data::ghist::{Bins, GHistIndex};
use crate::objective::GradPair;
use crate::tree::gain::{GradStats, RegParams, threshold_l1};
use crate::tree::hist::feature_slices;
use crate::tree::sampler::{ColumnSampler, FeatureSet};
use crate::tree::{ChildLeaf, RegTree};
use rayon::prelude::*;
use std::cmp::Ordering;
const MAX_NODE_BATCH: usize = 256;
const PARALLEL_EVALUATE_ROWS: usize = 16_384;
const PARALLEL_HIST_ENTRIES: usize = 65_536;
pub(crate) struct VectorGradients<'a> {
pub(crate) split: &'a [GradPair],
pub(crate) n_split: usize,
pub(crate) value: Option<&'a [GradPair]>,
pub(crate) n_outputs: usize,
}
pub(crate) struct MultiTreeBuilder<'a> {
config: BuilderConfig<'a>,
}
#[derive(Debug, Clone)]
struct Candidate {
loss_chg: f32,
feature: u32,
default_left: bool,
loc: SplitLocation,
left: Vec<GradStats>,
right: Vec<GradStats>,
}
impl Candidate {
fn none() -> Self {
Candidate {
loss_chg: 0.0,
feature: 0,
default_left: false,
loc: SplitLocation::Numeric(SplitPos::BelowBins),
left: Vec::new(),
right: Vec::new(),
}
}
fn update(
&mut self,
loss_chg: f32,
feature: u32,
default_left: bool,
loc: impl FnOnce() -> SplitLocation,
left: &[GradStats],
right: &[GradStats],
) -> bool {
if !need_replace(self.loss_chg, self.feature, loss_chg, feature) {
return false;
}
self.loss_chg = loss_chg;
self.feature = feature;
self.default_left = default_left;
self.loc = loc();
self.left.clear();
self.left.extend_from_slice(left);
self.right.clear();
self.right.extend_from_slice(right);
true
}
fn merge(&mut self, other: Candidate) {
if need_replace(self.loss_chg, self.feature, other.loss_chg, other.feature) {
*self = other;
}
}
fn is_categorical(&self) -> bool {
self.loc.is_categorical()
}
}
#[derive(Clone, Copy)]
struct FeatureBins<'h> {
nid: usize,
hist: &'h [GradStats],
feature: u32,
fs: usize,
fe: usize,
}
struct Entry {
nid: usize,
order: usize,
depth: usize,
rows: Vec<u32>,
hist: Vec<GradStats>,
best: Candidate,
allowed: Option<InteractionState>,
}
struct Expanded {
entry: Entry,
children: (usize, usize),
rows: (Vec<u32>, Vec<u32>),
child_valid: bool,
}
struct Grow<'g, 'a> {
b: &'g MultiTreeBuilder<'a>,
ghist: &'g GHistIndex,
grad: &'g VectorGradients<'g>,
tree: RegTree,
stats: Vec<GradStats>,
gain: Vec<f64>,
weights: Vec<f32>,
lower: Vec<f32>,
upper: Vec<f32>,
leaf_rows: Vec<LeafRows>,
num_leaves: usize,
}
impl<'a> MultiTreeBuilder<'a> {
pub(crate) fn new(params: &'a TrainingParams) -> Self {
MultiTreeBuilder {
config: BuilderConfig::new(params),
}
}
pub(crate) fn build(
&self,
ghist: &GHistIndex,
grad: &VectorGradients,
rows: &[u32],
sampler: &mut ColumnSampler,
) -> (RegTree, Vec<LeafRows>) {
let s = grad.n_split;
debug_assert!(grad.n_outputs > 1 && s >= 1);
debug_assert_eq!(grad.split.len(), ghist.n_rows() * s);
debug_assert!(grad.value.is_some() || s == grad.n_outputs);
let mut root = vec![GradStats::default(); s];
for &r in rows {
let g = &grad.split[r as usize * s..][..s];
for (acc, gp) in root.iter_mut().zip(g) {
acc.add(GradStats::from_pair(*gp));
}
}
let root_hess = root.iter().fold(0.0f32, |acc, t| acc + t.hess as f32);
let n_bounds = if self.config.cons.is_active() { s } else { 0 };
let mut grow = Grow {
b: self,
ghist,
grad,
tree: RegTree::with_vector_root(grad.n_outputs, root_hess),
stats: root.clone(),
gain: Vec::new(),
weights: Vec::new(),
lower: vec![-f32::MAX; n_bounds],
upper: vec![f32::MAX; n_bounds],
leaf_rows: Vec::new(),
num_leaves: 1,
};
let root_weight: Vec<f32> = (0..s).map(|t| grow.weight(0, t, root[t])).collect();
grow.gain.push(self.gain_given_weights(&root, &root_weight));
grow.weights.extend_from_slice(&root_weight);
let root_hist = grow.build_hist(rows);
let features = sampler.sample(0);
let best = grow.evaluate(0, &root_hist, &features, None, rows.len());
let mut queue = vec![Entry {
nid: 0,
order: 0,
depth: 0,
rows: rows.to_vec(),
hist: root_hist,
best,
allowed: None,
}];
grow.run(&mut queue, sampler);
for entry in queue {
grow.record_leaf(entry.nid, entry.rows);
}
grow.finish()
}
fn gain_given_weights(&self, stats: &[GradStats], weights: &[f32]) -> f64 {
stats.iter().zip(weights).fold(0.0, |gain, (st, &w)| {
gain + xgb_gain_given_weight(*st, &self.config.reg, w)
})
}
}
impl Grow<'_, '_> {
fn n_split(&self) -> usize {
self.grad.n_split
}
fn constrained(&self) -> bool {
!self.lower.is_empty()
}
fn weight(&self, nid: usize, t: usize, stats: GradStats) -> f32 {
self.bound(nid, t, xgb_calc_weight(stats, &self.b.config.reg) as f32)
}
fn bound(&self, nid: usize, t: usize, w: f32) -> f32 {
if !self.constrained() {
return w;
}
let i = nid * self.n_split() + t;
apply_bounds(w, self.lower[i], self.upper[i])
}
fn split_weights(
&self,
nid: usize,
t: usize,
dir: i8,
left: GradStats,
right: GradStats,
) -> (f32, f32) {
let wl = self.weight(nid, t, left);
let wr = self.weight(nid, t, right);
if !self.constrained() {
return (wl, wr);
}
let ordered = dir == 0 || (dir > 0 && wl <= wr) || (dir < 0 && wl >= wr);
if ordered {
return (wl, wr);
}
let pooled_reg = RegParams {
lambda: 2.0 * self.b.config.reg.lambda,
alpha: 2.0 * self.b.config.reg.alpha,
..self.b.config.reg
};
let mut both = left;
both.add(right);
let pooled = self.bound(nid, t, xgb_calc_weight(both, &pooled_reg) as f32);
(pooled, pooled)
}
fn split_gain(&self, nid: usize, dir: i8, left: &[GradStats], right: &[GradStats]) -> f64 {
let reg = &self.b.config.reg;
let constrained = self.constrained();
let (mut left_hess, mut right_hess, mut gain) = (0.0f64, 0.0f64, 0.0f64);
for (l, r) in left.iter().zip(right) {
left_hess += l.hess;
right_hess += r.hess;
if !constrained {
gain += calc_gain(reg, *l);
gain += calc_gain(reg, *r);
}
}
let k = left.len() as f64;
let (lh, rh) = (left_hess / k, right_hess / k);
let mcw = reg.min_child_weight;
if !(lh > 0.0 && rh > 0.0 && lh >= mcw && rh >= mcw) {
return f64::NEG_INFINITY;
}
if !constrained {
return gain;
}
for (t, (l, r)) in left.iter().zip(right).enumerate() {
let (wl, wr) = self.split_weights(nid, t, dir, *l, *r);
gain += xgb_gain_given_weight(*l, reg, wl);
gain += xgb_gain_given_weight(*r, reg, wr);
}
gain
}
fn node_stats(&self, nid: usize) -> &[GradStats] {
let s = self.n_split();
&self.stats[nid * s..(nid + 1) * s]
}
fn build_hist(&self, rows: &[u32]) -> Vec<GradStats> {
let s = self.n_split();
let ghist = self.ghist;
let gp = self.grad.split;
let mut hist = vec![GradStats::default(); ghist.total_bins() * s];
let add = |slot: &mut [GradStats], g: &[GradPair]| {
for (h, p) in slot.iter_mut().zip(g) {
h.grad += f64::from(p.grad);
h.hess += f64::from(p.hess);
}
};
if let Some(columns) = ghist.column_bins()
&& rows.len().saturating_mul(ghist.n_cols()) >= PARALLEL_HIST_ENTRIES
&& rayon_available()
{
let n_rows = ghist.n_rows();
feature_slices(ghist, &mut hist, s)
.into_par_iter()
.enumerate()
.for_each(|(f, (fs, slice))| {
let column = |r: u32| -> usize {
match &columns {
Bins::U16(c) => usize::from(c[f * n_rows + r as usize]),
Bins::U32(c) => c[f * n_rows + r as usize] as usize,
}
};
for &r in rows {
let b = column(r) - fs;
add(&mut slice[b * s..(b + 1) * s], &gp[r as usize * s..][..s]);
}
});
return hist;
}
let row_ptr = ghist.row_ptr();
let mut accumulate = |bins: &mut dyn Iterator<Item = usize>, g: &[GradPair]| {
for b in bins {
add(&mut hist[b * s..(b + 1) * s], g);
}
};
match ghist.bins() {
Bins::U16(bins) => {
for &r in rows {
let r = r as usize;
accumulate(
&mut bins[row_ptr[r]..row_ptr[r + 1]]
.iter()
.map(|&b| usize::from(b)),
&gp[r * s..][..s],
);
}
}
Bins::U32(bins) => {
for &r in rows {
let r = r as usize;
accumulate(
&mut bins[row_ptr[r]..row_ptr[r + 1]].iter().map(|&b| b as usize),
&gp[r * s..][..s],
);
}
}
}
hist
}
fn evaluate(
&self,
nid: usize,
hist: &[GradStats],
features: &[u32],
allowed: Option<&InteractionState>,
n_rows: usize,
) -> Candidate {
let features: Vec<u32> = features
.iter()
.copied()
.filter(|&f| permits(allowed, f))
.collect();
let one = |f: u32| {
let mut best = Candidate::none();
self.evaluate_feature(nid, hist, f, &mut best);
best
};
let mut best = Candidate::none();
if n_rows >= PARALLEL_EVALUATE_ROWS && features.len() > 1 && rayon_available() {
let per_feature: Vec<Candidate> = features.par_iter().map(|&f| one(f)).collect();
for c in per_feature {
best.merge(c);
}
} else {
for &f in &features {
best.merge(one(f));
}
}
best
}
fn evaluate_feature(&self, nid: usize, hist: &[GradStats], f: u32, best: &mut Candidate) {
let cuts = self.ghist.cuts();
let (fs, fe) = cuts.feature_bins(f as usize);
if fe <= fs {
return;
}
let bins = FeatureBins {
nid,
hist,
feature: f,
fs,
fe,
};
if cuts.is_categorical(f as usize) {
if fe - fs < MAX_CAT_TO_ONEHOT {
self.enumerate_one_hot(&bins, best);
} else {
self.enumerate_partition(&bins, best);
}
} else if self.enumerate_numeric(&bins, true, best) {
self.enumerate_numeric(&bins, false, best);
}
}
fn enumerate_numeric(&self, bins: &FeatureBins, forward: bool, best: &mut Candidate) -> bool {
let FeatureBins {
nid,
hist,
feature: f,
fs,
fe,
} = *bins;
let s = self.n_split();
let parent = self.node_stats(nid);
let parent_gain = self.gain[nid];
let dir = self.b.config.cons.dir(f as usize);
let mut acc = vec![GradStats::default(); s];
let mut rest = vec![GradStats::default(); s];
let bins: Box<dyn Iterator<Item = usize>> = if forward {
Box::new(fs..fe)
} else {
Box::new((fs..fe).rev())
};
for i in bins {
for t in 0..s {
acc[t].add(hist[i * s + t]);
rest[t] = parent[t].sub(acc[t]);
}
if forward {
let loss = (self.split_gain(nid, dir, &acc, &rest) - parent_gain) as f32;
let loc = || SplitLocation::Numeric(SplitPos::Bin(i));
best.update(loss, f, false, loc, &acc, &rest);
} else {
let loss = (self.split_gain(nid, dir, &rest, &acc) - parent_gain) as f32;
let loc = || SplitLocation::Numeric(SplitPos::backward(fs, i - fs));
best.update(loss, f, true, loc, &rest, &acc);
}
}
forward && acc.as_slice() != parent
}
fn enumerate_one_hot(&self, bins: &FeatureBins, best: &mut Candidate) {
let FeatureBins {
nid,
hist,
feature: f,
fs,
fe,
} = *bins;
let s = self.n_split();
let parent = self.node_stats(nid);
let parent_gain = self.gain[nid];
let dir = self.b.config.cons.dir(f as usize);
let cuts = self.ghist.cuts();
let mut missing = parent.to_vec();
for (t, m) in missing.iter_mut().enumerate() {
let mut present = GradStats::default();
for i in fs..fe {
present.add(hist[i * s + t]);
}
*m = m.sub(present);
}
let mut left = vec![GradStats::default(); s];
let mut right = vec![GradStats::default(); s];
let mut local = Candidate::none();
for i in fs..fe {
let cat = || SplitLocation::Categories(vec![cuts.cut_value(i) as u32]);
for missing_left in [true, false] {
for t in 0..s {
right[t] = hist[i * s + t];
if !missing_left {
right[t].add(missing[t]);
}
left[t] = parent[t].sub(right[t]);
}
let loss = (self.split_gain(nid, dir, &left, &right) - parent_gain) as f32;
local.update(loss, f, missing_left, cat, &left, &right);
}
}
if local.is_categorical() {
best.merge(local);
}
}
fn enumerate_partition(&self, bins: &FeatureBins, best: &mut Candidate) {
let FeatureBins {
nid, hist, fs, fe, ..
} = *bins;
let s = self.n_split();
let reg = &self.b.config.reg;
let parent = self.node_stats(nid);
let n_bins = fe - fs;
let parent_w: Vec<f32> = parent
.iter()
.map(|&p| xgb_calc_weight(p, reg) as f32)
.collect();
let scores: Vec<f64> = (0..n_bins)
.map(|b| {
(0..s).fold(0.0f64, |sc, t| {
let w = xgb_calc_weight(hist[(fs + b) * s + t], reg) as f32;
sc + f64::from(parent_w[t] * w)
})
})
.collect();
let mut sorted: Vec<usize> = (0..n_bins).collect();
sorted.sort_by(|&l, &r| scores[l].partial_cmp(&scores[r]).unwrap_or(Ordering::Equal));
for forward in [true, false] {
self.enumerate_part(bins, &sorted, forward, best);
}
}
fn enumerate_part(
&self,
bins: &FeatureBins,
sorted: &[usize],
forward: bool,
best: &mut Candidate,
) {
let FeatureBins {
nid,
hist,
feature: f,
fs,
fe,
} = *bins;
let s = self.n_split();
let parent = self.node_stats(nid);
let parent_gain = self.gain[nid];
let dir = self.b.config.cons.dir(f as usize);
let n_bins_feature = fe - fs;
let n_bins = MAX_CAT_THRESHOLD.min(n_bins_feature);
let mut left = vec![GradStats::default(); s];
let mut right = vec![GradStats::default(); s];
let mut local = Candidate::none();
let mut best_partition = None;
for step in 0..n_bins.saturating_sub(1) {
let j = if forward {
step
} else {
n_bins_feature - 1 - step
};
let bin = fs + sorted[j];
for t in 0..s {
if forward {
right[t].add(hist[bin * s + t]);
left[t] = parent[t].sub(right[t]);
} else {
left[t].add(hist[bin * s + t]);
right[t] = parent[t].sub(left[t]);
}
}
let loss = (self.split_gain(nid, dir, &left, &right) - parent_gain) as f32;
let placeholder = || SplitLocation::Numeric(SplitPos::BelowBins);
if local.update(loss, f, forward, placeholder, &left, &right) {
best_partition = Some(if forward { step + 1 } else { j });
}
}
if let Some(partition) = best_partition {
let cuts = self.ghist.cuts();
let mut cats: Vec<u32> = sorted[..partition]
.iter()
.map(|&c| cuts.cut_value(fs + c) as u32)
.collect();
cats.sort_unstable();
local.loc = SplitLocation::Categories(cats);
}
if local.is_categorical() {
best.merge(local);
}
}
fn apply(&mut self, entry: Entry, child_valid: bool) -> Expanded {
let s = self.n_split();
let nid = entry.nid;
let best = &entry.best;
let dir = self.b.config.cons.dir(best.feature as usize);
let mut w_left = Vec::with_capacity(s);
let mut w_right = Vec::with_capacity(s);
for t in 0..s {
let (l, r) = self.split_weights(nid, t, dir, best.left[t], best.right[t]);
w_left.push(l);
w_right.push(r);
}
let (route, (xgb_left, xgb_right)) = self.expand(nid, best);
self.store_children(nid, dir, best, [&w_left, &w_right], [xgb_left, xgb_right]);
let categorical = best.is_categorical();
let (rows_tree_left, rows_tree_right) = partition_rows(self.ghist, &entry.rows, route);
let rows = if categorical {
(rows_tree_right, rows_tree_left)
} else {
(rows_tree_left, rows_tree_right)
};
Expanded {
entry: Entry {
rows: Vec::new(),
..entry
},
children: (xgb_left, xgb_right),
rows,
child_valid,
}
}
fn expand<'c>(&mut self, nid: usize, best: &'c Candidate) -> (SplitRoute<'c>, (usize, usize)) {
let left_hess: f64 = best.left.iter().map(|g| g.hess).sum();
let right_hess: f64 = best.right.iter().map(|g| g.hess).sum();
let categorical = best.is_categorical();
let route = SplitRoute {
feature: best.feature,
location: &best.loc,
default_left: best.default_left != categorical,
};
let (tree_left_hess, tree_right_hess) = if categorical {
(right_hess, left_hess)
} else {
(left_hess, right_hess)
};
let (tree_left, tree_right) = self.tree.expand(
nid,
route.rule(self.ghist.cuts()),
ChildLeaf::new(0.0, tree_left_hess as f32),
ChildLeaf::new(0.0, tree_right_hess as f32),
);
self.tree.set_split_gain(nid, best.loss_chg);
self.tree.set_sum_hess(nid, (left_hess + right_hess) as f32);
let children = if categorical {
(tree_right, tree_left)
} else {
(tree_left, tree_right)
};
(route, children)
}
fn store_children(
&mut self,
nid: usize,
dir: i8,
best: &Candidate,
[w_left, w_right]: [&[f32]; 2],
[xgb_left, xgb_right]: [usize; 2],
) {
let s = self.n_split();
let n_nodes = self.tree.num_nodes();
self.stats.resize(n_nodes * s, GradStats::default());
self.weights.resize(n_nodes * s, 0.0);
self.gain.resize(n_nodes, 0.0);
let gl = self.b.gain_given_weights(&best.left, w_left);
let gr = self.b.gain_given_weights(&best.right, w_right);
for (child, stats, w, gain) in [
(xgb_left, &best.left, w_left, gl),
(xgb_right, &best.right, w_right, gr),
] {
self.stats[child * s..(child + 1) * s].copy_from_slice(stats);
self.weights[child * s..(child + 1) * s].copy_from_slice(w);
self.gain[child] = gain;
}
if self.constrained() {
self.lower.resize(n_nodes * s, -f32::MAX);
self.upper.resize(n_nodes * s, f32::MAX);
for t in 0..s {
let (lo, hi) = (self.lower[nid * s + t], self.upper[nid * s + t]);
let mid = w_left[t] + 0.5 * (w_right[t] - w_left[t]);
let (mut l_lo, mut l_hi, mut r_lo, mut r_hi) = (lo, hi, lo, hi);
if dir < 0 {
l_lo = mid;
r_hi = mid;
} else if dir > 0 {
l_hi = mid;
r_lo = mid;
}
self.lower[xgb_left * s + t] = l_lo;
self.upper[xgb_left * s + t] = l_hi;
self.lower[xgb_right * s + t] = r_lo;
self.upper[xgb_right * s + t] = r_hi;
}
}
}
fn children(&self, e: Expanded, features: [FeatureSet; 2]) -> [Entry; 2] {
let Expanded {
entry,
children: (xl, xr),
rows: (rows_l, rows_r),
..
} = e;
let b = &entry.best;
let left_hess: f64 = b.left.iter().map(|g| g.hess).sum();
let right_hess: f64 = b.right.iter().map(|g| g.hess).sum();
let build_right = right_hess < left_hess;
let built = self.build_hist(if build_right { &rows_r } else { &rows_l });
let (hist_l, hist_r) = with_sibling(entry.hist, built, !build_right);
let allowed = self
.b
.config
.next_allowed(entry.allowed.as_ref(), b.feature);
let [f_l, f_r] = features;
let n_rows = rows_l.len() + rows_r.len();
let eval_l = || self.evaluate(xl, &hist_l, &f_l, allowed.as_ref(), rows_l.len());
let eval_r = || self.evaluate(xr, &hist_r, &f_r, allowed.as_ref(), rows_r.len());
let (best_l, best_r) = if n_rows >= PARALLEL_EVALUATE_ROWS && rayon_available() {
rayon::join(eval_l, eval_r)
} else {
(eval_l(), eval_r())
};
let depth = entry.depth + 1;
let left = Entry {
nid: xl,
order: xl.min(xr),
depth,
rows: rows_l,
hist: hist_l,
best: best_l,
allowed: allowed.clone(),
};
let right = Entry {
nid: xr,
order: xl.max(xr),
depth,
rows: rows_r,
hist: hist_r,
best: best_r,
allowed,
};
[left, right]
}
fn expandable(&self, e: &Entry) -> bool {
let loss = e.best.loss_chg;
!(loss <= K_RT_EPS_F32
|| loss < self.b.config.params.gamma as f32
|| e.depth == limit_or_unbounded(self.b.config.params.max_depth)
|| self.num_leaves == limit_or_unbounded(self.b.config.params.max_leaves))
}
fn child_valid(&self, e: &Entry) -> bool {
!(e.depth + 1 >= limit_or_unbounded(self.b.config.params.max_depth)
|| self.num_leaves >= limit_or_unbounded(self.b.config.params.max_leaves))
}
fn pop(&mut self, queue: &mut Vec<Entry>) -> Vec<Entry> {
if queue.is_empty() {
return Vec::new();
}
if self.b.config.params.grow_policy == GrowPolicy::LossGuide {
let mut top = 0;
for (i, e) in queue.iter().enumerate().skip(1) {
let t = &queue[top];
if e.best.loss_chg > t.best.loss_chg
|| (e.best.loss_chg == t.best.loss_chg && e.order < t.order)
{
top = i;
}
}
let e = queue.swap_remove(top);
if self.expandable(&e) {
self.num_leaves += 1;
return vec![e];
}
self.record_leaf(e.nid, e.rows);
return Vec::new();
}
queue.sort_by_key(|e| std::cmp::Reverse(e.order));
let level = queue[queue.len() - 1].depth;
let mut result = Vec::new();
while let Some(e) = queue.pop_if(|e| e.depth == level) {
if self.expandable(&e) {
self.num_leaves += 1;
result.push(e);
} else {
self.record_leaf(e.nid, e.rows);
}
if result.len() >= MAX_NODE_BATCH {
break;
}
}
result
}
fn run(&mut self, queue: &mut Vec<Entry>, sampler: &mut ColumnSampler) {
let mut batch = self.pop(queue);
while !batch.is_empty() {
let mut expanded = Vec::with_capacity(batch.len());
for entry in batch {
let valid = self.child_valid(&entry);
expanded.push(self.apply(entry, valid));
}
let mut work = Vec::new();
for e in expanded {
if e.child_valid {
let depth = e.entry.depth + 1;
let features = [sampler.sample(depth), sampler.sample(depth)];
work.push((e, features));
} else {
let (xl, xr) = e.children;
let (rl, rr) = e.rows;
self.record_leaf(xl, rl);
self.record_leaf(xr, rr);
}
}
let parallel = work.len() > 1 && rayon_available();
let this = &*self;
let children: Vec<[Entry; 2]> = if parallel {
work.into_par_iter()
.map(|(e, f)| this.children(e, f))
.collect()
} else {
work.into_iter().map(|(e, f)| this.children(e, f)).collect()
};
for child in children.into_iter().flatten() {
if child.best.loss_chg > K_RT_EPS_F32 {
queue.push(child);
} else {
self.record_leaf(child.nid, child.rows);
}
}
batch = self.pop(queue);
}
}
fn record_leaf(&mut self, node: usize, rows: Vec<u32>) {
self.leaf_rows.push(LeafRows { node, rows });
}
fn finish(mut self) -> (RegTree, Vec<LeafRows>) {
self.leaf_rows.sort_by_key(|l| l.node);
let s = self.n_split();
let k = self.grad.n_outputs;
match self.grad.value {
None => {
for leaf in &self.leaf_rows {
let w = &self.weights[leaf.node * s..(leaf.node + 1) * s];
self.tree.set_leaf_vector(leaf.node, w);
}
}
Some(value) => {
let mut sums = vec![GradStats::default(); k];
let mut w = vec![0.0f32; k];
for leaf in &self.leaf_rows {
sums.fill(GradStats::default());
for &r in &leaf.rows {
for (acc, gp) in sums.iter_mut().zip(&value[r as usize * k..][..k]) {
acc.add(GradStats::from_pair(*gp));
}
}
for (w, st) in w.iter_mut().zip(&sums) {
*w = xgb_calc_weight(*st, &self.b.config.reg) as f32;
}
self.tree.set_leaf_vector(leaf.node, &w);
}
}
}
(self.tree, self.leaf_rows)
}
}
fn calc_gain(reg: &RegParams, st: GradStats) -> f64 {
if st.hess <= 0.0 {
return 0.0;
}
if reg.max_delta_step == 0.0 {
let t = threshold_l1(st.grad, reg.alpha);
t * t / (st.hess + reg.lambda)
} else {
let w = xgb_calc_weight(st, reg);
-(2.0 * st.grad * w + (st.hess + reg.lambda) * (w * w) + 2.0 * reg.alpha * w.abs())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::data::DMatrix;
use crate::tree::builder::{HistTreeBuilder, all_rows, test_support};
fn binned(x: &[f32], n: usize, cols: usize) -> GHistIndex {
test_support::binned(&DMatrix::from_dense(x, n, cols).unwrap(), 256)
}
fn build(
params: &TrainingParams,
ghist: &GHistIndex,
split: &[GradPair],
n_split: usize,
value: Option<&[GradPair]>,
n_outputs: usize,
) -> (RegTree, Vec<LeafRows>) {
let grad = VectorGradients {
split,
n_split,
value,
n_outputs,
};
MultiTreeBuilder::new(params).build(
ghist,
&grad,
&all_rows(ghist.n_rows()),
&mut ColumnSampler::all(ghist.n_cols()),
)
}
fn two_target_data(n: usize) -> (GHistIndex, Vec<GradPair>) {
let mut x = Vec::new();
let mut g = Vec::new();
for i in 0..n {
let a = (i % 17) as f32 / 17.0;
let b = (i % 11) as f32 / 11.0;
x.extend([a, b]);
g.push(GradPair::new(if a < 0.5 { -1.0 } else { 1.0 }, 1.0));
g.push(GradPair::new(if b < 0.3 { 2.0 } else { -0.5 }, 1.0));
}
(binned(&x, n, 2), g)
}
#[test]
fn duplicated_target_matches_scalar_structure() {
let (ghist, g2) = two_target_data(400);
let g1: Vec<GradPair> = g2.iter().step_by(2).copied().collect();
let dup: Vec<GradPair> = g1.iter().flat_map(|&g| [g, g]).collect();
let params = TrainingParams::builder().max_depth(3).build().unwrap();
let (vector, _) = build(¶ms, &ghist, &dup, 2, None, 2);
let scalar = HistTreeBuilder::new(¶ms).build(
&ghist,
&g1,
&all_rows(400),
&mut ColumnSampler::all(2),
);
assert_eq!(vector.num_nodes(), scalar.num_nodes());
for r in 0..400usize {
let row = [(r % 17) as f32 / 17.0, (r % 11) as f32 / 11.0];
let vl = vector.leaf_id_dense(&row, f32::NAN);
let sl = scalar.leaf_id_dense(&row, f32::NAN);
let w = vector.leaf_vector(vl);
assert_eq!(w[0], w[1]);
assert_eq!(w[0], scalar.node(sl).leaf_value);
}
}
#[test]
fn splits_serve_every_target() {
let (ghist, g) = two_target_data(400);
let params = TrainingParams::builder().max_depth(2).build().unwrap();
let (tree, leaves) = build(¶ms, &ghist, &g, 2, None, 2);
let features: Vec<u32> = tree
.nodes()
.iter()
.filter(|n| !n.is_leaf())
.map(|n| n.split_feature)
.collect();
assert!(features.contains(&0) && features.contains(&1));
let mut all: Vec<u32> = leaves.iter().flat_map(|l| l.rows.clone()).collect();
all.sort_unstable();
assert_eq!(all, (0..400).collect::<Vec<_>>());
assert_eq!(tree.node(0).sum_hess, 800.0);
}
#[test]
fn min_child_weight_uses_mean_hessian() {
let n = 60;
let x: Vec<f32> = (0..n).map(|i| i as f32).collect();
let ghist = binned(&x, n, 1);
let g: Vec<GradPair> = (0..n)
.flat_map(|i| {
let gr = if i < 30 { -1.0 } else { 1.0 };
[GradPair::new(gr, 1.0), GradPair::new(0.0, 0.0)]
})
.collect();
let loose = TrainingParams::builder()
.max_depth(1)
.min_child_weight(15.0)
.build()
.unwrap();
let (split, _) = build(&loose, &ghist, &g, 2, None, 2);
assert_eq!(split.num_nodes(), 3);
let strict = TrainingParams::builder()
.max_depth(1)
.min_child_weight(20.0)
.build()
.unwrap();
let (stump, _) = build(&strict, &ghist, &g, 2, None, 2);
assert_eq!(stump.num_nodes(), 1);
}
#[test]
fn reduced_gradients_refit_leaves_from_values() {
let (ghist, g) = two_target_data(400);
let mean: Vec<GradPair> = g
.as_chunks::<2>()
.0
.iter()
.map(|[a, b]| GradPair::new(f32::midpoint(a.grad, b.grad), 1.0))
.collect();
let params = TrainingParams::builder()
.max_depth(2)
.lambda(0.0)
.build()
.unwrap();
let (tree, leaves) = build(¶ms, &ghist, &mean, 1, Some(&g), 2);
assert_eq!(tree.size_leaf_vector(), 2);
for leaf in &leaves {
let n = leaf.rows.len() as f32;
for t in 0..2 {
let sum: f32 = leaf.rows.iter().map(|&r| g[r as usize * 2 + t].grad).sum();
let w = tree.leaf_vector(leaf.node)[t];
assert!((w + sum / n).abs() < 1e-5, "leaf {} target {t}", leaf.node);
}
}
}
#[test]
fn lossguide_respects_max_leaves() {
let (ghist, g) = two_target_data(400);
let params = TrainingParams::builder()
.grow_policy(GrowPolicy::LossGuide)
.unlimited_depth()
.max_leaves(3)
.build()
.unwrap();
let (tree, leaves) = build(¶ms, &ghist, &g, 2, None, 2);
assert_eq!(tree.num_leaves(), 3);
assert_eq!(leaves.len(), 3);
}
}