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 pub fn learn_weighted(
99 &mut self,
100 features: &[f64],
101 target: bool,
102 weight: f64,
103 ) -> Result<(), RillError> {
104 crate::weighted::validate_weight(weight)?;
105 validate_features(self.feature_count, features)?;
106 if weight == 0.0 {
107 return Ok(());
108 }
109 let next_samples = checked_increment(self.samples_seen, "logistic regression sample")?;
110 let p = sigmoid(self.logit(features)?);
111 let grad = self.loss.gradient_wrt_logit(p, target) * weight;
112 ensure_finite("weighted loss gradient", grad)?;
113 let grad_weights = features
114 .iter()
115 .map(|&feature| {
116 let gradient = grad * feature;
117 ensure_finite("weighted weight gradient", gradient)?;
118 Ok(gradient)
119 })
120 .collect::<Result<Vec<_>, RillError>>()?;
121 self.optimizer
122 .step(&mut self.weights, &mut self.intercept, &grad_weights, grad)?;
123 self.samples_seen = next_samples;
124 Ok(())
125 }
126}
127
128impl crate::weighted::WeightedOnlineBinaryClassifier for LogisticRegression {
129 fn learn_weighted(
130 &mut self,
131 features: &[f64],
132 target: bool,
133 weight: f64,
134 ) -> Result<(), RillError> {
135 LogisticRegression::learn_weighted(self, features, target, weight)
136 }
137}
138
139impl OnlineBinaryClassifier for LogisticRegression {
140 fn feature_count(&self) -> usize {
141 self.feature_count
142 }
143
144 fn samples_seen(&self) -> u64 {
145 self.samples_seen
146 }
147
148 fn predict_proba(&self, features: &[f64]) -> Result<f64, RillError> {
149 let z = self.logit(features)?;
150 Ok(sigmoid(z))
151 }
152
153 fn learn(&mut self, features: &[f64], target: bool) -> Result<(), RillError> {
154 validate_features(self.feature_count, features)?;
155 let next_samples = checked_increment(self.samples_seen, "logistic regression sample")?;
156 let z = self.logit(features)?;
157 let p = sigmoid(z);
158 let grad = self.loss.gradient_wrt_logit(p, target);
160 ensure_finite("loss gradient", grad)?;
161 let grad_weights = features
162 .iter()
163 .map(|&feature| {
164 let gradient = grad * feature;
165 ensure_finite("weight gradient", gradient)?;
166 Ok(gradient)
167 })
168 .collect::<Result<Vec<_>, RillError>>()?;
169 let grad_intercept = grad;
170 self.optimizer.step(
171 &mut self.weights,
172 &mut self.intercept,
173 &grad_weights,
174 grad_intercept,
175 )?;
176 self.samples_seen = next_samples;
177 Ok(())
178 }
179
180 fn reset(&mut self) {
181 self.weights.fill(0.0);
182 self.intercept = 0.0;
183 self.optimizer.reset();
184 self.samples_seen = 0;
185 }
186}
187
188#[cfg(feature = "serde")]
189impl ValidateState for LogisticRegression {
190 fn validate_state(&self) -> Result<(), RillError> {
191 if self.feature_count == 0 {
192 return Err(RillError::EmptyFeatures);
193 }
194 if self.weights.len() != self.feature_count {
195 return Err(RillError::InvalidState(format!(
196 "logistic regression weights length {} does not match feature_count {}",
197 self.weights.len(),
198 self.feature_count
199 )));
200 }
201 if self.optimizer.param_count() != self.feature_count + 1 {
202 return Err(RillError::InvalidState(format!(
203 "logistic regression optimizer param_count {} does not match feature_count+1 {}",
204 self.optimizer.param_count(),
205 self.feature_count + 1
206 )));
207 }
208 ensure_finite("intercept", self.intercept)?;
209 for &w in &self.weights {
210 ensure_finite("weights", w)?;
211 }
212 self.optimizer.validate_state()?;
213 Ok(())
214 }
215}
216
217#[cfg(test)]
218mod tests {
219 use super::*;
220 use crate::optim::SgdConfig;
221 use rand::SeedableRng;
222
223 fn make_model(d: usize, lr: f64) -> LogisticRegression {
224 LogisticRegression::new(
225 d,
226 LogisticRegressionConfig {
227 optimizer: Optimizer::sgd(
228 d,
229 SgdConfig {
230 learning_rate: lr,
231 l2: 0.0,
232 },
233 )
234 .unwrap(),
235 loss: BinaryLogLoss::new(),
236 },
237 )
238 .unwrap()
239 }
240
241 #[test]
242 fn predict_proba_in_range() {
243 let model = make_model(2, 0.1);
244 let p = model.predict_proba(&[1.0, 2.0]).unwrap();
245 assert!(p > 0.0 && p < 1.0);
246 assert!((p - 0.5).abs() < 1e-12);
248 }
249
250 #[test]
251 fn learn_separable_data() {
252 let mut model = make_model(2, 0.5);
253 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(3);
254 for _ in 0..1000 {
255 let x1 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
257 let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
258 let y = x1 > 0.0;
259 model.learn(&[x1, x2], y).unwrap();
260 }
261 let p_pos = model.predict_proba(&[2.0, 0.0]).unwrap();
262 let p_neg = model.predict_proba(&[-2.0, 0.0]).unwrap();
263 assert!(p_pos > 0.7, "p_pos = {p_pos}");
264 assert!(p_neg < 0.3, "p_neg = {p_neg}");
265 }
266
267 #[test]
268 fn predict_does_not_update_state() {
269 let model = make_model(1, 0.1);
270 let _ = model.predict_proba(&[1.0]).unwrap();
271 assert_eq!(model.samples_seen(), 0);
272 }
273
274 #[test]
275 fn dimension_mismatch_rejected() {
276 let mut model = make_model(3, 0.1);
277 assert!(model.predict_proba(&[1.0, 2.0]).is_err());
278 assert!(model.learn(&[1.0, 2.0], true).is_err());
279 }
280
281 #[test]
282 fn reset_clears_state() {
283 let mut model = make_model(1, 0.1);
284 model.learn(&[1.0], true).unwrap();
285 model.reset();
286 assert_eq!(model.samples_seen(), 0);
287 assert!((model.predict_proba(&[1.0]).unwrap() - 0.5).abs() < 1e-12);
288 }
289
290 #[test]
291 fn weighted_learning_zero_is_noop_and_positive_updates() {
292 let mut model = make_model(1, 0.1);
293 let before = model.clone();
294 model.learn_weighted(&[1.0], true, 0.0).unwrap();
295 assert_eq!(model.weights(), before.weights());
296 assert_eq!(model.samples_seen(), 0);
297 model.learn_weighted(&[1.0], true, 2.0).unwrap();
298 assert_eq!(model.samples_seen(), 1);
299 assert!(model.weights()[0] > 0.0);
300 }
301}