use std::collections::BTreeSet;
use crate::config::TrainingParams;
use crate::objective::GradPair;
use crate::tree::RegTree;
use crate::tree::constraints::{Bounds, MonotoneConstraints, calc_weight_bounded};
use crate::tree::gain::{GradStats, RegParams, threshold_l1};
pub(super) fn rayon_available() -> bool {
rayon::current_num_threads() > 1
}
pub(crate) struct LeafRows {
pub node: usize,
pub rows: Vec<u32>,
}
pub(super) struct BuilderConfig<'a> {
pub(super) params: &'a TrainingParams,
pub(super) reg: RegParams,
pub(super) cons: MonotoneConstraints,
interaction_sets: Option<Vec<Vec<u32>>>,
}
impl<'a> BuilderConfig<'a> {
pub(super) fn new(params: &'a TrainingParams) -> Self {
let groups = ¶ms.interaction_constraints;
let interaction_sets = (!groups.is_empty()).then(|| {
groups
.iter()
.map(|group| {
let mut group = group.clone();
group.sort_unstable();
group.dedup();
group
})
.collect()
});
BuilderConfig {
params,
reg: RegParams::from_params(params),
cons: MonotoneConstraints::from_params(¶ms.monotone_constraints),
interaction_sets,
}
}
pub(super) fn next_allowed(
&self,
parent: Option<&InteractionState>,
feature: u32,
) -> Option<InteractionState> {
let groups = self.interaction_sets.as_deref()?;
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(),
})
}
}
#[inline]
pub(crate) fn xgb_calc_weight(stats: GradStats, reg: &RegParams) -> f64 {
if stats.hess <= 0.0 {
return 0.0;
}
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
}
#[inline(always)]
pub(super) fn apply_bounds(w: f32, lower: f32, upper: f32) -> f32 {
if w < lower {
lower
} else if w > upper {
upper
} else {
w
}
}
#[inline]
pub(super) fn xgb_weight(stats: GradStats, reg: &RegParams, bounds: Bounds) -> f32 {
let w = xgb_calc_weight(stats, reg) as f32;
apply_bounds(w, bounds.lower as f32, bounds.upper as f32)
}
#[inline]
pub(super) 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
}
#[derive(Clone)]
pub(super) struct InteractionState {
path: Vec<u32>,
allowed: Vec<u32>,
}
pub(super) fn permits(state: Option<&InteractionState>, feature: u32) -> bool {
state.is_none_or(|state| state.allowed.binary_search(&feature).is_ok())
}
#[inline(never)]
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,
) {
let n_nodes = tree.num_nodes();
for (id, (&stats, &bounds)) in stats[..n_nodes].iter().zip(&bounds[..n_nodes]).enumerate() {
if tree.node(id).is_leaf() {
let w = calc_weight_bounded(stats, reg, bounds);
tree.set_leaf_value(id, w as f32);
}
}
}