use std::cell::RefCell;
use super::shared::{apply_bounds, xgb_gain_given_weight, xgb_weight};
use super::{Children, Score, SplitPos, children_valid};
use crate::tree::constraints::{Bounds, calc_weight_bounded, gain_at_weight, satisfies};
use crate::tree::gain::{GradStats, RegParams, calc_gain};
#[derive(Debug, Clone, Copy)]
pub(super) struct SplitScorer<'a> {
pub(super) reg: &'a RegParams,
pub(super) root_gain: f32,
pub(super) bounds: Bounds,
pub(super) dir: i8,
}
impl SplitScorer<'_> {
#[inline]
pub(super) fn loss_chg(&self, left: GradStats, right: GradStats) -> Option<Score<f32>> {
let reg = self.reg;
if !children_valid(left, right, reg.min_child_weight) {
return None;
}
let wl = xgb_weight(left, reg, self.bounds);
let wr = xgb_weight(right, reg, self.bounds);
if !satisfies(self.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(Score {
loss_chg: gain - self.root_gain,
w_left: wl,
w_right: wr,
})
}
#[inline]
pub(super) fn score_run<const ACC_LEFT: bool>(
&self,
total: GradStats,
acc_grad: &[f64],
acc_hess: &[f64],
loss: &mut [f32],
) {
match self.dir {
d if d > 0 => self.score_run_dir::<ACC_LEFT, 1>(total, acc_grad, acc_hess, loss),
d if d < 0 => self.score_run_dir::<ACC_LEFT, { -1 }>(total, acc_grad, acc_hess, loss),
_ => self.score_run_dir::<ACC_LEFT, 0>(total, acc_grad, acc_hess, loss),
}
}
#[inline(always)]
fn score_run_dir<const ACC_LEFT: bool, const DIR: i8>(
&self,
total: GradStats,
acc_grad: &[f64],
acc_hess: &[f64],
loss: &mut [f32],
) {
let RegParams {
lambda,
alpha,
max_delta_step,
min_child_weight,
} = *self.reg;
let (lower, upper) = (self.bounds.lower as f32, self.bounds.upper as f32);
let root_gain = self.root_gain;
let weight = |g: f64, h: f64| -> f32 {
let t = if g > alpha {
g - alpha
} else if g < -alpha {
g + alpha
} else {
0.0
};
let mut w = -t / (h + lambda);
if max_delta_step != 0.0 && w.abs() > max_delta_step {
w = max_delta_step.copysign(w);
}
apply_bounds(w as f32, lower, upper)
};
let gain = |g: f64, h: f64, w: f32| -> f64 {
-(2.0 * g * f64::from(w)
+ (h + lambda) * f64::from(w * w)
+ 2.0 * alpha * f64::from(w.abs()))
};
let n = loss.len();
let (acc_grad, acc_hess) = (&acc_grad[..n], &acc_hess[..n]);
for i in 0..n {
let (ag, ah) = (acc_grad[i], acc_hess[i]);
let (og, oh) = (total.grad - ag, total.hess - ah);
let (lg, lh, rg, rh) = if ACC_LEFT {
(ag, ah, og, oh)
} else {
(og, oh, ag, ah)
};
let valid = lh > 0.0 && rh > 0.0 && lh >= min_child_weight && rh >= min_child_weight;
let wl = weight(lg, lh);
let wr = weight(rg, rh);
let monotone = match DIR {
1 => f64::from(wl) <= f64::from(wr),
-1 => f64::from(wl) >= f64::from(wr),
_ => true,
};
let chg = (gain(lg, lh, wl) as f32 + gain(rg, rh, wr) as f32) - root_gain;
loss[i] = if valid && monotone {
chg
} else {
f32::NEG_INFINITY
};
}
}
#[inline]
pub(super) fn approx_exact(&self) -> bool {
let reg = self.reg;
self.dir == 0
&& self.bounds.lower == f64::NEG_INFINITY
&& self.bounds.upper == f64::INFINITY
&& reg.alpha == 0.0
&& reg.max_delta_step == 0.0
&& reg.lambda + reg.min_child_weight >= 1e-3
&& self.root_gain.is_finite()
}
#[cfg(test)]
pub(super) fn cannot_beat(&self, left: GradStats, right: GradStats, incumbent: f64) -> bool {
self.screen()
.is_some_and(|screen| screen.cannot_beat(left, right, incumbent))
}
#[inline]
pub(super) fn screen(&self) -> Option<Screen> {
self.approx_exact().then(|| Screen {
lambda: self.reg.lambda,
root: f64::from(self.root_gain),
})
}
}
#[derive(Debug, Clone, Copy)]
pub(super) struct Screen {
lambda: f64,
root: f64,
}
impl Screen {
#[cfg(test)]
pub(super) fn cannot_beat(self, left: GradStats, right: GradStats, incumbent: f64) -> bool {
let Screen { lambda, root } = self;
let (hl, hr) = (left.hess + lambda, right.hess + lambda);
if !(hl > 0.0 && hr > 0.0) {
return false;
}
let absolute = UNDERFLOW_MARGIN * (hl + hr + 1.0);
let bound = incumbent + root - SCREEN_KAPPA * (root.abs() + incumbent.abs()) - absolute;
let n = left.grad * left.grad * hr + right.grad * right.grad * hl;
n * (1.0 + SCREEN_KAPPA) < bound * (hl * hr)
}
#[inline]
pub(super) fn bound(self, total_hess: f64, incumbent: f64) -> ScreenBound {
const SLACK: f64 = 1.0 + 1.0 / 1_048_576.0;
let Screen { lambda, root } = self;
let allowance = UNDERFLOW_MARGIN * ((total_hess + 2.0 * lambda) * SLACK + 1.0) * SLACK;
ScreenBound {
lambda,
limit: incumbent + root - SCREEN_KAPPA * (root.abs() + incumbent.abs()) - allowance,
}
}
}
const SCREEN_KAPPA: f64 = 1.0 / 524_288.0;
#[derive(Debug, Clone, Copy)]
pub(super) struct ScreenBound {
lambda: f64,
limit: f64,
}
impl ScreenBound {
#[inline]
pub(super) fn rules_out(self, left: GradStats, right: GradStats) -> bool {
let (hl, hr) = (left.hess + self.lambda, right.hess + self.lambda);
if !(hl > 0.0 && hr > 0.0) {
return false;
}
let n = left.grad * left.grad * hr + right.grad * right.grad * hl;
n * (1.0 + SCREEN_KAPPA) < self.limit * (hl * hr)
}
}
impl SplitScorer<'_> {
#[inline]
pub(super) fn approx_run<const ACC_LEFT: bool>(
&self,
total: GradStats,
acc: &[GradStats],
approx: &mut [f32],
) {
if self.reg.min_child_weight > 0.0 {
self.approx_run_with::<ACC_LEFT, true>(total, acc, approx);
} else {
self.approx_run_with::<ACC_LEFT, false>(total, acc, approx);
}
}
#[inline(always)]
#[allow(
clippy::needless_bitwise_bool,
reason = "non-short-circuit validity keeps the loop branch-free so it vectorizes"
)]
fn approx_run_with<const ACC_LEFT: bool, const MCW_POSITIVE: bool>(
&self,
total: GradStats,
acc: &[GradStats],
approx: &mut [f32],
) {
let RegParams {
lambda,
min_child_weight,
..
} = *self.reg;
let root_gain = self.root_gain;
let gain = |g: f64, h: f64| -> f32 {
let g = g as f32;
g * (g / (h + lambda) as f32)
};
let valid_child = |h: f64| {
if MCW_POSITIVE {
h >= min_child_weight
} else {
h > 0.0
}
};
let n = approx.len();
let acc = &acc[..n];
for i in 0..n {
let (ag, ah) = (acc[i].grad, acc[i].hess);
let (og, oh) = (total.grad - ag, total.hess - ah);
let (lg, lh, rg, rh) = if ACC_LEFT {
(ag, ah, og, oh)
} else {
(og, oh, ag, ah)
};
let valid = valid_child(lh) & valid_child(rh);
let chg = (gain(lg, lh) + gain(rg, rh)) - root_gain;
approx[i] = if valid { chg } else { f32::NEG_INFINITY };
}
}
#[inline]
pub(super) fn candidate_gain(
&self,
left: GradStats,
right: GradStats,
constrained: bool,
) -> Option<Score> {
let reg = self.reg;
if left.hess < reg.min_child_weight || right.hess < reg.min_child_weight {
return None;
}
let parent = f64::from(self.root_gain);
if constrained {
let wl = calc_weight_bounded(left, reg, self.bounds);
let wr = calc_weight_bounded(right, reg, self.bounds);
if !satisfies(self.dir, wl, wr) {
return None;
}
let g = gain_at_weight(left, reg, wl) + gain_at_weight(right, reg, wr) - parent;
Some(Score {
loss_chg: g,
w_left: wl,
w_right: wr,
})
} else {
let g = calc_gain(left, reg) + calc_gain(right, reg) - parent;
Some(Score {
loss_chg: g,
w_left: 0.0,
w_right: 0.0,
})
}
}
}
#[inline]
pub(super) fn for_each_numeric_split(
bins: &[GradStats],
first: usize,
total: GradStats,
dense: bool,
mut offer: impl FnMut(SplitPos, Children),
) {
let mut acc = GradStats::default();
for (offset, &bin) in bins.iter().enumerate() {
acc.add(bin);
offer(
SplitPos::Bin(first + offset),
Children::new(false, acc, total.sub(acc)),
);
}
if dense || acc == total {
return;
}
let mut suffix = GradStats::default();
for offset in (0..bins.len()).rev() {
suffix.add(bins[offset]);
offer(
SplitPos::backward(first, offset),
Children::new(true, total.sub(suffix), suffix),
);
}
}
const SCAN_RUN: usize = 64;
pub(super) enum NumericScan {
Empty,
Best {
loss_chg: f32,
pos: SplitPos,
children: Children,
},
Nan,
}
pub(super) fn scan_numeric_splits(
bins: &[GradStats],
first: usize,
total: GradStats,
dense: bool,
scorer: &SplitScorer,
scratch: &mut ScanScratch,
) -> NumericScan {
if bins.len() <= FILTER_BINS
&& scorer.approx_exact()
&& let Some(scan) = scan_filtered(bins, first, total, dense, scorer, scratch)
{
return scan;
}
scan_batched(bins, first, total, dense, scorer)
}
fn scan_batched(
bins: &[GradStats],
first: usize,
total: GradStats,
dense: bool,
scorer: &SplitScorer,
) -> NumericScan {
let mut run = RunMax {
best: f32::NEG_INFINITY,
at: None,
};
let mut grad = [0f64; SCAN_RUN];
let mut hess = [0f64; SCAN_RUN];
let mut loss = [0f32; SCAN_RUN];
let mut found = None;
let mut acc = GradStats::default();
for (index, chunk) in bins.chunks(SCAN_RUN).enumerate() {
let n = chunk.len();
for (k, &bin) in chunk.iter().enumerate() {
acc.add(bin);
grad[k] = acc.grad;
hess[k] = acc.hess;
}
scorer.score_run::<true>(total, &grad[..n], &hess[..n], &mut loss[..n]);
if !run.scan(&loss[..n], &grad, &hess, index * SCAN_RUN) {
return NumericScan::Nan;
}
}
if let Some((offset, left)) = run.at.take() {
found = Some((
SplitPos::Bin(first + offset),
Children::new(false, left, total.sub(left)),
));
}
if !(dense || acc == total) {
let mut suffix = GradStats::default();
let mut end = bins.len();
let mut base = 0;
while end > 0 {
let n = end.min(SCAN_RUN);
for k in 0..n {
suffix.add(bins[end - 1 - k]);
grad[k] = suffix.grad;
hess[k] = suffix.hess;
}
scorer.score_run::<false>(total, &grad[..n], &hess[..n], &mut loss[..n]);
if !run.scan(&loss[..n], &grad, &hess, base) {
return NumericScan::Nan;
}
base += n;
end -= n;
}
if let Some((step, right)) = run.at {
let offset = bins.len() - 1 - step;
found = Some((
SplitPos::backward(first, offset),
Children::new(true, total.sub(right), right),
));
}
}
match found {
Some((pos, children)) => NumericScan::Best {
loss_chg: run.best,
pos,
children,
},
None => NumericScan::Empty,
}
}
const FILTER_BINS: usize = 256;
const FILTER_RUN: usize = 16;
const APPROX_MARGIN: f64 = 1.0 / 131_072.0;
const UNDERFLOW_MARGIN: f64 = f64::from_bits((1023 - 144) << 52);
fn scan_filtered(
bins: &[GradStats],
first: usize,
total: GradStats,
dense: bool,
scorer: &SplitScorer,
scratch: &mut ScanScratch,
) -> Option<NumericScan> {
let n = bins.len();
let mut prefix = Prefix::default();
for (&bin, s) in bins.iter().zip(&mut scratch.acc[..n]) {
prefix.add(bin);
*s = prefix.acc;
}
scan_filtered_rest(bins, first, total, dense, scorer, scratch, prefix)
}
#[derive(Clone, Copy)]
struct Prefix {
acc: GradStats,
nonneg: bool,
}
impl Default for Prefix {
fn default() -> Self {
Prefix {
acc: GradStats::default(),
nonneg: true,
}
}
}
impl Prefix {
#[inline(always)]
fn add(&mut self, bin: GradStats) {
self.acc.add(bin);
self.nonneg &= bin.hess >= 0.0;
}
}
pub(super) fn scan_numeric_pair(
a: &NumericInput,
b: &NumericInput,
scratch: [&mut ScanScratch; 2],
) -> [NumericScan; 2] {
let [sa, sb] = scratch;
let filtered = |x: &NumericInput| x.bins.len() <= FILTER_BINS && x.scorer.approx_exact();
if !(filtered(a) && filtered(b)) {
return [a.scan(sa), b.scan(sb)];
}
let (na, nb) = (a.bins.len(), b.bins.len());
let common = na.min(nb);
let (mut acc_a, mut acc_b) = (Prefix::default(), Prefix::default());
{
let (sa, sb) = (&mut sa.acc[..common], &mut sb.acc[..common]);
let (ba, bb) = (&a.bins[..common], &b.bins[..common]);
for i in 0..common {
acc_a.add(ba[i]);
acc_b.add(bb[i]);
(sa[i], sb[i]) = (acc_a.acc, acc_b.acc);
}
}
for (x, s, acc) in [(a, &mut *sa, &mut acc_a), (b, &mut *sb, &mut acc_b)] {
let n = x.bins.len();
for i in common..n {
acc.add(x.bins[i]);
s.acc[i] = acc.acc;
}
}
let finish = |x: &NumericInput, s: &mut ScanScratch, acc| {
scan_filtered_rest(x.bins, x.first, x.total, x.dense, &x.scorer, s, acc)
.unwrap_or_else(|| scan_batched(x.bins, x.first, x.total, x.dense, &x.scorer))
};
[finish(a, sa, acc_a), finish(b, sb, acc_b)]
}
pub(super) struct NumericInput<'a> {
pub(super) bins: &'a [GradStats],
pub(super) first: usize,
pub(super) total: GradStats,
pub(super) dense: bool,
pub(super) scorer: SplitScorer<'a>,
}
impl NumericInput<'_> {
pub(super) fn scan(&self, scratch: &mut ScanScratch) -> NumericScan {
scan_numeric_splits(
self.bins,
self.first,
self.total,
self.dense,
&self.scorer,
scratch,
)
}
}
#[allow(
clippy::needless_bitwise_bool,
reason = "branch-free overflow and threshold tests vectorize"
)]
fn scan_filtered_rest(
bins: &[GradStats],
first: usize,
total: GradStats,
dense: bool,
scorer: &SplitScorer,
scratch: &mut ScanScratch,
prefix: Prefix,
) -> Option<NumericScan> {
let Prefix { acc, nonneg } = prefix;
let n = bins.len();
let ScanScratch { acc: stats, approx } = scratch;
let (stats, approx) = (&mut stats[..2 * n], &mut approx[..2 * n]);
let extremes = |sums: &[GradStats]| match (nonneg, sums.first(), sums.last()) {
(true, Some(first), Some(last)) => (first.hess, last.hess),
_ => min_max_hess(sums),
};
let (mut hess_lo, mut hess_hi) = extremes(&stats[..n]);
scorer.approx_run::<true>(total, &stats[..n], &mut approx[..n]);
let mut m = n;
if !(dense || acc == total) {
let mut suffix = GradStats::default();
for (&bin, s) in bins.iter().rev().zip(&mut stats[n..]) {
suffix.add(bin);
*s = suffix;
}
let (lo, hi) = extremes(&stats[n..]);
hess_lo = hess_lo.min(lo);
hess_hi = hess_hi.max(hi);
scorer.approx_run::<false>(total, &stats[n..], &mut approx[n..]);
m = 2 * n;
}
let mut max = f32::NEG_INFINITY;
let mut overflow = false;
for &a in &approx[..m] {
overflow |= a.is_nan() | (a == f32::INFINITY);
max = max.max(a);
}
if overflow {
return None;
}
if max == f32::NEG_INFINITY {
return Some(NumericScan::Empty);
}
let root = f64::from(scorer.root_gain);
let scale = (f64::from(max) + root + root.abs()).max(0.0);
let max_d = hess_hi.max(total.hess - hess_lo) + scorer.reg.lambda;
if scale >= 1e30 || max_d.is_nan() || max_d >= 1e30 {
return None;
}
let absolute = UNDERFLOW_MARGIN * (2.0 * max_d + 1.0);
let threshold = f64::from(max) - 2.0 * (APPROX_MARGIN * scale + absolute);
let mut cutoff = threshold as f32;
if f64::from(cutoff) > threshold {
cutoff = cutoff.next_down();
}
let mut best = f32::NEG_INFINITY;
let mut found = None;
for (chunk, run) in approx[..m].chunks(FILTER_RUN).enumerate() {
if !run.iter().fold(false, |any, &a| any | (a >= cutoff)) {
continue;
}
for (k, &a) in run.iter().enumerate() {
if a < cutoff || f64::from(a) < threshold {
continue;
}
let i = chunk * FILTER_RUN + k;
let acc = stats[i];
let (pos, children) = if i < n {
(
SplitPos::Bin(first + i),
Children::new(false, acc, total.sub(acc)),
)
} else {
let offset = n - 1 - (i - n);
(
SplitPos::backward(first, offset),
Children::new(true, total.sub(acc), acc),
)
};
let Some(score) = scorer.loss_chg(children.left, children.right) else {
continue;
};
let l = score.loss_chg;
if l.is_nan() {
return Some(NumericScan::Nan);
}
if l > best && l.is_finite() {
best = l;
found = Some((pos, children));
}
}
}
Some(match found {
Some((pos, children)) => NumericScan::Best {
loss_chg: best,
pos,
children,
},
None => NumericScan::Empty,
})
}
#[inline]
fn min_max_hess(stats: &[GradStats]) -> (f64, f64) {
let mut lo = [f64::INFINITY; 4];
let mut hi = [f64::NEG_INFINITY; 4];
let (quads, rest) = stats.as_chunks::<4>();
for quad in quads {
for k in 0..4 {
lo[k] = lo[k].min(quad[k].hess);
hi[k] = hi[k].max(quad[k].hess);
}
}
for (k, s) in rest.iter().enumerate() {
lo[k] = lo[k].min(s.hess);
hi[k] = hi[k].max(s.hess);
}
(
lo[0].min(lo[1]).min(lo[2].min(lo[3])),
hi[0].max(hi[1]).max(hi[2].max(hi[3])),
)
}
thread_local! {
static SCAN_SCRATCH: RefCell<[ScanScratch; 2]> =
RefCell::new([ScanScratch::new(), ScanScratch::new()]);
}
pub(super) fn with_scan_scratch<R>(f: impl FnOnce(&mut [ScanScratch; 2]) -> R) -> R {
SCAN_SCRATCH.with(|cell| match cell.try_borrow_mut() {
Ok(mut scratch) => f(&mut scratch),
Err(_) => f(&mut [ScanScratch::new(), ScanScratch::new()]),
})
}
pub(super) struct ScanScratch {
acc: Vec<GradStats>,
approx: Vec<f32>,
}
impl ScanScratch {
pub(super) fn new() -> Self {
ScanScratch {
acc: vec![GradStats::default(); 2 * FILTER_BINS],
approx: vec![0.0; 2 * FILTER_BINS],
}
}
}
struct RunMax {
best: f32,
at: Option<(usize, GradStats)>,
}
impl RunMax {
#[inline]
fn scan(&mut self, loss: &[f32], grad: &[f64], hess: &[f64], base: usize) -> bool {
for (k, &l) in loss.iter().enumerate() {
if l.is_nan() {
return false;
}
if l > self.best && l.is_finite() {
self.best = l;
self.at = Some((base + k, GradStats::new(grad[k], hess[k])));
}
}
true
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::objective::GradPair;
use crate::tree::builder::shared::xgb_node_gain;
use crate::tree::builder::{BestSplit, SplitLocation, xgb_update};
fn sequential_and_scanned(
bins: &[GradStats],
total: GradStats,
dense: bool,
scorer: &SplitScorer,
) -> (BestSplit, BestSplit) {
let offer = |best: &mut BestSplit, pos, children: Children| {
if let Some(score) = scorer.loss_chg(children.left, children.right) {
xgb_update(best, 0, pos, children, score);
}
};
let mut expected = BestSplit::none();
let mut nan = false;
for_each_numeric_split(bins, 0, total, dense, |pos, children| {
nan |= scorer
.loss_chg(children.left, children.right)
.is_some_and(|score| score.loss_chg.is_nan());
offer(&mut expected, pos, children);
});
let mut actual = BestSplit::none();
match scan_numeric_splits(bins, 0, total, dense, scorer, &mut ScanScratch::new()) {
NumericScan::Empty => {}
NumericScan::Best { pos, children, .. } => offer(&mut actual, pos, children),
NumericScan::Nan => {
assert!(nan, "no candidate scores NaN");
actual = expected.clone();
}
}
(expected, actual)
}
fn numeric_bin(b: &BestSplit) -> Option<usize> {
match &b.location {
SplitLocation::Numeric(pos) => pos.bin(),
SplitLocation::Categories(categories) => panic!("categorical split {categories:?}"),
}
}
fn split_key(b: &BestSplit) -> (u64, Option<usize>, bool, [u64; 4]) {
(
b.loss_chg.to_bits(),
numeric_bin(b),
b.default_left,
[b.left.grad, b.left.hess, b.right.grad, b.right.hess].map(f64::to_bits),
)
}
fn random_bins(
rng: &mut crate::rng::Rng,
n: usize,
(grad_scale, hess_scale): (f32, f32),
repeat: bool,
) -> Vec<GradStats> {
let mut bins = vec![GradStats::default(); n];
for bin in &mut bins {
if rng.below(4) == 0 {
continue;
}
for _ in 0..rng.range(1..20) {
let g = (rng.f32() * 4.0 - 2.0) * grad_scale;
let h = (0.05 + rng.f32()) * hess_scale;
bin.add(GradStats::from_pair(GradPair::new(g, h)));
}
}
if repeat {
let half = n / 2;
for i in 0..half {
bins[n - 1 - i] = bins[i];
}
}
bins
}
#[test]
fn numeric_scan_matches_sequential_search() {
let mut rng = crate::rng::Rng::new(7);
for trial in 0..4000 {
let n = 2 + rng.range(0..300);
let bins = random_bins(&mut rng, n, (1.0, 1.0), trial % 5 == 0);
let mut total = GradStats::default();
for &bin in &bins {
total.add(bin);
}
let dense = trial % 3 == 0;
if !dense && rng.below(2) == 0 {
total.add(GradStats::new(f64::from(rng.f32()) * 8.0 - 4.0, 3.0));
}
let reg = RegParams {
lambda: [0.0, 1.0, 0.1][rng.range(0..3)],
alpha: if rng.below(6) == 0 { 0.5 } else { 0.0 },
max_delta_step: if rng.below(6) == 0 { 0.7 } else { 0.0 },
min_child_weight: [0.0, 1.0, 5.0][rng.range(0..3)],
};
let dir = [0, 0, 0, 1, -1][rng.range(0..5)];
let scorer = SplitScorer {
reg: ®,
root_gain: xgb_node_gain(total, ®, Bounds::default()),
bounds: Bounds::default(),
dir,
};
let (expected, actual) = sequential_and_scanned(&bins, total, dense, &scorer);
assert_eq!(split_key(&actual), split_key(&expected), "trial {trial}");
}
}
#[test]
fn numeric_scan_matches_sequential_search_at_extreme_scales() {
let mut rng = crate::rng::Rng::new(13);
let scales = [1e-30f32, 1e-20, 1e-10, 1e-3, 1.0, 1e10, 1e20, 1e30];
for trial in 0..4000 {
let n = 2 + rng.range(0..40);
let grad_scale = scales[rng.range(0..scales.len())];
let hess_scale = [1.0f32, 1e10, 1e20, 1e30, 1e36, 1e38][rng.range(0..6)];
let bins = random_bins(&mut rng, n, (grad_scale, hess_scale), trial % 5 == 0);
let mut total = GradStats::default();
for &bin in &bins {
total.add(bin);
}
let dense = trial % 3 == 0;
if !dense && rng.below(2) == 0 {
let g = f64::from(rng.f32() * 8.0 - 4.0) * f64::from(grad_scale);
total.add(GradStats::new(g, 3.0 * f64::from(hess_scale)));
}
let reg = RegParams {
lambda: [1.0, 0.1][rng.range(0..2)],
alpha: 0.0,
max_delta_step: 0.0,
min_child_weight: [0.0, 1.0][rng.range(0..2)],
};
let scorer = SplitScorer {
reg: ®,
root_gain: xgb_node_gain(total, ®, Bounds::default()),
bounds: Bounds::default(),
dir: 0,
};
let (expected, actual) = sequential_and_scanned(&bins, total, dense, &scorer);
assert_eq!(split_key(&actual), split_key(&expected), "trial {trial}");
}
}
fn subnormal_weight_square_bins() -> (Vec<GradStats>, GradStats, RegParams) {
let bins: Vec<GradStats> = [-1.02e-22f32, -1e-24, 1.03e-22]
.iter()
.map(|&label| GradStats::from_pair(GradPair::new((0.0 - label) * 1e38, 1e38)))
.collect();
let mut total = GradStats::default();
for &bin in &bins {
total.add(bin);
}
let reg = RegParams {
lambda: 1.0,
alpha: 0.0,
max_delta_step: 0.0,
min_child_weight: 1.0,
};
(bins, total, reg)
}
#[test]
fn numeric_scan_keeps_the_winner_when_weights_square_to_subnormals() {
let (bins, total, reg) = subnormal_weight_square_bins();
let scorer = SplitScorer {
reg: ®,
root_gain: xgb_node_gain(total, ®, Bounds::default()),
bounds: Bounds::default(),
dir: 0,
};
assert!(scorer.approx_exact());
let (expected, actual) = sequential_and_scanned(&bins, total, true, &scorer);
assert_eq!(numeric_bin(&expected), Some(0));
assert_eq!(split_key(&actual), split_key(&expected));
}
#[test]
fn cannot_beat_allows_for_subnormal_weight_squares() {
let (bins, total, reg) = subnormal_weight_square_bins();
let scorer = SplitScorer {
reg: ®,
root_gain: xgb_node_gain(total, ®, Bounds::default()),
bounds: Bounds::default(),
dir: 0,
};
let (left, right) = (bins[0], total.sub(bins[0]));
let exact = scorer.loss_chg(left, right).unwrap().loss_chg;
for incumbent in [exact * 0.98, exact * 0.99, exact.next_down()] {
assert!(
!scorer.cannot_beat(left, right, f64::from(incumbent)),
"{incumbent} < {exact}"
);
}
}
#[test]
fn cannot_beat_is_conservative() {
let mut rng = crate::rng::Rng::new(11);
let mut ruled_out = 0;
for _ in 0..20_000 {
let scale = [1e-3f64, 1.0, 1e3][rng.range(0..3)];
let stats = |rng: &mut crate::rng::Rng| {
GradStats::new((rng.f64() * 2.0 - 1.0) * scale, 0.01 + rng.f64() * scale)
};
let (left, right) = (stats(&mut rng), stats(&mut rng));
let total = GradStats::new(left.grad + right.grad, left.hess + right.hess);
let reg = RegParams {
lambda: [1.0, 0.1][rng.range(0..2)],
alpha: 0.0,
max_delta_step: 0.0,
min_child_weight: 0.0,
};
let scorer = SplitScorer {
reg: ®,
root_gain: xgb_node_gain(total, ®, Bounds::default()),
bounds: Bounds::default(),
dir: 0,
};
let Some(score) = scorer.loss_chg(left, right) else {
continue;
};
let exact = score.loss_chg;
let mut incumbent = exact;
for _ in 0..4 {
incumbent = incumbent.next_down();
}
for _ in 0..8 {
let beats = exact > incumbent;
if scorer.cannot_beat(left, right, f64::from(incumbent)) {
assert!(!beats, "{left:?} {right:?} {incumbent} {exact}");
}
incumbent = incumbent.next_up();
}
if scorer.cannot_beat(left, right, f64::from(exact) + f64::from(exact.abs()) + 1.0) {
ruled_out += 1;
}
}
assert!(ruled_out > 10_000, "{ruled_out}");
}
#[test]
fn screen_bound_implies_cannot_beat() {
let mut rng = crate::rng::Rng::new(23);
let mut ruled_out = 0;
for _ in 0..20_000 {
let g_scale = [1e-30f64, 1e-3, 1.0, 1e3, 1e30][rng.range(0..5)];
let h_scale = [1e-30f64, 1e-3, 1.0, 1e3, 1e30][rng.range(0..5)];
let total = GradStats::new(
(rng.f64() * 2.0 - 1.0) * g_scale,
0.01 * h_scale + rng.f64() * h_scale,
);
let left = GradStats::new((rng.f64() * 2.0 - 1.0) * g_scale, rng.f64() * total.hess);
let right = total.sub(left);
if right.hess < 0.0 {
continue;
}
let lambda = [1.0, 1e-3, 0.0][rng.range(0..3)];
let reg = RegParams {
lambda,
alpha: 0.0,
max_delta_step: 0.0,
min_child_weight: 1e-3,
};
let scorer = SplitScorer {
reg: ®,
root_gain: xgb_node_gain(total, ®, Bounds::default()),
bounds: Bounds::default(),
dir: 0,
};
let Some(screen) = scorer.screen() else {
continue;
};
for (l, r) in [(left, right), (right, left)] {
let Some(score) = scorer.loss_chg(l, r) else {
continue;
};
let mut incumbent = score.loss_chg;
for _ in 0..4 {
incumbent = incumbent.next_down();
}
for _ in 0..8 {
let incumbent64 = f64::from(incumbent);
if screen.bound(total.hess, incumbent64).rules_out(l, r) {
assert!(
screen.cannot_beat(l, r, incumbent64),
"{l:?} {r:?} {total:?} {incumbent}"
);
}
incumbent = incumbent.next_up();
}
let far = f64::from(score.loss_chg) + f64::from(score.loss_chg.abs()) + 1.0;
if screen.bound(total.hess, far).rules_out(l, r) {
assert!(screen.cannot_beat(l, r, far));
ruled_out += 1;
}
}
}
assert!(ruled_out > 10_000, "{ruled_out}");
}
#[test]
fn cannot_beat_is_conservative_at_extreme_scales() {
let mut rng = crate::rng::Rng::new(17);
let grad_scales = [1e-30f64, 1e-20, 1e-10, 1.0, 1e10, 1e20, 1e30];
let hess_scales = [1.0f64, 1e10, 1e20, 1e30, 1e38];
for _ in 0..20_000 {
let grad_scale = grad_scales[rng.range(0..grad_scales.len())];
let hess_scale = hess_scales[rng.range(0..hess_scales.len())];
let mut stats = || {
GradStats::new(
(rng.f64() * 2.0 - 1.0) * grad_scale,
(0.01 + rng.f64()) * hess_scale,
)
};
let (left, right) = (stats(), stats());
let total = GradStats::new(left.grad + right.grad, left.hess + right.hess);
let reg = RegParams {
lambda: [1.0, 0.1][rng.range(0..2)],
alpha: 0.0,
max_delta_step: 0.0,
min_child_weight: 0.0,
};
let scorer = SplitScorer {
reg: ®,
root_gain: xgb_node_gain(total, ®, Bounds::default()),
bounds: Bounds::default(),
dir: 0,
};
let Some(score) = scorer.loss_chg(left, right) else {
continue;
};
let exact = score.loss_chg;
if !exact.is_finite() {
continue;
}
let mut incumbent = exact;
for _ in 0..8 {
incumbent = incumbent.next_down();
assert!(
!scorer.cannot_beat(left, right, f64::from(incumbent)),
"{left:?} {right:?} {incumbent} {exact}"
);
}
}
}
}