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        for w in &mut self.weights {
209            *w = 0.0;
210        }
211        self.intercept = 0.0;
212        self.optimizer.reset();
213        self.samples_seen = 0;
214    }
215}
216
217#[cfg(feature = "serde")]
218impl ValidateState for LinearRegression {
219    fn validate_state(&self) -> Result<(), RillError> {
220        if self.feature_count == 0 {
221            return Err(RillError::EmptyFeatures);
222        }
223        if self.weights.len() != self.feature_count {
224            return Err(RillError::InvalidState(format!(
225                "linear regression weights length {} does not match feature_count {}",
226                self.weights.len(),
227                self.feature_count
228            )));
229        }
230        if self.optimizer.param_count() != self.feature_count + 1 {
231            return Err(RillError::InvalidState(format!(
232                "linear regression optimizer param_count {} does not match feature_count+1 {}",
233                self.optimizer.param_count(),
234                self.feature_count + 1
235            )));
236        }
237        ensure_finite("intercept", self.intercept)?;
238        for &w in &self.weights {
239            ensure_finite("weights", w)?;
240        }
241        self.optimizer.validate_state()?;
242        Ok(())
243    }
244}
245
246#[cfg(test)]
247mod tests {
248    use super::*;
249    use crate::optim::{AdaGradConfig, SgdConfig};
250    use rand::SeedableRng;
251
252    fn make_sgd(lr: f64, l2: f64, d: usize) -> Optimizer {
253        Optimizer::sgd(
254            d,
255            SgdConfig {
256                learning_rate: lr,
257                l2,
258            },
259        )
260        .unwrap()
261    }
262
263    #[test]
264    fn predict_cold_start_returns_intercept() {
265        let model = LinearRegression::new(
266            2,
267            LinearRegressionConfig {
268                optimizer: make_sgd(0.1, 0.0, 2),
269                loss: RegressionLoss::default(),
270            },
271        )
272        .unwrap();
273        assert_eq!(model.predict(&[1.0, 2.0]).unwrap(), 0.0);
274    }
275
276    #[test]
277    fn learn_reduces_loss_on_linear_data() {
278        let mut model = LinearRegression::new(
279            2,
280            LinearRegressionConfig {
281                optimizer: make_sgd(0.05, 0.0, 2),
282                loss: RegressionLoss::default(),
283            },
284        )
285        .unwrap();
286        // y = 2*x1 - 0.5*x2 + 1
287        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(7);
288        let mut first_loss = 0.0;
289        let mut last_loss = 0.0;
290        for i in 0..500 {
291            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
292            let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
293            let y = 2.0 * x1 - 0.5 * x2 + 1.0;
294            let pred = model.predict(&[x1, x2]).unwrap();
295            let l = crate::loss::SquaredError::loss(pred, y);
296            if i < 10 {
297                first_loss += l;
298            }
299            if i >= 490 {
300                last_loss += l;
301            }
302            model.learn(&[x1, x2], y).unwrap();
303        }
304        assert!(last_loss < first_loss, "loss should decrease");
305        // weights should be approximately [2, -0.5]
306        assert!((model.weights()[0] - 2.0).abs() < 0.3);
307        assert!((model.weights()[1] + 0.5).abs() < 0.3);
308        assert!((model.intercept() - 1.0).abs() < 0.3);
309    }
310
311    #[test]
312    fn predict_does_not_update_state() {
313        let model = LinearRegression::new(
314            1,
315            LinearRegressionConfig {
316                optimizer: make_sgd(0.1, 0.0, 1),
317                loss: RegressionLoss::default(),
318            },
319        )
320        .unwrap();
321        let _ = model.predict(&[1.0]).unwrap();
322        assert_eq!(model.samples_seen(), 0);
323    }
324
325    #[test]
326    fn dimension_mismatch_rejected() {
327        let mut model = LinearRegression::new(
328            3,
329            LinearRegressionConfig {
330                optimizer: make_sgd(0.1, 0.0, 3),
331                loss: RegressionLoss::default(),
332            },
333        )
334        .unwrap();
335        assert!(model.predict(&[1.0, 2.0]).is_err());
336        assert!(model.learn(&[1.0, 2.0], 1.0).is_err());
337    }
338
339    #[test]
340    fn optimizer_feature_count_mismatch_rejected() {
341        let config = LinearRegressionConfig {
342            optimizer: make_sgd(0.1, 0.0, 3),
343            loss: RegressionLoss::default(),
344        };
345        assert!(LinearRegression::new(2, config).is_err());
346    }
347
348    #[test]
349    fn adagrad_works() {
350        let mut model = LinearRegression::new(
351            1,
352            LinearRegressionConfig {
353                optimizer: Optimizer::adagrad(
354                    1,
355                    AdaGradConfig {
356                        learning_rate: 0.5,
357                        l2: 0.0,
358                        epsilon: 1e-8,
359                    },
360                )
361                .unwrap(),
362                loss: RegressionLoss::default(),
363            },
364        )
365        .unwrap();
366        for _ in 0..200 {
367            model.learn(&[1.0], 5.0).unwrap();
368        }
369        assert!((model.predict(&[1.0]).unwrap() - 5.0).abs() < 0.5);
370    }
371
372    #[test]
373    fn reset_clears_state() {
374        let mut model = LinearRegression::new(
375            1,
376            LinearRegressionConfig {
377                optimizer: make_sgd(0.1, 0.0, 1),
378                loss: RegressionLoss::default(),
379            },
380        )
381        .unwrap();
382        model.learn(&[1.0], 5.0).unwrap();
383        model.reset();
384        assert_eq!(model.samples_seen(), 0);
385        assert_eq!(model.predict(&[1.0]).unwrap(), 0.0);
386    }
387
388    #[test]
389    fn weighted_learning_scales_gradient_and_zero_is_noop() {
390        let mut weighted = LinearRegression::new(
391            1,
392            LinearRegressionConfig {
393                optimizer: make_sgd(0.1, 0.0, 1),
394                loss: RegressionLoss::default(),
395            },
396        )
397        .unwrap();
398        let before = weighted.clone();
399        weighted.learn_weighted(&[2.0], 3.0, 0.0).unwrap();
400        assert_eq!(weighted.weights(), before.weights());
401        assert_eq!(weighted.samples_seen(), before.samples_seen());
402        weighted.learn_weighted(&[2.0], 3.0, 2.0).unwrap();
403        assert_eq!(weighted.samples_seen(), 1);
404        assert!(weighted.weights()[0] > 0.0);
405    }
406
407    #[test]
408    fn non_finite_target_rejected() {
409        let mut model = LinearRegression::new(
410            1,
411            LinearRegressionConfig {
412                optimizer: make_sgd(0.1, 0.0, 1),
413                loss: RegressionLoss::default(),
414            },
415        )
416        .unwrap();
417        assert!(model.learn(&[1.0], f64::NAN).is_err());
418    }
419
420    #[test]
421    fn huber_loss_works() {
422        let mut model = LinearRegression::new(
423            1,
424            LinearRegressionConfig {
425                optimizer: make_sgd(0.1, 0.0, 1),
426                loss: RegressionLoss::Huber(crate::loss::HuberLoss::new(1.0).unwrap()),
427            },
428        )
429        .unwrap();
430        model.learn(&[1.0], 1.0).unwrap();
431        assert_eq!(model.samples_seen(), 1);
432    }
433}