use super::{HistTreeBuilder, NodeCtx};
use crate::data::ghist::GHistIndex;
use crate::data::quantile::HistCuts;
use crate::tree::builder::categorical::sweep_categorical;
use crate::tree::builder::shared::{InteractionState, permits, rayon_available, xgb_node_gain};
use crate::tree::builder::split::{
NumericInput, NumericScan, SplitScorer, for_each_numeric_split, scan_numeric_pair,
with_scan_scratch,
};
use crate::tree::builder::{BestSplit, need_replace, xgb_update};
use crate::tree::gain::GradStats;
use crate::tree::reuse::CategoricalPenalty;
use rayon::prelude::*;
const PARALLEL_SCAN_BINS: usize = 4096;
const SCAN_TASK_BINS: usize = 2048;
impl HistTreeBuilder<'_> {
pub(super) fn evaluate(
&self,
ghist: &GHistIndex,
hist: &[GradStats],
feature_subset: &[u32],
allowed: Option<&InteractionState>,
node: NodeCtx,
) -> BestSplit {
let filtered: Vec<u32>;
let feature_subset: &[u32] = match allowed {
Some(_) => {
filtered = feature_subset
.iter()
.copied()
.filter(|&f| permits(allowed, f))
.collect();
&filtered
}
None => feature_subset,
};
if let Some(options) = &self.options {
return options.evaluate(
ghist,
hist,
feature_subset,
&self.config.cons,
&self.config.reg,
node,
);
}
let cuts = ghist.cuts();
let mut best = BestSplit::none();
let dense = ghist.dense_stride().is_some();
let total = node.stats;
let node_scorer = SplitScorer {
reg: &self.config.reg,
root_gain: xgb_node_gain(total, &self.config.reg, node.bounds),
bounds: node.bounds,
dir: 0,
};
let input = |f: u32| {
let (fs, fe) = cuts.feature_bins(f as usize);
NumericInput {
bins: &hist[fs..fe],
first: fs,
total,
dense,
scorer: SplitScorer {
dir: self.config.cons.dir(f as usize),
..node_scorer
},
}
};
let scan_chunk = |chunk: &[u32]| -> Vec<Option<NumericScan>> {
let plain = |f: u32| {
let (fs, fe) = cuts.feature_bins(f as usize);
!cuts.is_categorical(f as usize) && fe > fs + 1
};
let mut out: Vec<Option<NumericScan>> = chunk.iter().map(|_| None).collect();
with_scan_scratch(|[sa, sb]| {
let mut pending = None;
for (i, &f) in chunk.iter().enumerate() {
if !plain(f) {
continue;
}
match pending.take() {
None => pending = Some(i),
Some(j) => {
let [x, y] = scan_numeric_pair(&input(chunk[j]), &input(f), [sa, sb]);
(out[j], out[i]) = (Some(x), Some(y));
}
}
}
if let Some(j) = pending {
out[j] = Some(input(chunk[j]).scan(sa));
}
});
out
};
let mut scans = if self.reuse.is_some() {
None
} else {
Some(
Self::parallel_scans(cuts, feature_subset, scan_chunk)
.unwrap_or_else(|| scan_chunk(feature_subset)),
)
};
for (i, &f) in feature_subset.iter().enumerate() {
let (fs, fe) = cuts.feature_bins(f as usize);
let scorer = SplitScorer {
dir: self.config.cons.dir(f as usize),
..node_scorer
};
if cuts.is_categorical(f as usize) {
let cats: Vec<(u32, GradStats)> = (fs..fe)
.map(|i| (cuts.cut_value(i) as u32, hist[i]))
.collect();
sweep_categorical(
&mut best,
&cats,
total,
&scorer,
f,
self.reuse.as_ref().map(|r| r as &dyn CategoricalPenalty),
);
continue;
}
if fe <= fs + 1 {
continue; }
if let Some(reuse) = &self.reuse {
for_each_numeric_split(&hist[fs..fe], fs, total, dense, |pos, children| {
let Some(mut score) = scorer.loss_chg(children.left, children.right) else {
return;
};
score.loss_chg -= reuse.bin_penalty(f, pos.bin());
xgb_update(&mut best, f, pos, children, score);
});
continue;
}
let scanned = scans.as_mut().and_then(|scans| scans[i].take());
match scanned.unwrap_or_else(|| with_scan_scratch(|[s, _]| input(f).scan(s))) {
NumericScan::Empty => {}
NumericScan::Best {
loss_chg,
pos,
children,
} => {
if need_replace(best.loss_chg as f32, best.feature, loss_chg, f)
&& let Some(score) = scorer.loss_chg(children.left, children.right)
{
xgb_update(&mut best, f, pos, children, score);
}
}
NumericScan::Nan => {
for_each_numeric_split(&hist[fs..fe], fs, total, dense, |pos, children| {
if let Some(score) = scorer.loss_chg(children.left, children.right) {
xgb_update(&mut best, f, pos, children, score);
}
});
}
}
}
best
}
fn parallel_scans(
cuts: &HistCuts,
feature_subset: &[u32],
scan_chunk: impl Fn(&[u32]) -> Vec<Option<NumericScan>> + Sync,
) -> Option<Vec<Option<NumericScan>>> {
if !rayon_available() {
return None;
}
let bins = |f: u32| {
let (fs, fe) = cuts.feature_bins(f as usize);
fe - fs
};
let candidates: usize = feature_subset.iter().map(|&f| bins(f)).sum();
if candidates < PARALLEL_SCAN_BINS {
return None;
}
let per_task = (SCAN_TASK_BINS * feature_subset.len()).div_ceil(candidates);
Some(
feature_subset
.par_chunks(per_task.max(1))
.flat_map_iter(&scan_chunk)
.collect(),
)
}
}