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        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        // cold start: weights=0, intercept=0 -> sigmoid(0) = 0.5
247        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            // class 1: x1 > 0, class 0: x1 < 0
256            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}