Skip to main content

rill_ml/loss/
log_loss.rs

1//! Binary log loss (cross-entropy) for logistic regression.
2//!
3//! `loss = -(y * log(p) + (1 - y) * log(1 - p))`
4//!
5//! Probabilities are clipped to `[epsilon, 1 - epsilon]` for numerical
6//! stability. The gradient w.r.t. the raw logit `z` is `(p - y)`.
7
8use crate::error::{RillError, ensure_finite};
9
10/// Default clipping epsilon for probabilities.
11pub(crate) const DEFAULT_EPSILON: f64 = 1e-15;
12
13/// Binary cross-entropy (log) loss.
14#[derive(Debug, Clone)]
15#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
16pub struct BinaryLogLoss {
17    epsilon: f64,
18}
19
20impl BinaryLogLoss {
21    /// Create a new log loss with the default epsilon.
22    pub const fn new() -> Self {
23        Self {
24            epsilon: DEFAULT_EPSILON,
25        }
26    }
27
28    /// Create a new log loss with a custom epsilon.
29    pub fn with_epsilon(epsilon: f64) -> Result<Self, RillError> {
30        ensure_finite("epsilon", epsilon)?;
31        if epsilon <= 0.0 || epsilon >= 0.5 {
32            return Err(RillError::InvalidParameter {
33                name: "epsilon",
34                value: epsilon,
35            });
36        }
37        Ok(Self { epsilon })
38    }
39
40    /// The configured epsilon.
41    pub const fn epsilon(&self) -> f64 {
42        self.epsilon
43    }
44
45    /// Clip a probability to `[epsilon, 1 - epsilon]`.
46    pub fn clip(&self, p: f64) -> f64 {
47        p.clamp(self.epsilon, 1.0 - self.epsilon)
48    }
49
50    /// Compute the log loss given a probability and boolean target.
51    pub fn loss(&self, probability: f64, target: bool) -> f64 {
52        let p = self.clip(probability);
53        let y = if target { 1.0 } else { 0.0 };
54        -(y * p.ln() + (1.0 - y) * (1.0 - p).ln())
55    }
56
57    /// Gradient of the loss w.r.t. the logit `z = log(p / (1 - p))`.
58    ///
59    /// This equals `p - y` where `p = sigmoid(z)`.
60    pub fn gradient_wrt_logit(&self, probability: f64, target: bool) -> f64 {
61        let y = if target { 1.0 } else { 0.0 };
62        probability - y
63    }
64}
65
66impl Default for BinaryLogLoss {
67    fn default() -> Self {
68        Self::new()
69    }
70}
71
72/// Numerically stable sigmoid function.
73///
74/// For `z >= 0`: `1 / (1 + exp(-z))`.
75/// For `z < 0`: `exp(z) / (1 + exp(z))`.
76pub(crate) fn sigmoid(z: f64) -> f64 {
77    if z >= 0.0 {
78        1.0 / (1.0 + (-z).exp())
79    } else {
80        let e = z.exp();
81        e / (1.0 + e)
82    }
83}
84
85#[cfg(test)]
86mod tests {
87    use super::*;
88
89    #[test]
90    fn sigmoid_bounds() {
91        assert!((sigmoid(0.0) - 0.5).abs() < 1e-12);
92        assert!(sigmoid(-100.0) > 0.0 && sigmoid(-100.0) < 1e-10);
93        assert!(sigmoid(100.0) <= 1.0 && sigmoid(100.0) > 1.0 - 1e-10);
94        assert!(!sigmoid(1000.0).is_nan());
95        assert!(!sigmoid(-1000.0).is_nan());
96    }
97
98    #[test]
99    fn log_loss_correct_target() {
100        let l = BinaryLogLoss::new();
101        // p=0.9, target=true -> -log(0.9)
102        assert!((l.loss(0.9, true) - (-0.9_f64.ln())).abs() < 1e-12);
103        // p=0.1, target=false -> -log(0.9)
104        assert!((l.loss(0.1, false) - (-0.9_f64.ln())).abs() < 1e-12);
105    }
106
107    #[test]
108    fn log_loss_clips_extreme_probabilities() {
109        let l = BinaryLogLoss::new();
110        let loss = l.loss(0.0, true);
111        assert!(loss.is_finite() && loss > 0.0);
112        let loss = l.loss(1.0, false);
113        assert!(loss.is_finite() && loss > 0.0);
114    }
115
116    #[test]
117    fn gradient_wrt_logit_matches_probability_minus_target() {
118        let l = BinaryLogLoss::new();
119        assert!((l.gradient_wrt_logit(0.7, true) - (-0.3)).abs() < 1e-12);
120        assert!((l.gradient_wrt_logit(0.7, false) - 0.7).abs() < 1e-12);
121    }
122
123    #[test]
124    fn invalid_epsilon_rejected() {
125        assert!(BinaryLogLoss::with_epsilon(0.0).is_err());
126        assert!(BinaryLogLoss::with_epsilon(0.5).is_err());
127        assert!(BinaryLogLoss::with_epsilon(-0.1).is_err());
128    }
129}