1use crate::error::{
6 RillError, checked_finite_add, checked_increment, ensure_finite, validate_features,
7};
8use crate::loss::log_loss::{BinaryLogLoss, sigmoid};
9use crate::optim::Optimizer;
10#[cfg(feature = "serde")]
11use crate::persistence::ValidateState;
12use crate::traits::OnlineBinaryClassifier;
13
14#[derive(Debug, Clone)]
16#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
17#[non_exhaustive]
18pub struct LogisticRegressionConfig {
19 pub optimizer: Optimizer,
21 pub loss: BinaryLogLoss,
23}
24
25impl Default for LogisticRegressionConfig {
26 fn default() -> Self {
27 Self {
28 optimizer: Optimizer::sgd(1, Default::default()).expect("default optimizer"),
29 loss: BinaryLogLoss::new(),
30 }
31 }
32}
33
34#[derive(Debug, Clone)]
39#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
40pub struct LogisticRegression {
41 feature_count: usize,
42 weights: Vec<f64>,
43 intercept: f64,
44 optimizer: Optimizer,
45 loss: BinaryLogLoss,
46 samples_seen: u64,
47}
48
49impl LogisticRegression {
50 pub fn new(feature_count: usize, config: LogisticRegressionConfig) -> Result<Self, RillError> {
52 if feature_count == 0 {
53 return Err(RillError::EmptyFeatures);
54 }
55 if config.optimizer.param_count() != feature_count + 1 {
56 return Err(RillError::DimensionMismatch {
57 expected: feature_count + 1,
58 actual: config.optimizer.param_count(),
59 });
60 }
61 Ok(Self {
62 feature_count,
63 weights: vec![0.0; feature_count],
64 intercept: 0.0,
65 optimizer: config.optimizer,
66 loss: config.loss,
67 samples_seen: 0,
68 })
69 }
70
71 pub fn weights(&self) -> &[f64] {
73 &self.weights
74 }
75
76 pub const fn intercept(&self) -> f64 {
78 self.intercept
79 }
80
81 fn logit(&self, features: &[f64]) -> Result<f64, RillError> {
83 validate_features(self.feature_count, features)?;
84 let dot = self.weights.iter().zip(features.iter()).try_fold(
85 0.0,
86 |sum, (&weight, &feature)| {
87 let term = weight * feature;
88 ensure_finite("logit term", term)?;
89 checked_finite_add(sum, term, "logit")
90 },
91 )?;
92 checked_finite_add(dot, self.intercept, "logit")
93 }
94}
95
96impl OnlineBinaryClassifier for LogisticRegression {
97 fn feature_count(&self) -> usize {
98 self.feature_count
99 }
100
101 fn samples_seen(&self) -> u64 {
102 self.samples_seen
103 }
104
105 fn predict_proba(&self, features: &[f64]) -> Result<f64, RillError> {
106 let z = self.logit(features)?;
107 Ok(sigmoid(z))
108 }
109
110 fn learn(&mut self, features: &[f64], target: bool) -> Result<(), RillError> {
111 validate_features(self.feature_count, features)?;
112 let next_samples = checked_increment(self.samples_seen, "logistic regression sample")?;
113 let z = self.logit(features)?;
114 let p = sigmoid(z);
115 let grad = self.loss.gradient_wrt_logit(p, target);
117 ensure_finite("loss gradient", grad)?;
118 let grad_weights = features
119 .iter()
120 .map(|&feature| {
121 let gradient = grad * feature;
122 ensure_finite("weight gradient", gradient)?;
123 Ok(gradient)
124 })
125 .collect::<Result<Vec<_>, RillError>>()?;
126 let grad_intercept = grad;
127 self.optimizer.step(
128 &mut self.weights,
129 &mut self.intercept,
130 &grad_weights,
131 grad_intercept,
132 )?;
133 self.samples_seen = next_samples;
134 Ok(())
135 }
136
137 fn reset(&mut self) {
138 for w in &mut self.weights {
139 *w = 0.0;
140 }
141 self.intercept = 0.0;
142 self.optimizer.reset();
143 self.samples_seen = 0;
144 }
145}
146
147#[cfg(feature = "serde")]
148impl ValidateState for LogisticRegression {
149 fn validate_state(&self) -> Result<(), RillError> {
150 if self.feature_count == 0 {
151 return Err(RillError::EmptyFeatures);
152 }
153 if self.weights.len() != self.feature_count {
154 return Err(RillError::InvalidState(format!(
155 "logistic regression weights length {} does not match feature_count {}",
156 self.weights.len(),
157 self.feature_count
158 )));
159 }
160 if self.optimizer.param_count() != self.feature_count + 1 {
161 return Err(RillError::InvalidState(format!(
162 "logistic regression optimizer param_count {} does not match feature_count+1 {}",
163 self.optimizer.param_count(),
164 self.feature_count + 1
165 )));
166 }
167 ensure_finite("intercept", self.intercept)?;
168 for &w in &self.weights {
169 ensure_finite("weights", w)?;
170 }
171 self.optimizer.validate_state()?;
172 Ok(())
173 }
174}
175
176#[cfg(test)]
177mod tests {
178 use super::*;
179 use crate::optim::SgdConfig;
180 use rand::SeedableRng;
181
182 fn make_model(d: usize, lr: f64) -> LogisticRegression {
183 LogisticRegression::new(
184 d,
185 LogisticRegressionConfig {
186 optimizer: Optimizer::sgd(
187 d,
188 SgdConfig {
189 learning_rate: lr,
190 l2: 0.0,
191 },
192 )
193 .unwrap(),
194 loss: BinaryLogLoss::new(),
195 },
196 )
197 .unwrap()
198 }
199
200 #[test]
201 fn predict_proba_in_range() {
202 let model = make_model(2, 0.1);
203 let p = model.predict_proba(&[1.0, 2.0]).unwrap();
204 assert!(p > 0.0 && p < 1.0);
205 assert!((p - 0.5).abs() < 1e-12);
207 }
208
209 #[test]
210 fn learn_separable_data() {
211 let mut model = make_model(2, 0.5);
212 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(3);
213 for _ in 0..1000 {
214 let x1 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
216 let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
217 let y = x1 > 0.0;
218 model.learn(&[x1, x2], y).unwrap();
219 }
220 let p_pos = model.predict_proba(&[2.0, 0.0]).unwrap();
221 let p_neg = model.predict_proba(&[-2.0, 0.0]).unwrap();
222 assert!(p_pos > 0.7, "p_pos = {p_pos}");
223 assert!(p_neg < 0.3, "p_neg = {p_neg}");
224 }
225
226 #[test]
227 fn predict_does_not_update_state() {
228 let model = make_model(1, 0.1);
229 let _ = model.predict_proba(&[1.0]).unwrap();
230 assert_eq!(model.samples_seen(), 0);
231 }
232
233 #[test]
234 fn dimension_mismatch_rejected() {
235 let mut model = make_model(3, 0.1);
236 assert!(model.predict_proba(&[1.0, 2.0]).is_err());
237 assert!(model.learn(&[1.0, 2.0], true).is_err());
238 }
239
240 #[test]
241 fn reset_clears_state() {
242 let mut model = make_model(1, 0.1);
243 model.learn(&[1.0], true).unwrap();
244 model.reset();
245 assert_eq!(model.samples_seen(), 0);
246 assert!((model.predict_proba(&[1.0]).unwrap() - 0.5).abs() < 1e-12);
247 }
248}