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 for w in &mut self.weights {
182 *w = 0.0;
183 }
184 self.intercept = 0.0;
185 self.optimizer.reset();
186 self.samples_seen = 0;
187 }
188}
189
190#[cfg(feature = "serde")]
191impl ValidateState for LogisticRegression {
192 fn validate_state(&self) -> Result<(), RillError> {
193 if self.feature_count == 0 {
194 return Err(RillError::EmptyFeatures);
195 }
196 if self.weights.len() != self.feature_count {
197 return Err(RillError::InvalidState(format!(
198 "logistic regression weights length {} does not match feature_count {}",
199 self.weights.len(),
200 self.feature_count
201 )));
202 }
203 if self.optimizer.param_count() != self.feature_count + 1 {
204 return Err(RillError::InvalidState(format!(
205 "logistic regression optimizer param_count {} does not match feature_count+1 {}",
206 self.optimizer.param_count(),
207 self.feature_count + 1
208 )));
209 }
210 ensure_finite("intercept", self.intercept)?;
211 for &w in &self.weights {
212 ensure_finite("weights", w)?;
213 }
214 self.optimizer.validate_state()?;
215 Ok(())
216 }
217}
218
219#[cfg(test)]
220mod tests {
221 use super::*;
222 use crate::optim::SgdConfig;
223 use rand::SeedableRng;
224
225 fn make_model(d: usize, lr: f64) -> LogisticRegression {
226 LogisticRegression::new(
227 d,
228 LogisticRegressionConfig {
229 optimizer: Optimizer::sgd(
230 d,
231 SgdConfig {
232 learning_rate: lr,
233 l2: 0.0,
234 },
235 )
236 .unwrap(),
237 loss: BinaryLogLoss::new(),
238 },
239 )
240 .unwrap()
241 }
242
243 #[test]
244 fn predict_proba_in_range() {
245 let model = make_model(2, 0.1);
246 let p = model.predict_proba(&[1.0, 2.0]).unwrap();
247 assert!(p > 0.0 && p < 1.0);
248 assert!((p - 0.5).abs() < 1e-12);
250 }
251
252 #[test]
253 fn learn_separable_data() {
254 let mut model = make_model(2, 0.5);
255 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(3);
256 for _ in 0..1000 {
257 let x1 = rand::Rng::gen_range(&mut rng, -2.0..2.0);
259 let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
260 let y = x1 > 0.0;
261 model.learn(&[x1, x2], y).unwrap();
262 }
263 let p_pos = model.predict_proba(&[2.0, 0.0]).unwrap();
264 let p_neg = model.predict_proba(&[-2.0, 0.0]).unwrap();
265 assert!(p_pos > 0.7, "p_pos = {p_pos}");
266 assert!(p_neg < 0.3, "p_neg = {p_neg}");
267 }
268
269 #[test]
270 fn predict_does_not_update_state() {
271 let model = make_model(1, 0.1);
272 let _ = model.predict_proba(&[1.0]).unwrap();
273 assert_eq!(model.samples_seen(), 0);
274 }
275
276 #[test]
277 fn dimension_mismatch_rejected() {
278 let mut model = make_model(3, 0.1);
279 assert!(model.predict_proba(&[1.0, 2.0]).is_err());
280 assert!(model.learn(&[1.0, 2.0], true).is_err());
281 }
282
283 #[test]
284 fn reset_clears_state() {
285 let mut model = make_model(1, 0.1);
286 model.learn(&[1.0], true).unwrap();
287 model.reset();
288 assert_eq!(model.samples_seen(), 0);
289 assert!((model.predict_proba(&[1.0]).unwrap() - 0.5).abs() < 1e-12);
290 }
291
292 #[test]
293 fn weighted_learning_zero_is_noop_and_positive_updates() {
294 let mut model = make_model(1, 0.1);
295 let before = model.clone();
296 model.learn_weighted(&[1.0], true, 0.0).unwrap();
297 assert_eq!(model.weights(), before.weights());
298 assert_eq!(model.samples_seen(), 0);
299 model.learn_weighted(&[1.0], true, 2.0).unwrap();
300 assert_eq!(model.samples_seen(), 1);
301 assert!(model.weights()[0] > 0.0);
302 }
303}