1use 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#[derive(Debug, Clone)]
19#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
20#[non_exhaustive]
21pub struct LinearRegressionConfig {
22 pub optimizer: Optimizer,
24 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#[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 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 pub fn weights(&self) -> &[f64] {
96 &self.weights
97 }
98
99 pub const fn intercept(&self) -> f64 {
101 self.intercept
102 }
103
104 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 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 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 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 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}