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