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
119impl OnlineRegressor for LinearRegression {
120    fn feature_count(&self) -> usize {
121        self.feature_count
122    }
123
124    fn samples_seen(&self) -> u64 {
125        self.samples_seen
126    }
127
128    fn predict(&self, features: &[f64]) -> Result<f64, RillError> {
129        self.predict_inner(features)
130    }
131
132    fn learn(&mut self, features: &[f64], target: f64) -> Result<(), RillError> {
133        validate_features(self.feature_count, features)?;
134        ensure_finite_target(target)?;
135        let next_samples = checked_increment(self.samples_seen, "linear regression sample")?;
136
137        let prediction = self.predict_inner(features)?;
138        let grad = self.loss.gradient(prediction, target);
139        ensure_finite("loss gradient", grad)?;
140
141        // gradient w.r.t. each weight w_i is grad * x_i
142        let grad_weights = features
143            .iter()
144            .map(|&feature| {
145                let gradient = grad * feature;
146                ensure_finite("weight gradient", gradient)?;
147                Ok(gradient)
148            })
149            .collect::<Result<Vec<_>, RillError>>()?;
150        let grad_intercept = grad;
151
152        self.optimizer.step(
153            &mut self.weights,
154            &mut self.intercept,
155            &grad_weights,
156            grad_intercept,
157        )?;
158        self.samples_seen = next_samples;
159        Ok(())
160    }
161
162    fn reset(&mut self) {
163        for w in &mut self.weights {
164            *w = 0.0;
165        }
166        self.intercept = 0.0;
167        self.optimizer.reset();
168        self.samples_seen = 0;
169    }
170}
171
172#[cfg(feature = "serde")]
173impl ValidateState for LinearRegression {
174    fn validate_state(&self) -> Result<(), RillError> {
175        if self.feature_count == 0 {
176            return Err(RillError::EmptyFeatures);
177        }
178        if self.weights.len() != self.feature_count {
179            return Err(RillError::InvalidState(format!(
180                "linear regression weights length {} does not match feature_count {}",
181                self.weights.len(),
182                self.feature_count
183            )));
184        }
185        if self.optimizer.param_count() != self.feature_count + 1 {
186            return Err(RillError::InvalidState(format!(
187                "linear regression optimizer param_count {} does not match feature_count+1 {}",
188                self.optimizer.param_count(),
189                self.feature_count + 1
190            )));
191        }
192        ensure_finite("intercept", self.intercept)?;
193        for &w in &self.weights {
194            ensure_finite("weights", w)?;
195        }
196        self.optimizer.validate_state()?;
197        Ok(())
198    }
199}
200
201#[cfg(test)]
202mod tests {
203    use super::*;
204    use crate::optim::{AdaGradConfig, SgdConfig};
205    use rand::SeedableRng;
206
207    fn make_sgd(lr: f64, l2: f64, d: usize) -> Optimizer {
208        Optimizer::sgd(
209            d,
210            SgdConfig {
211                learning_rate: lr,
212                l2,
213            },
214        )
215        .unwrap()
216    }
217
218    #[test]
219    fn predict_cold_start_returns_intercept() {
220        let model = LinearRegression::new(
221            2,
222            LinearRegressionConfig {
223                optimizer: make_sgd(0.1, 0.0, 2),
224                loss: RegressionLoss::default(),
225            },
226        )
227        .unwrap();
228        assert_eq!(model.predict(&[1.0, 2.0]).unwrap(), 0.0);
229    }
230
231    #[test]
232    fn learn_reduces_loss_on_linear_data() {
233        let mut model = LinearRegression::new(
234            2,
235            LinearRegressionConfig {
236                optimizer: make_sgd(0.05, 0.0, 2),
237                loss: RegressionLoss::default(),
238            },
239        )
240        .unwrap();
241        // y = 2*x1 - 0.5*x2 + 1
242        let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(7);
243        let mut first_loss = 0.0;
244        let mut last_loss = 0.0;
245        for i in 0..500 {
246            let x1 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
247            let x2 = rand::Rng::gen_range(&mut rng, -1.0..1.0);
248            let y = 2.0 * x1 - 0.5 * x2 + 1.0;
249            let pred = model.predict(&[x1, x2]).unwrap();
250            let l = crate::loss::SquaredError::loss(pred, y);
251            if i < 10 {
252                first_loss += l;
253            }
254            if i >= 490 {
255                last_loss += l;
256            }
257            model.learn(&[x1, x2], y).unwrap();
258        }
259        assert!(last_loss < first_loss, "loss should decrease");
260        // weights should be approximately [2, -0.5]
261        assert!((model.weights()[0] - 2.0).abs() < 0.3);
262        assert!((model.weights()[1] + 0.5).abs() < 0.3);
263        assert!((model.intercept() - 1.0).abs() < 0.3);
264    }
265
266    #[test]
267    fn predict_does_not_update_state() {
268        let model = LinearRegression::new(
269            1,
270            LinearRegressionConfig {
271                optimizer: make_sgd(0.1, 0.0, 1),
272                loss: RegressionLoss::default(),
273            },
274        )
275        .unwrap();
276        let _ = model.predict(&[1.0]).unwrap();
277        assert_eq!(model.samples_seen(), 0);
278    }
279
280    #[test]
281    fn dimension_mismatch_rejected() {
282        let mut model = LinearRegression::new(
283            3,
284            LinearRegressionConfig {
285                optimizer: make_sgd(0.1, 0.0, 3),
286                loss: RegressionLoss::default(),
287            },
288        )
289        .unwrap();
290        assert!(model.predict(&[1.0, 2.0]).is_err());
291        assert!(model.learn(&[1.0, 2.0], 1.0).is_err());
292    }
293
294    #[test]
295    fn optimizer_feature_count_mismatch_rejected() {
296        let config = LinearRegressionConfig {
297            optimizer: make_sgd(0.1, 0.0, 3),
298            loss: RegressionLoss::default(),
299        };
300        assert!(LinearRegression::new(2, config).is_err());
301    }
302
303    #[test]
304    fn adagrad_works() {
305        let mut model = LinearRegression::new(
306            1,
307            LinearRegressionConfig {
308                optimizer: Optimizer::adagrad(
309                    1,
310                    AdaGradConfig {
311                        learning_rate: 0.5,
312                        l2: 0.0,
313                        epsilon: 1e-8,
314                    },
315                )
316                .unwrap(),
317                loss: RegressionLoss::default(),
318            },
319        )
320        .unwrap();
321        for _ in 0..200 {
322            model.learn(&[1.0], 5.0).unwrap();
323        }
324        assert!((model.predict(&[1.0]).unwrap() - 5.0).abs() < 0.5);
325    }
326
327    #[test]
328    fn reset_clears_state() {
329        let mut model = LinearRegression::new(
330            1,
331            LinearRegressionConfig {
332                optimizer: make_sgd(0.1, 0.0, 1),
333                loss: RegressionLoss::default(),
334            },
335        )
336        .unwrap();
337        model.learn(&[1.0], 5.0).unwrap();
338        model.reset();
339        assert_eq!(model.samples_seen(), 0);
340        assert_eq!(model.predict(&[1.0]).unwrap(), 0.0);
341    }
342
343    #[test]
344    fn non_finite_target_rejected() {
345        let mut model = LinearRegression::new(
346            1,
347            LinearRegressionConfig {
348                optimizer: make_sgd(0.1, 0.0, 1),
349                loss: RegressionLoss::default(),
350            },
351        )
352        .unwrap();
353        assert!(model.learn(&[1.0], f64::NAN).is_err());
354    }
355
356    #[test]
357    fn huber_loss_works() {
358        let mut model = LinearRegression::new(
359            1,
360            LinearRegressionConfig {
361                optimizer: make_sgd(0.1, 0.0, 1),
362                loss: RegressionLoss::Huber(crate::loss::HuberLoss::new(1.0).unwrap()),
363            },
364        )
365        .unwrap();
366        model.learn(&[1.0], 1.0).unwrap();
367        assert_eq!(model.samples_seen(), 1);
368    }
369}