Skip to main content

rill_ml/models/
logistic_regression.rs

1//! Online logistic regression for binary classification.
2//!
3//! Uses a numerically stable sigmoid and binary cross-entropy (log) loss.
4
5use 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/// Configuration for [`LogisticRegression`].
15#[derive(Debug, Clone)]
16#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
17#[non_exhaustive]
18pub struct LogisticRegressionConfig {
19    /// The optimizer (SGD or AdaGrad).
20    pub optimizer: Optimizer,
21    /// The log loss configuration.
22    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/// Online logistic regression model.
35///
36/// Predicts `P(y=1 | x) = sigmoid(w·x + b)`. Learning uses the gradient of
37/// the binary log loss w.r.t. the logit, which simplifies to `p - y`.
38#[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    /// Create a new logistic regression model.
51    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    /// The learned weights.
72    pub fn weights(&self) -> &[f64] {
73        &self.weights
74    }
75
76    /// The learned intercept (bias).
77    pub const fn intercept(&self) -> f64 {
78        self.intercept
79    }
80
81    /// Compute the logit `w·x + b`.
82    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    /// Learn with a finite, non-negative sample weight.
96    ///
97    /// `weight = 0` validates inputs but is a complete state no-op.
98    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        // gradient of log loss w.r.t. logit is (p - y)
159        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        // cold start: weights=0, intercept=0 -> sigmoid(0) = 0.5
249        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            // class 1: x1 > 0, class 0: x1 < 0
258            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}