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
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        // gradient of log loss w.r.t. logit is (p - y)
116        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        // cold start: weights=0, intercept=0 -> sigmoid(0) = 0.5
206        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            // class 1: x1 > 0, class 0: x1 < 0
215            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}