use super::{GradPair, Loss, 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;
use crate::objective::PseudoHuber;
#[derive(Debug, Clone, Copy)]
pub(crate) struct SquaredError {
scale_pos_weight: f32,
}
impl SquaredError {
pub(crate) fn new(scale_pos_weight: f32) -> Self {
SquaredError { scale_pos_weight }
}
}
impl Default for SquaredError {
fn default() -> Self {
SquaredError::new(1.0)
}
}
impl Loss for SquaredError {
fn name(&self) -> &'static str {
"reg:squarederror"
}
fn gradient(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
out: &mut [GradPair],
) {
let scale_pos_weight = self.scale_pos_weight;
super::elementwise_gradient(preds, labels, weights, out, |p, y, mut w| {
if y == 1.0 {
w *= scale_pos_weight;
}
GradPair::new((p - y) * w, w)
});
}
fn const_hess(&self) -> bool {
true
}
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![weighted_label_mean(info.label_values(), info.weights)]
}
fn pointwise_loss(&self) -> Option<super::PointwiseLoss<'_>> {
let scale_pos_weight = f64::from(self.scale_pos_weight);
Some(Box::new(move |margin, label| {
let weight = if label == 1.0 { scale_pos_weight } else { 1.0 };
weight * 0.5 * (f64::from(margin) - f64::from(label)).powi(2)
}))
}
fn default_metric(&self) -> EvalMetric {
EvalMetric::Rmse
}
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct PseudoHuberLoss {
param: PseudoHuber,
slope: f32,
}
impl PseudoHuberLoss {
pub(crate) fn new(param: PseudoHuber) -> Self {
PseudoHuberLoss {
param,
slope: param.slope() as f32,
}
}
}
impl Default for PseudoHuberLoss {
fn default() -> Self {
PseudoHuberLoss::new(PseudoHuber::default())
}
}
impl Loss for PseudoHuberLoss {
fn name(&self) -> &'static str {
"reg:pseudohubererror"
}
fn gradient(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
out: &mut [GradPair],
) {
let slope_sq = self.slope * self.slope;
super::elementwise_gradient(preds, labels, weights, out, |p, y, w| {
let z = p - y;
let scale_sqrt = (1.0 + z * z / slope_sq).sqrt();
let scale = slope_sq + z * z;
GradPair::new((z / scale_sqrt) * w, (slope_sq / (scale * scale_sqrt)) * w)
});
}
fn pointwise_loss(&self) -> Option<super::PointwiseLoss<'_>> {
let slope_sq = f64::from(self.slope).powi(2);
Some(Box::new(move |margin, label| {
let z = f64::from(margin) - f64::from(label);
z * z / ((1.0 + z * z / slope_sq).sqrt() + 1.0)
}))
}
fn default_metric(&self) -> EvalMetric {
EvalMetric::Mphe(self.param)
}
}
#[derive(Debug, Clone, Copy, Default)]
#[non_exhaustive]
pub struct SquaredLogError;
const SQUARED_LOG_MIN_PRED: f32 = (-1.0f64 + 1e-6) as f32;
impl Loss for SquaredLogError {
fn name(&self) -> &'static str {
"reg:squaredlogerror"
}
fn gradient(
&self,
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
out: &mut [GradPair],
) {
super::elementwise_gradient(preds, labels, weights, out, |p, y, w| {
let p = p.max(SQUARED_LOG_MIN_PRED);
let (log_p, log_y) = (p.ln_1p(), y.ln_1p());
let grad = (log_p - log_y) / (p + 1.0);
let shifted = f64::from(p + 1.0);
let hess = ((f64::from(-log_p + log_y + 1.0) / (shifted * shifted)) as f32).max(1e-6);
GradPair::new(grad * w, hess * w)
});
}
fn validate_info(&self, info: &MetaInfo) -> Result<()> {
check_label_domain(info, |y| y <= -1.0)
}
fn default_metric(&self) -> EvalMetric {
EvalMetric::Rmsle
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::objective::{base_margins, gradient_pairs};
#[test]
fn gradient_matches_closed_form() {
let obj = SquaredError::default();
let preds = [2.0f32, 0.0, -1.0];
let labels = [1.0f32, 0.5, -3.0];
let out = gradient_pairs(&obj, &preds, &labels, None);
assert_eq!(out[0], GradPair::new(1.0, 1.0)); assert_eq!(out[1], GradPair::new(-0.5, 1.0)); assert_eq!(out[2], GradPair::new(2.0, 1.0)); }
#[test]
fn weighted_gradient_scales() {
let obj = SquaredError::default();
let preds = [2.0f32];
let labels = [1.0f32];
let w = [4.0f32];
let out = gradient_pairs(&obj, &preds, &labels, Some(&w));
assert_eq!(out[0], GradPair::new(4.0, 4.0));
}
#[test]
fn scale_pos_weight_reweights_rows_labeled_one() {
let obj = SquaredError::new(3.0);
let out = gradient_pairs(&obj, &[2.0, 2.0], &[1.0, 1.5], Some(&[2.0, 2.0]));
assert_eq!(out[0], GradPair::new(6.0, 6.0)); assert_eq!(out[1], GradPair::new(1.0, 2.0)); assert_eq!(base_margins(&obj, &[1.0, 4.0], None), vec![1.75]);
}
#[test]
fn base_margins_is_label_mean() {
let obj = SquaredError::default();
assert_eq!(base_margins(&obj, &[1.0, 2.0, 3.0], None), vec![2.0]);
}
#[test]
fn pseudo_huber_slope_scales_gradient() {
let obj = PseudoHuberLoss::new(PseudoHuber::new(2.0).unwrap());
let out = gradient_pairs(&obj, &[2.0], &[0.0], None);
let root2 = 2f32.sqrt();
assert!(
(out[0].grad - 2.0 / root2).abs() < 1e-6,
"grad {}",
out[0].grad
);
assert!(
(out[0].hess - 1.0 / (2.0 * root2)).abs() < 1e-6,
"hess {}",
out[0].hess
);
let out = gradient_pairs(&PseudoHuberLoss::default(), &[2.0], &[0.0], None);
let s = 5f32;
assert_eq!(out[0], GradPair::new(2.0 / s.sqrt(), 1.0 / (s * s.sqrt())));
}
#[test]
fn pseudo_huber_intercept_is_newton_step() {
let obj = PseudoHuberLoss::default();
let labels = [0.0f32, 4.0];
let margins = base_margins(&obj, &labels, None);
let s = 17f32; let g1 = -4.0f32 / s.sqrt();
let h1 = 1.0f32 / (s * s.sqrt());
let expected = (-f64::from(g1) / (1.0 + f64::from(h1))) as f32;
assert_eq!(margins, vec![expected]);
assert!(
margins[0] < 1.0,
"Newton step {} should undershoot the mean",
margins[0]
);
}
#[test]
fn pseudo_huber_loss_survives_large_slopes() {
let obj = PseudoHuberLoss::new(PseudoHuber::new(1e9).unwrap());
let loss = obj.pointwise_loss().unwrap();
assert!((loss(0.0, 1.0) - 0.5).abs() < 1e-12, "{}", loss(0.0, 1.0));
let obj = PseudoHuberLoss::new(PseudoHuber::new(2.0).unwrap());
let loss = obj.pointwise_loss().unwrap();
let naive = 4.0 * ((1.0f64 + 9.0 / 4.0).sqrt() - 1.0);
assert!((loss(3.0, 0.0) - naive).abs() < 1e-12);
}
#[test]
fn squared_log_gradient_zero_at_label_and_finite_below_minus_one() {
let obj = SquaredLogError;
let out = gradient_pairs(&obj, &[3.0, -1.0, -7.0], &[3.0, 0.5, 0.5], None);
assert_eq!(out[0].grad, 0.0);
assert_eq!(out[0].hess, 1.0 / 16.0);
assert!(out[1].grad.is_finite() && out[1].grad < 0.0);
assert_eq!(out[1], out[2], "margins below the clamp share its gradient");
}
#[test]
fn squared_log_hessian_floor_is_weighted() {
let obj = SquaredLogError;
let out = gradient_pairs(&obj, &[1e4], &[0.0], Some(&[2.0]));
assert_eq!(out[0].hess, 2e-6);
assert!(out[0].grad > 0.0);
}
#[test]
fn squared_log_rejects_labels_at_minus_one() {
let obj = SquaredLogError;
assert!(
obj.validate_info(&MetaInfo::new(&[-0.5, 2.0], None, None))
.is_ok()
);
assert!(
obj.validate_info(&MetaInfo::new(&[-1.0], None, None))
.is_err()
);
}
}