use super::{
GradPair, Loss, MIN_HESS, OutputDomain, check_base_score_domain, check_label_domain,
newton_intercepts, weighted_label_mean,
};
use crate::K_RT_EPS_F32;
use crate::data::MetaInfo;
use crate::error::Result;
use crate::metric::EvalMetric;
#[derive(Debug, Clone, Copy)]
pub struct LogisticLoss {
scale_pos_weight: f32,
variant: LogisticVariant,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum LogisticVariant {
Binary,
Regression,
Raw,
}
impl LogisticLoss {
pub fn new(scale_pos_weight: f32) -> Self {
LogisticLoss {
scale_pos_weight,
variant: LogisticVariant::Binary,
}
}
pub fn regression(scale_pos_weight: f32) -> Self {
LogisticLoss {
scale_pos_weight,
variant: LogisticVariant::Regression,
}
}
pub fn raw(scale_pos_weight: f32) -> Self {
LogisticLoss {
scale_pos_weight,
variant: LogisticVariant::Raw,
}
}
fn link(self, p: f32) -> f32 {
if self.variant == LogisticVariant::Raw {
return p;
}
let p = p.clamp(K_RT_EPS_F32, 1.0 - K_RT_EPS_F32);
-(1.0 / p - 1.0).ln()
}
}
impl Default for LogisticLoss {
fn default() -> Self {
Self::new(1.0)
}
}
impl Loss for LogisticLoss {
fn name(&self) -> &str {
match self.variant {
LogisticVariant::Binary => "binary:logistic",
LogisticVariant::Regression => "reg:logistic",
LogisticVariant::Raw => "binary:logitraw",
}
}
fn gradient(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
out: &mut [GradPair],
) {
let scale_pos_weight = self.scale_pos_weight;
super::rowwise_gradient(
labels.len(),
1,
preds,
labels,
weights,
out,
|preds, labels, weights, out| {
crate::simd::logistic_gradient(
preds,
labels,
weights,
scale_pos_weight,
MIN_HESS,
out,
);
},
);
}
fn pred_transform(&self, preds: &mut [f32]) {
if self.variant != LogisticVariant::Raw {
crate::simd::sigmoid_inplace(preds);
}
}
fn base_margins_info(&self, info: &MetaInfo) -> Vec<f32> {
if (self.scale_pos_weight - 1.0).abs() > K_RT_EPS_F32 {
return newton_intercepts(self, info);
}
vec![self.link(weighted_label_mean(info.label_values(), info.weights))]
}
fn probs_to_margins(&self, scores: &mut [f32]) {
for s in scores {
*s = self.link(*s);
}
}
fn validate_base_score(&self, base_score: f64) -> Result<()> {
if self.variant == LogisticVariant::Raw {
return Ok(());
}
check_base_score_domain(base_score, OutputDomain::Probability)
}
fn pointwise_loss(&self) -> Option<super::PointwiseLoss<'_>> {
let scale_pos_weight = f64::from(self.scale_pos_weight);
Some(Box::new(move |margin, label| {
let (m, y) = (f64::from(margin), f64::from(label));
let softplus = m.max(0.0) + (-m.abs()).exp().ln_1p();
let weight = if label == 1.0 { scale_pos_weight } else { 1.0 };
weight * (softplus - y * m)
}))
}
fn validate_info(&self, info: &MetaInfo) -> Result<()> {
check_label_domain(info, |y| !(0.0..=1.0).contains(&y))
}
fn default_metric(&self) -> EvalMetric {
match self.variant {
LogisticVariant::Regression => EvalMetric::Rmse,
LogisticVariant::Binary | LogisticVariant::Raw => EvalMetric::LogLoss,
}
}
}
#[derive(Debug, Clone, Copy, Default)]
#[non_exhaustive]
pub struct Hinge;
impl Loss for Hinge {
fn name(&self) -> &'static str {
"binary:hinge"
}
fn gradient(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
out: &mut [GradPair],
) {
super::elementwise_gradient(preds, labels, weights, out, |p, y, w| {
let z = f64::from(y) * 2.0 - 1.0;
if f64::from(p) * z < 1.0 {
GradPair::new((-z * f64::from(w)) as f32, w)
} else {
GradPair::new(0.0, f32::MIN_POSITIVE)
}
});
}
fn pred_transform(&self, preds: &mut [f32]) {
for p in preds {
*p = if *p > 0.0 { 1.0 } else { 0.0 };
}
}
fn margins_to_probs(&self, _margins: &mut [f32]) {
}
fn default_metric(&self) -> EvalMetric {
EvalMetric::Error
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::objective::{base_margins, gradient_pairs};
use approx::assert_relative_eq;
#[test]
fn sigmoid_symmetry() {
let mut values = [0.0, 2.0, -2.0, 80.0, -80.0];
LogisticLoss::default().pred_transform(&mut values);
assert_relative_eq!(values[0], 0.5, epsilon = 1e-6);
assert_relative_eq!(values[1] + values[2], 1.0, epsilon = 1e-6);
assert!(values[3] <= 1.0 && values[3] > 0.999);
assert!(values[4] >= 0.0 && values[4] < 0.001);
}
#[test]
fn gradient_matches_closed_form() {
let obj = LogisticLoss::default();
let preds = [0.0f32];
let labels = [1.0f32];
let out = gradient_pairs(&obj, &preds, &labels, None);
assert_relative_eq!(out[0].grad, -0.5, epsilon = 1e-6); assert_relative_eq!(out[0].hess, 0.25, epsilon = 1e-6); }
#[test]
fn scale_pos_weight_scales_positive() {
let obj = LogisticLoss::new(3.0);
let preds = [0.0f32, 0.0];
let labels = [1.0f32, 0.0];
let out = gradient_pairs(&obj, &preds, &labels, None);
assert_relative_eq!(out[0].grad, -1.5, epsilon = 1e-6);
assert_relative_eq!(out[0].hess, 0.75, epsilon = 1e-6);
assert_relative_eq!(out[1].grad, 0.5, epsilon = 1e-6);
assert_relative_eq!(out[1].hess, 0.25, epsilon = 1e-6);
}
#[test]
fn base_margins_is_logit_of_rate() {
let obj = LogisticLoss::default();
assert_eq!(base_margins(&obj, &[1.0, 0.0], None), vec![0.0]);
let quarter = base_margins(&obj, &[1.0, 0.0, 0.0, 0.0], None);
assert_eq!(quarter, vec![-(1.0f32 / 0.25 - 1.0).ln()]);
}
#[test]
fn probs_to_margins_clamps_to_xgboost_bounds() {
let obj = LogisticLoss::default();
let mut scores = [0.0, 1e-6, 1.0, 1.0 - 1e-6];
obj.probs_to_margins(&mut scores);
assert_eq!(scores[0], scores[1]);
assert_eq!(scores[2], scores[3]);
assert!(scores[0].is_finite());
}
#[test]
fn logitraw_keeps_margins_and_uses_mean_intercept() {
let raw = LogisticLoss::raw(1.0);
let mut values = [-3.0f32, 0.5];
raw.pred_transform(&mut values);
assert_eq!(values, [-3.0, 0.5]);
assert_eq!(base_margins(&raw, &[1.0, 0.0, 0.0, 0.0], None), vec![0.25]);
let (preds, labels) = ([0.3, -1.2], [1.0, 0.0]);
assert_eq!(
gradient_pairs(&raw, &preds, &labels, None),
gradient_pairs(&LogisticLoss::new(1.0), &preds, &labels, None)
);
}
#[test]
fn hinge_gradient_and_threshold() {
let obj = Hinge;
let out = gradient_pairs(
&obj,
&[0.5, 1.0, -0.5, -2.0],
&[1.0, 1.0, 0.0, 0.0],
Some(&[2.0, 2.0, 3.0, 3.0]),
);
assert_eq!(out[0], GradPair::new(-2.0, 2.0));
assert_eq!(out[1], GradPair::new(0.0, f32::MIN_POSITIVE));
assert_eq!(out[2], GradPair::new(3.0, 3.0));
assert_eq!(out[3], GradPair::new(0.0, f32::MIN_POSITIVE));
let mut p = [0.0f32, 1e-7, -1.0];
obj.pred_transform(&mut p);
assert_eq!(p, [0.0, 1.0, 0.0]);
}
#[test]
fn hinge_intercept_is_thresholded_newton_step() {
let obj = Hinge;
assert_eq!(base_margins(&obj, &[1.0, 1.0, 1.0, 0.0], None), vec![1.0]);
assert_eq!(base_margins(&obj, &[0.0, 0.0, 1.0], None), vec![0.0]);
let mut stored = [0.5f32];
obj.margins_to_probs(&mut stored);
assert_eq!(stored, [0.5]);
}
}