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 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 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 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}