use crate::config::TrainingParams;
use crate::objective::GradPair;
use crate::tree::constraints::gain_at_weight;
#[derive(Debug, Clone, Copy, PartialEq, Default)]
#[repr(C)]
pub struct GradStats {
pub grad: f64,
pub hess: f64,
}
impl GradStats {
#[inline]
pub fn new(grad: f64, hess: f64) -> Self {
GradStats { grad, hess }
}
#[inline]
pub fn from_pair(gp: GradPair) -> Self {
GradStats::new(f64::from(gp.grad), f64::from(gp.hess))
}
#[inline]
pub fn add(&mut self, other: GradStats) {
self.grad += other.grad;
self.hess += other.hess;
}
#[inline]
#[must_use]
pub fn sub(&self, other: GradStats) -> GradStats {
GradStats {
grad: self.grad - other.grad,
hess: self.hess - other.hess,
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct RegParams {
pub lambda: f64,
pub alpha: f64,
pub max_delta_step: f64,
pub min_child_weight: f64,
}
impl RegParams {
pub fn from_params(p: &TrainingParams) -> Self {
RegParams {
lambda: f64::from(p.lambda as f32),
alpha: f64::from(p.alpha as f32),
max_delta_step: f64::from(p.effective_max_delta_step() as f32),
min_child_weight: f64::from(p.min_child_weight as f32),
}
}
}
#[inline]
pub fn threshold_l1(g: f64, alpha: f64) -> f64 {
if g > alpha {
g - alpha
} else if g < -alpha {
g + alpha
} else {
0.0
}
}
pub fn calc_weight(stats: GradStats, reg: &RegParams) -> f64 {
if stats.hess < reg.min_child_weight || 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 = w.clamp(-reg.max_delta_step, reg.max_delta_step);
}
w
}
pub fn calc_gain(stats: GradStats, reg: &RegParams) -> f64 {
if stats.hess < reg.min_child_weight || stats.hess <= 0.0 {
return 0.0;
}
if reg.max_delta_step == 0.0 {
let t = threshold_l1(stats.grad, reg.alpha);
(t * t) / (stats.hess + reg.lambda)
} else {
gain_at_weight(stats, reg, calc_weight(stats, reg))
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_relative_eq;
fn reg(lambda: f64, alpha: f64) -> RegParams {
RegParams {
lambda,
alpha,
max_delta_step: 0.0,
min_child_weight: 0.0,
}
}
#[test]
fn threshold_l1_shrinks_toward_zero() {
assert_eq!(threshold_l1(5.0, 2.0), 3.0);
assert_eq!(threshold_l1(-5.0, 2.0), -3.0);
assert_eq!(threshold_l1(1.0, 2.0), 0.0);
}
#[test]
fn weight_and_gain_closed_form() {
let s = GradStats::new(-4.0, 2.0);
let r = reg(1.0, 0.0);
assert_relative_eq!(calc_weight(s, &r), 4.0 / 3.0, epsilon = 1e-12);
assert_relative_eq!(calc_gain(s, &r), 16.0 / 3.0, epsilon = 1e-12);
}
#[test]
fn l1_reduces_weight_and_gain() {
let s = GradStats::new(-4.0, 2.0);
let plain = calc_gain(s, ®(1.0, 0.0));
let l1 = calc_gain(s, ®(1.0, 1.0));
assert!(l1 < plain);
}
#[test]
fn max_delta_step_clamps_weight() {
let s = GradStats::new(-100.0, 1.0);
let r = RegParams {
lambda: 0.0,
alpha: 0.0,
max_delta_step: 1.0,
min_child_weight: 0.0,
};
assert_relative_eq!(calc_weight(s, &r), 1.0, epsilon = 1e-12);
}
#[test]
fn min_child_weight_zeros_out() {
let s = GradStats::new(-4.0, 0.5);
let r = RegParams {
lambda: 1.0,
alpha: 0.0,
max_delta_step: 0.0,
min_child_weight: 1.0,
};
assert_eq!(calc_weight(s, &r), 0.0);
assert_eq!(calc_gain(s, &r), 0.0);
}
}