Skip to main content

rill_ml/loss/
huber.rs

1//! Huber loss, robust to outliers.
2//!
3//! For `|residual| <= delta`: `0.5 * residual^2`.
4//! For `|residual| > delta`: `delta * (|residual| - 0.5 * delta)`.
5//!
6//! The gradient w.r.t. the prediction is:
7//! - `residual` if `|residual| <= delta`
8//! - `delta * sign(residual)` otherwise
9
10use crate::error::{RillError, ensure_finite};
11
12/// Huber loss with configurable delta.
13#[derive(Debug, Clone)]
14#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
15pub struct HuberLoss {
16    delta: f64,
17}
18
19impl HuberLoss {
20    /// Create a new Huber loss with the given delta.
21    ///
22    /// Returns an error if `delta` is not finite and strictly positive.
23    pub fn new(delta: f64) -> Result<Self, RillError> {
24        ensure_finite("delta", delta)?;
25        if delta <= 0.0 {
26            return Err(RillError::InvalidParameter {
27                name: "delta",
28                value: delta,
29            });
30        }
31        Ok(Self { delta })
32    }
33
34    /// The configured delta.
35    pub const fn delta(&self) -> f64 {
36        self.delta
37    }
38
39    /// Compute the Huber loss.
40    pub fn loss(&self, prediction: f64, target: f64) -> f64 {
41        let residual = prediction - target;
42        let abs_r = residual.abs();
43        if abs_r <= self.delta {
44            0.5 * residual * residual
45        } else {
46            self.delta * (abs_r - 0.5 * self.delta)
47        }
48    }
49
50    /// Compute the gradient w.r.t. the prediction.
51    pub fn gradient(&self, prediction: f64, target: f64) -> f64 {
52        let residual = prediction - target;
53        let abs_r = residual.abs();
54        if abs_r <= self.delta {
55            residual
56        } else {
57            self.delta * residual.signum()
58        }
59    }
60}
61
62impl Default for HuberLoss {
63    fn default() -> Self {
64        Self { delta: 1.0 }
65    }
66}
67
68#[cfg(test)]
69mod tests {
70    use super::*;
71
72    #[test]
73    fn quadratic_region_matches_squared() {
74        let h = HuberLoss::new(1.0).unwrap();
75        // residual = 0.5, within delta
76        assert!((h.loss(1.5, 1.0) - 0.5 * 0.25).abs() < 1e-12);
77        assert!((h.gradient(1.5, 1.0) - 0.5).abs() < 1e-12);
78    }
79
80    #[test]
81    fn linear_region_clipped() {
82        let h = HuberLoss::new(1.0).unwrap();
83        // residual = 3, outside delta=1
84        let expected_loss = 1.0 * (3.0 - 0.5);
85        assert!((h.loss(4.0, 1.0) - expected_loss).abs() < 1e-12);
86        assert!((h.gradient(4.0, 1.0) - 1.0).abs() < 1e-12);
87    }
88
89    #[test]
90    fn invalid_delta_rejected() {
91        assert!(HuberLoss::new(0.0).is_err());
92        assert!(HuberLoss::new(-1.0).is_err());
93        assert!(HuberLoss::new(f64::NAN).is_err());
94    }
95}