use super::shared::xgb_node_gain;
use super::split::{SplitScorer, for_each_numeric_split};
use crate::data::quantile::HistCuts;
use crate::tree::constraints::Bounds;
use crate::tree::gain::{GradStats, RegParams};
#[derive(Debug, Clone, Copy)]
pub(crate) struct SplitRank {
pub(crate) candidates: usize,
pub(crate) better: Option<usize>,
pub(crate) loss_chg: f32,
}
pub(crate) fn rank_split(
reg: &RegParams,
cuts: &HistCuts,
hist: &[GradStats],
total: GradStats,
dense: bool,
(feature, split_cond, default_left): (u32, f32, bool),
) -> SplitRank {
let bounds = Bounds::default();
let scorer = SplitScorer {
reg,
root_gain: xgb_node_gain(total, reg, bounds),
bounds,
dir: 0,
};
let mut gains = Vec::new();
let mut current = None;
for f in 0..cuts.n_features() {
let (fs, fe) = cuts.feature_bins(f);
if cuts.is_categorical(f) || fe <= fs + 1 {
continue;
}
for_each_numeric_split(&hist[fs..fe], fs, total, dense, |pos, children| {
let Some(score) = scorer.loss_chg(children.left, children.right) else {
return;
};
if !score.loss_chg.is_finite() {
return;
}
gains.push(score.loss_chg);
if f as u32 == feature
&& pos.threshold(cuts) == split_cond
&& children.default_left == default_left
{
current = Some(score.loss_chg);
}
});
}
SplitRank {
candidates: gains.len(),
better: current.map(|c| gains.iter().filter(|&&g| g > c).count()),
loss_chg: current.unwrap_or(0.0),
}
}