Skip to main content

rill_ml/models/
linear_regression.rs

1//! Online linear regression using SGD or AdaGrad.
2//!
3//! The model learns `y ≈ w·x + b` incrementally, one sample at a time.
4//! Prediction is side-effect free; learning computes the gradient of the
5//! configured loss and applies one optimizer step.
6
7use crate::error::{
8    RillError, checked_finite_add, checked_increment, ensure_finite, ensure_finite_target,
9    validate_features,
10};
11use crate::loss::RegressionLoss;
12use crate::optim::Optimizer;
13#[cfg(feature = "serde")]
14use crate::persistence::ValidateState;
15use crate::traits::OnlineRegressor;
16
17/// Configuration for [`LinearRegression`].
18#[derive(Debug, Clone)]
19#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
20#[non_exhaustive]
21pub struct LinearRegressionConfig {
22    /// The optimizer to use (SGD or AdaGrad).
23    pub optimizer: Optimizer,
24    /// The loss function (SquaredError or Huber).
25    pub loss: RegressionLoss,
26}
27
28impl Default for LinearRegressionConfig {
29    fn default() -> Self {
30        Self {
31            optimizer: Optimizer::sgd(1, Default::default()).expect("default optimizer"),
32            loss: RegressionLoss::default(),
33        }
34    }
35}
36
37/// Online linear regression model.
38///
39/// # Examples
40///
41/// ```
42/// use rill_ml::{
43///     models::{LinearRegression, LinearRegressionConfig},
44///     optim::{Optimizer, SgdConfig},
45///     OnlineRegressor,
46/// };
47///
48/// let feature_count = 2;
49/// let mut sgd = SgdConfig::default();
50/// sgd.learning_rate = 0.1;
51/// sgd.l2 = 0.0;
52/// let mut lr_config = LinearRegressionConfig::default();
53/// lr_config.optimizer = Optimizer::sgd(feature_count, sgd).unwrap();
54/// let mut model = LinearRegression::new(feature_count, lr_config).unwrap();
55///
56/// let prediction = model.predict(&[1.0, 2.0]).unwrap();
57/// model.learn(&[1.0, 2.0], 3.0).unwrap();
58/// ```
59#[derive(Debug, Clone)]
60#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
61pub struct LinearRegression {
62    feature_count: usize,
63    weights: Vec<f64>,
64    intercept: f64,
65    optimizer: Optimizer,
66    loss: RegressionLoss,
67    samples_seen: u64,
68}
69
70impl LinearRegression {
71    /// Create a new linear regression model.
72    ///
73    /// The optimizer's feature count must match `feature_count`.
74    pub fn new(feature_count: usize, config: LinearRegressionConfig) -> Result<Self, RillError> {
75        if feature_count == 0 {
76            return Err(RillError::EmptyFeatures);
77        }
78        if config.optimizer.param_count() != feature_count + 1 {
79            return Err(RillError::DimensionMismatch {
80                expected: feature_count + 1,
81                actual: config.optimizer.param_count(),
82            });
83        }
84        Ok(Self {
85            feature_count,
86            weights: vec![0.0; feature_count],
87            intercept: 0.0,
88            optimizer: config.optimizer,
89            loss: config.loss,
90            samples_seen: 0,
91        })
92    }
93
94    /// The learned weights.
95    pub fn weights(&self) -> &[f64] {
96        &self.weights
97    }
98
99    /// The learned intercept (bias).
100    pub const fn intercept(&self) -> f64 {
101        self.intercept
102    }
103
104    /// Compute the prediction `w·x + b` without updating state.
105    fn predict_inner(&self, features: &[f64]) -> Result<f64, RillError> {
106        validate_features(self.feature_count, features)?;
107        let dot = self.weights.iter().zip(features.iter()).try_fold(
108            0.0,
109            |sum, (&weight, &feature)| {
110                let term = weight * feature;
111                ensure_finite("linear prediction term", term)?;
112                checked_finite_add(sum, term, "linear prediction")
113            },
114        )?;
115        checked_finite_add(dot, self.intercept, "linear prediction")
116    }
117
118    /// Learn with a finite, non-negative sample weight.
119    ///
120    /// `weight = 0` validates the sample but leaves all model and optimizer
121    /// state unchanged. Positive weights scale the loss gradient.
122    pub fn learn_weighted(
123        &mut self,
124        features: &[f64],
125        target: f64,
126        weight: f64,
127    ) -> Result<(), RillError> {
128        crate::weighted::validate_weight(weight)?;
129        validate_features(self.feature_count, features)?;
130        ensure_finite_target(target)?;
131        if weight == 0.0 {
132            return Ok(());
133        }
134        let next_samples = checked_increment(self.samples_seen, "linear regression sample")?;
135        let prediction = self.predict_inner(features)?;
136        let grad = self.loss.gradient(prediction, target) * weight;
137        ensure_finite("weighted loss gradient", grad)?;
138        let grad_weights = features
139            .iter()
140            .map(|&feature| {
141                let gradient = grad * feature;
142                ensure_finite("weighted weight gradient", gradient)?;
143                Ok(gradient)
144            })
145            .collect::<Result<Vec<_>, RillError>>()?;
146        self.optimizer
147            .step(&mut self.weights, &mut self.intercept, &grad_weights, grad)?;
148        self.samples_seen = next_samples;
149        Ok(())
150    }
151}
152
153impl crate::weighted::WeightedOnlineRegressor for LinearRegression {
154    fn learn_weighted(
155        &mut self,
156        features: &[f64],
157        target: f64,
158        weight: f64,
159    ) -> Result<(), RillError> {
160        LinearRegression::learn_weighted(self, features, target, weight)
161    }
162}
163
164impl OnlineRegressor for LinearRegression {
165    fn feature_count(&self) -> usize {
166        self.feature_count
167    }
168
169    fn samples_seen(&self) -> u64 {
170        self.samples_seen
171    }
172
173    fn predict(&self, features: &[f64]) -> Result<f64, RillError> {
174        self.predict_inner(features)
175    }
176
177    fn learn(&mut self, features: &[f64], target: f64) -> Result<(), RillError> {
178        validate_features(self.feature_count, features)?;
179        ensure_finite_target(target)?;
180        let next_samples = checked_increment(self.samples_seen, "linear regression sample")?;
181
182        let prediction = self.predict_inner(features)?;
183        let grad = self.loss.gradient(prediction, target);
184        ensure_finite("loss gradient", grad)?;
185
186        // gradient w.r.t. each weight w_i is grad * x_i
187        let grad_weights = features
188            .iter()
189            .map(|&feature| {
190                let gradient = grad * feature;
191                ensure_finite("weight gradient", gradient)?;
192                Ok(gradient)
193            })
194            .collect::<Result<Vec<_>, RillError>>()?;
195        let grad_intercept = grad;
196
197        self.optimizer.step(
198            &mut self.weights,
199            &mut self.intercept,
200            &grad_weights,
201            grad_intercept,
202        )?;
203        self.samples_seen = next_samples;
204        Ok(())
205    }
206
207    fn reset(&mut self) {
208        self.weights.fill(0.0);
209        self.intercept = 0.0;
210        self.optimizer.reset();
211        self.samples_seen = 0;
212    }
213}
214
215#[cfg(feature = "serde")]
216impl ValidateState for LinearRegression {
217    fn validate_state(&self) -> Result<(), RillError> {
218        if self.feature_count == 0 {
219            return Err(RillError::EmptyFeatures);
220        }
221        if self.weights.len() != self.feature_count {
222            return Err(RillError::InvalidState(format!(
223                "linear regression weights length {} does not match feature_count {}",
224                self.weights.len(),
225                self.feature_count
226            )));
227        }
228        if self.optimizer.param_count() != self.feature_count + 1 {
229            return Err(RillError::InvalidState(format!(
230                "linear regression optimizer param_count {} does not match feature_count+1 {}",
231                self.optimizer.param_count(),
232                self.feature_count + 1
233            )));
234        }
235        ensure_finite("intercept", self.intercept)?;
236        for &w in &self.weights {
237            ensure_finite("weights", w)?;
238        }
239        self.optimizer.validate_state()?;
240        Ok(())
241    }
242}
243
244#[cfg(test)]
245mod tests {
246    use super::*;
247    use crate::optim::{AdaGradConfig, SgdConfig};
248    use rand::SeedableRng;
249
250    fn make_sgd(lr: f64, l2: f64, d: usize) -> Optimizer {
251        Optimizer::sgd(
252            d,
253            SgdConfig {
254                learning_rate: lr,
255                l2,
256            },
257        )
258        .unwrap()
259    }
260
261    #[test]
262    fn predict_cold_start_returns_intercept() {
263        let model = LinearRegression::new(
264            2,
265            LinearRegressionConfig {
266                optimizer: make_sgd(0.1, 0.0, 2),
267                loss: RegressionLoss::default(),
268            },
269        )
270        .unwrap();
271        assert_eq!(model.predict(&[1.0, 2.0]).unwrap(), 0.0);
272    }
273
274    #[test]
275    fn learn_reduces_loss_on_linear_data() {
276        let mut model = LinearRegression::new(
277            2,
278            LinearRegressionConfig {
279                optimizer: make_sgd(0.05, 0.0, 2),
280                loss: RegressionLoss::default(),
281            },
282        )
283        .unwrap();
284        // y = 2*x1 - 0.5*x2 + 1
285        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(7);
286        let mut first_loss = 0.0;
287        let mut last_loss = 0.0;
288        for i in 0..500 {
289            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
290            let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
291            let y = 2.0 * x1 - 0.5 * x2 + 1.0;
292            let pred = model.predict(&[x1, x2]).unwrap();
293            let l = crate::loss::SquaredError::loss(pred, y);
294            if i < 10 {
295                first_loss += l;
296            }
297            if i >= 490 {
298                last_loss += l;
299            }
300            model.learn(&[x1, x2], y).unwrap();
301        }
302        assert!(last_loss < first_loss, "loss should decrease");
303        // weights should be approximately [2, -0.5]
304        assert!((model.weights()[0] - 2.0).abs() < 0.3);
305        assert!((model.weights()[1] + 0.5).abs() < 0.3);
306        assert!((model.intercept() - 1.0).abs() < 0.3);
307    }
308
309    #[test]
310    fn predict_does_not_update_state() {
311        let model = LinearRegression::new(
312            1,
313            LinearRegressionConfig {
314                optimizer: make_sgd(0.1, 0.0, 1),
315                loss: RegressionLoss::default(),
316            },
317        )
318        .unwrap();
319        let _ = model.predict(&[1.0]).unwrap();
320        assert_eq!(model.samples_seen(), 0);
321    }
322
323    #[test]
324    fn dimension_mismatch_rejected() {
325        let mut model = LinearRegression::new(
326            3,
327            LinearRegressionConfig {
328                optimizer: make_sgd(0.1, 0.0, 3),
329                loss: RegressionLoss::default(),
330            },
331        )
332        .unwrap();
333        assert!(model.predict(&[1.0, 2.0]).is_err());
334        assert!(model.learn(&[1.0, 2.0], 1.0).is_err());
335    }
336
337    #[test]
338    fn optimizer_feature_count_mismatch_rejected() {
339        let config = LinearRegressionConfig {
340            optimizer: make_sgd(0.1, 0.0, 3),
341            loss: RegressionLoss::default(),
342        };
343        assert!(LinearRegression::new(2, config).is_err());
344    }
345
346    #[test]
347    fn adagrad_works() {
348        let mut model = LinearRegression::new(
349            1,
350            LinearRegressionConfig {
351                optimizer: Optimizer::adagrad(
352                    1,
353                    AdaGradConfig {
354                        learning_rate: 0.5,
355                        l2: 0.0,
356                        epsilon: 1e-8,
357                    },
358                )
359                .unwrap(),
360                loss: RegressionLoss::default(),
361            },
362        )
363        .unwrap();
364        for _ in 0..200 {
365            model.learn(&[1.0], 5.0).unwrap();
366        }
367        assert!((model.predict(&[1.0]).unwrap() - 5.0).abs() < 0.5);
368    }
369
370    #[test]
371    fn reset_clears_state() {
372        let mut model = LinearRegression::new(
373            1,
374            LinearRegressionConfig {
375                optimizer: make_sgd(0.1, 0.0, 1),
376                loss: RegressionLoss::default(),
377            },
378        )
379        .unwrap();
380        model.learn(&[1.0], 5.0).unwrap();
381        model.reset();
382        assert_eq!(model.samples_seen(), 0);
383        assert_eq!(model.predict(&[1.0]).unwrap(), 0.0);
384    }
385
386    #[test]
387    fn weighted_learning_scales_gradient_and_zero_is_noop() {
388        let mut weighted = LinearRegression::new(
389            1,
390            LinearRegressionConfig {
391                optimizer: make_sgd(0.1, 0.0, 1),
392                loss: RegressionLoss::default(),
393            },
394        )
395        .unwrap();
396        let before = weighted.clone();
397        weighted.learn_weighted(&[2.0], 3.0, 0.0).unwrap();
398        assert_eq!(weighted.weights(), before.weights());
399        assert_eq!(weighted.samples_seen(), before.samples_seen());
400        weighted.learn_weighted(&[2.0], 3.0, 2.0).unwrap();
401        assert_eq!(weighted.samples_seen(), 1);
402        assert!(weighted.weights()[0] > 0.0);
403    }
404
405    #[test]
406    fn non_finite_target_rejected() {
407        let mut model = LinearRegression::new(
408            1,
409            LinearRegressionConfig {
410                optimizer: make_sgd(0.1, 0.0, 1),
411                loss: RegressionLoss::default(),
412            },
413        )
414        .unwrap();
415        assert!(model.learn(&[1.0], f64::NAN).is_err());
416    }
417
418    #[test]
419    fn huber_loss_works() {
420        let mut model = LinearRegression::new(
421            1,
422            LinearRegressionConfig {
423                optimizer: make_sgd(0.1, 0.0, 1),
424                loss: RegressionLoss::Huber(crate::loss::HuberLoss::new(1.0).unwrap()),
425            },
426        )
427        .unwrap();
428        model.learn(&[1.0], 1.0).unwrap();
429        assert_eq!(model.samples_seen(), 1);
430    }
431}