use super::shared::xgb_calc_weight;
use super::split::SplitScorer;
use super::{BestSplit, Children, Score, need_replace};
use crate::tree::gain::GradStats;
use crate::tree::reuse::CategoricalPenalty;
pub(super) const MAX_CAT_TO_ONEHOT: usize = 4;
pub(super) const MAX_CAT_THRESHOLD: usize = 64;
struct CatCandidate {
children: Children,
score: Score<f32>,
}
pub(super) fn sweep_categorical(
best: &mut BestSplit,
cats: &[(u32, GradStats)],
total: GradStats,
scorer: &SplitScorer,
feature: u32,
penalty: Option<&dyn CategoricalPenalty>,
) {
let n = cats.len();
let reg = scorer.reg;
let score = |children: Children, set: &dyn Fn() -> Vec<u32>| {
let mut score = scorer.loss_chg(children.left, children.right)?;
if let Some(penalty) = penalty {
score.loss_chg -= penalty.categorical_penalty(feature, &set()) as f32;
}
Some(score)
};
let offer =
|local: &mut Option<CatCandidate>, children: Children, set: &dyn Fn() -> Vec<u32>| {
let Some(score) = score(children, set) else {
return false;
};
let incumbent = local.as_ref().map_or(0.0, |c| c.score.loss_chg);
if !need_replace(incumbent, feature, score.loss_chg, feature) {
return false;
}
*local = Some(CatCandidate { children, score });
true
};
let merge = |best: &mut BestSplit, local: Option<CatCandidate>, mut set: Vec<u32>| {
let Some(CatCandidate { children, score }) = local else {
return;
};
if need_replace(best.loss_chg as f32, best.feature, score.loss_chg, feature) {
set.sort_unstable();
*best =
BestSplit::categorical(feature, children.swapped(), score.swapped().into(), set);
best.children_swapped = true;
}
};
if n < MAX_CAT_TO_ONEHOT {
let mut present = GradStats::default();
for &(_, stats) in cats {
present.add(stats);
}
let missing = total.sub(present);
let mut local = None;
let mut chosen = 0;
for &(cat, stats) in cats {
let single = || vec![cat];
let mut right = stats;
let missing_left = Children::new(true, total.sub(right), right);
if offer(&mut local, missing_left, &single) {
chosen = cat;
}
right.add(missing);
let missing_right = Children::new(false, total.sub(right), right);
if offer(&mut local, missing_right, &single) {
chosen = cat;
}
}
merge(best, local, vec![chosen]);
return;
}
let weight = |s: GradStats| -> f32 {
if s.hess < reg.min_child_weight {
0.0
} else {
xgb_calc_weight(s, reg) as f32
}
};
let keys: Vec<f32> = cats.iter().map(|&(_, s)| weight(s)).collect();
let mut sorted: Vec<usize> = (0..n).collect();
sorted.sort_by(|&l, &r| {
keys[l]
.partial_cmp(&keys[r])
.unwrap_or(std::cmp::Ordering::Equal)
});
let set_of =
|partition: usize| -> Vec<u32> { sorted[..partition].iter().map(|&c| cats[c].0).collect() };
let depth = MAX_CAT_THRESHOLD.min(n);
for forward in [true, false] {
let mut acc = GradStats::default();
let mut local = None;
let mut best_partition = 0;
for step in 0..depth - 1 {
let j = if forward { step } else { n - 1 - step };
acc.add(cats[sorted[j]].1);
let partition = if forward { step + 1 } else { j };
let (left, right) = if forward {
(total.sub(acc), acc)
} else {
(acc, total.sub(acc))
};
if offer(&mut local, Children::new(forward, left, right), &|| {
set_of(partition)
}) {
best_partition = partition;
}
}
merge(best, local, set_of(best_partition));
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tree::builder::SplitLocation;
use crate::tree::builder::shared::xgb_node_gain;
use crate::tree::constraints::Bounds;
use crate::tree::gain::RegParams;
fn unregularized() -> RegParams {
RegParams {
lambda: 0.0,
alpha: 0.0,
max_delta_step: 0.0,
min_child_weight: 0.0,
}
}
fn sweep(cats: &[(u32, GradStats)], total: GradStats) -> BestSplit {
let reg = unregularized();
let mut best = BestSplit::none();
let scorer = SplitScorer {
reg: ®,
root_gain: xgb_node_gain(total, ®, Bounds::default()),
bounds: Bounds::default(),
dir: 0,
};
sweep_categorical(&mut best, cats, total, &scorer, 0, None);
best
}
fn categories(best: &BestSplit) -> &[u32] {
match &best.location {
SplitLocation::Categories(categories) => categories,
SplitLocation::Numeric(pos) => panic!("numeric split at {pos:?}"),
}
}
#[test]
fn a_lone_category_splits_from_missing_values() {
let best = sweep(&[(0, GradStats::new(2.0, 2.0))], GradStats::new(0.0, 4.0));
assert_eq!(categories(&best), [0]);
assert!(!best.default_left, "missing values go to the other child");
assert!((best.loss_chg - 4.0).abs() < 1e-6, "{}", best.loss_chg);
}
#[test]
fn missing_values_join_the_side_that_fits_them() {
let cats = [
(0, GradStats::new(2.0, 1.0)),
(1, GradStats::new(1.0, 1.0)),
(2, GradStats::new(-1.0, 1.0)),
(3, GradStats::new(-2.0, 1.0)),
];
let expected = 4.5 + 12.0 - 1.8;
for (missing_grad, with_set) in [(-3.0, false), (3.0, true)] {
let best = sweep(&cats, GradStats::new(missing_grad, 5.0));
assert_eq!(categories(&best), [0, 1], "missing gradient {missing_grad}");
assert_eq!(
best.default_left, with_set,
"missing gradient {missing_grad}"
);
assert!(
(best.loss_chg - expected).abs() < 1e-5,
"missing gradient {missing_grad}: {}",
best.loss_chg
);
}
}
}