1use crate::error::{RillError, ensure_finite};
11
12#[derive(Debug, Clone)]
14#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
15pub struct HuberLoss {
16 delta: f64,
17}
18
19impl HuberLoss {
20 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 pub const fn delta(&self) -> f64 {
36 self.delta
37 }
38
39 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 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 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 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}