use super::{GradPair, Objective, weighted_label_mean};
fn log_link_transform(preds: &mut [f32]) {
crate::simd::exp_inplace(preds);
}
macro_rules! log_link_objective {
() => {
fn pred_transform(&self, preds: &mut [f32]) {
log_link_transform(preds);
}
fn prob_to_margin(&self, base_score: f32) -> f32 {
base_score.ln()
}
fn base_margins(
&self,
labels: &[f32],
weights: Option<&[f32]>,
_group: Option<&crate::data::GroupInfo>,
) -> Vec<f32> {
vec![self.prob_to_margin(weighted_label_mean(labels, weights))]
}
};
}
#[derive(Debug, Clone, Copy)]
pub struct PoissonObjective {
max_delta_step: f32,
}
impl PoissonObjective {
pub fn new(max_delta_step: f32) -> Self {
PoissonObjective { max_delta_step }
}
}
impl Default for PoissonObjective {
fn default() -> Self {
PoissonObjective {
max_delta_step: 0.7,
}
}
}
impl Objective for PoissonObjective {
fn name(&self) -> &'static str {
"count:poisson"
}
fn gradient(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
out: &mut [GradPair],
) {
super::check_gradient_inputs(labels.len(), 1, preds, labels, weights, out);
crate::simd::poisson_gradient(preds, labels, weights, self.max_delta_step, out);
}
log_link_objective!();
fn default_metric(&self) -> String {
"poisson-nloglik".to_string()
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct GammaObjective;
impl Objective for GammaObjective {
fn name(&self) -> &'static str {
"reg:gamma"
}
fn gradient(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
out: &mut [GradPair],
) {
super::check_gradient_inputs(labels.len(), 1, preds, labels, weights, out);
crate::simd::gamma_gradient(preds, labels, weights, out);
}
log_link_objective!();
fn default_metric(&self) -> String {
"gamma-nloglik".to_string()
}
}
#[derive(Debug, Clone, Copy)]
pub struct TweedieObjective {
rho: f32,
}
impl TweedieObjective {
pub fn new(rho: f32) -> Self {
TweedieObjective { rho }
}
}
impl Default for TweedieObjective {
fn default() -> Self {
TweedieObjective { rho: 1.5 }
}
}
impl Objective for TweedieObjective {
fn name(&self) -> &'static str {
"reg:tweedie"
}
fn gradient(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
out: &mut [GradPair],
) {
super::check_gradient_inputs(labels.len(), 1, preds, labels, weights, out);
crate::simd::tweedie_gradient(preds, labels, weights, self.rho, out);
}
log_link_objective!();
fn default_metric(&self) -> String {
format!("tweedie-nloglik@{}", self.rho)
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_relative_eq;
#[test]
fn poisson_gradient_at_log_mean_is_zero_sum() {
let obj = PoissonObjective::default();
let labels = [2.0f32, 5.0];
let preds = [2.0f32.ln(), 5.0f32.ln()];
let mut out = vec![GradPair::default(); 2];
obj.gradient(&preds, &labels, None, &mut out);
assert_relative_eq!(out[0].grad, 0.0, epsilon = 1e-5);
assert_relative_eq!(out[1].grad, 0.0, epsilon = 1e-5);
assert!(out[0].hess > 0.0);
}
#[test]
fn gamma_gradient_zero_at_log_y() {
let obj = GammaObjective;
let labels = [3.0f32];
let preds = [3.0f32.ln()];
let mut out = vec![GradPair::default(); 1];
obj.gradient(&preds, &labels, None, &mut out);
assert_relative_eq!(out[0].grad, 0.0, epsilon = 1e-5);
}
#[test]
fn tweedie_transform_is_exp() {
let obj = TweedieObjective::default();
let mut p = [0.0f32, 1.0];
obj.pred_transform(&mut p);
assert_relative_eq!(p[0], 1.0, epsilon = 1e-6);
assert_relative_eq!(p[1], 1.0f32.exp(), epsilon = 1e-6);
}
#[test]
fn log_link_intercept_is_ln_of_mean() {
let obj = PoissonObjective::default();
assert_eq!(obj.base_margins(&[2.0, 6.0], None, None), vec![4f32.ln()]);
let w = [3.0f32, 1.0];
assert_eq!(
obj.base_margins(&[2.0, 6.0], Some(&w), None),
vec![3f32.ln()]
);
assert_eq!(
GammaObjective.base_margins(&[0.0, 0.0], None, None),
vec![f32::NEG_INFINITY]
);
}
}