1use crate::error::{RillError, ensure_finite};
9
10pub(crate) const DEFAULT_EPSILON: f64 = 1e-15;
12
13#[derive(Debug, Clone)]
15#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
16pub struct BinaryLogLoss {
17 epsilon: f64,
18}
19
20impl BinaryLogLoss {
21 pub const fn new() -> Self {
23 Self {
24 epsilon: DEFAULT_EPSILON,
25 }
26 }
27
28 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 pub const fn epsilon(&self) -> f64 {
42 self.epsilon
43 }
44
45 pub fn clip(&self, p: f64) -> f64 {
47 p.clamp(self.epsilon, 1.0 - self.epsilon)
48 }
49
50 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 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
72pub(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 assert!((l.loss(0.9, true) - (-0.9_f64.ln())).abs() < 1e-12);
103 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}