use crate::config::Monotone;
use crate::tree::gain::{GradStats, RegParams, calc_weight, threshold_l1};
#[derive(Debug, Clone, Default)]
pub struct MonotoneConstraints {
dirs: Vec<i8>,
}
impl MonotoneConstraints {
pub fn from_params(list: &[Monotone]) -> Self {
let dirs = list
.iter()
.map(|m| match m {
Monotone::Increasing => 1,
Monotone::Decreasing => -1,
Monotone::None => 0,
})
.collect();
MonotoneConstraints { dirs }
}
pub fn is_active(&self) -> bool {
self.dirs.iter().any(|&d| d != 0)
}
#[inline]
pub fn dir(&self, feature: usize) -> i8 {
self.dirs.get(feature).copied().unwrap_or(0)
}
}
#[derive(Debug, Clone, Copy)]
pub struct Bounds {
pub lower: f64,
pub upper: f64,
}
impl Default for Bounds {
fn default() -> Self {
Bounds {
lower: f64::NEG_INFINITY,
upper: f64::INFINITY,
}
}
}
impl Bounds {
#[inline]
pub fn clamp(&self, w: f64) -> f64 {
w.clamp(self.lower, self.upper)
}
}
pub fn calc_weight_bounded(stats: GradStats, reg: &RegParams, bounds: Bounds) -> f64 {
bounds.clamp(calc_weight(stats, reg))
}
pub fn gain_at_weight(stats: GradStats, reg: &RegParams, w: f64) -> f64 {
if stats.hess < reg.min_child_weight || stats.hess <= 0.0 {
return 0.0;
}
let t = threshold_l1(stats.grad, reg.alpha);
-(2.0 * t * w + (stats.hess + reg.lambda) * w * w)
}
pub fn child_bounds(parent: Bounds, dir: i8, w_left: f64, w_right: f64) -> (Bounds, Bounds) {
if dir == 0 {
return (parent, parent);
}
let mid = f64::midpoint(w_left, w_right);
if dir > 0 {
(
Bounds {
lower: parent.lower,
upper: mid,
},
Bounds {
lower: mid,
upper: parent.upper,
},
)
} else {
(
Bounds {
lower: mid,
upper: parent.upper,
},
Bounds {
lower: parent.lower,
upper: mid,
},
)
}
}
#[inline]
pub fn satisfies(dir: i8, w_left: f64, w_right: f64) -> bool {
match dir {
d if d > 0 => w_left <= w_right,
d if d < 0 => w_left >= w_right,
_ => true,
}
}
#[cfg(test)]
mod tests {
use super::*;
fn reg() -> RegParams {
RegParams {
lambda: 1.0,
alpha: 0.0,
max_delta_step: 0.0,
min_child_weight: 0.0,
}
}
#[test]
fn unbounded_matches_closed_form() {
let s = GradStats::new(-4.0, 2.0);
let r = reg();
let w = calc_weight_bounded(s, &r, Bounds::default());
let g = gain_at_weight(s, &r, w);
assert!((g - 16.0 / 3.0).abs() < 1e-9);
}
#[test]
fn increasing_child_bounds_order() {
let (l, r) = child_bounds(Bounds::default(), 1, -1.0, 2.0);
assert_eq!(l.upper, 0.5); assert_eq!(r.lower, 0.5);
assert!(satisfies(1, -1.0, 2.0));
assert!(!satisfies(1, 2.0, -1.0));
}
#[test]
fn bounds_clamp_weight() {
let b = Bounds {
lower: 0.0,
upper: 1.0,
};
assert_eq!(b.clamp(-5.0), 0.0);
assert_eq!(b.clamp(5.0), 1.0);
assert_eq!(b.clamp(0.5), 0.5);
}
}