1use crate::error::{RillError, ensure_finite_target};
14#[cfg(feature = "serde")]
15use crate::persistence::ValidateState;
16use crate::traits::{OnlineBinaryClassifier, OnlineRegressor, Transformer};
17use crate::weighted::{WeightedOnlineBinaryClassifier, WeightedOnlineRegressor, validate_weight};
18
19#[derive(Debug, Clone)]
21#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
22pub struct RegressionPipeline<T, M> {
23 transformer: T,
24 model: M,
25}
26
27impl<T, M> RegressionPipeline<T, M>
28where
29 T: Transformer,
30 M: OnlineRegressor,
31{
32 pub fn new(transformer: T, model: M) -> Result<Self, RillError> {
37 if transformer.output_dim() != model.feature_count() {
38 return Err(RillError::DimensionMismatch {
39 expected: model.feature_count(),
40 actual: transformer.output_dim(),
41 });
42 }
43 Ok(Self { transformer, model })
44 }
45
46 pub fn transformer(&self) -> &T {
48 &self.transformer
49 }
50
51 pub fn model(&self) -> &M {
53 &self.model
54 }
55
56 pub fn learn_transactional(&mut self, features: &[f64], target: f64) -> Result<(), RillError>
63 where
64 T: Clone,
65 M: Clone,
66 {
67 let mut next_transformer = self.transformer.clone();
68 let mut next_model = self.model.clone();
69 next_transformer.update(features)?;
70 let transformed = next_transformer.transform(features)?;
71 next_model.learn(&transformed, target)?;
72 self.transformer = next_transformer;
73 self.model = next_model;
74 Ok(())
75 }
76}
77
78impl<T, M> OnlineRegressor for RegressionPipeline<T, M>
79where
80 T: Transformer,
81 M: OnlineRegressor,
82{
83 fn feature_count(&self) -> usize {
84 self.transformer.input_dim()
85 }
86
87 fn samples_seen(&self) -> u64 {
88 self.transformer.samples_seen()
89 }
90
91 fn predict(&self, features: &[f64]) -> Result<f64, RillError> {
92 let transformed = self.transformer.transform(features)?;
93 self.model.predict(&transformed)
94 }
95
96 fn learn(&mut self, features: &[f64], target: f64) -> Result<(), RillError> {
97 ensure_finite_target(target)?;
98 self.transformer.update(features)?;
99 let transformed = self.transformer.transform(features)?;
100 self.model.learn(&transformed, target)
101 }
102
103 fn reset(&mut self) {
104 self.transformer.reset();
105 self.model.reset();
106 }
107}
108
109impl<T, M> WeightedOnlineRegressor for RegressionPipeline<T, M>
112where
113 T: Transformer + Clone,
114 M: WeightedOnlineRegressor + Clone,
115{
116 fn learn_weighted(
117 &mut self,
118 features: &[f64],
119 target: f64,
120 weight: f64,
121 ) -> Result<(), RillError> {
122 validate_weight(weight)?;
123 ensure_finite_target(target)?;
124 if weight == 0.0 {
125 let _ = self.transformer.transform(features)?;
126 return Ok(());
127 }
128 let mut next_transformer = self.transformer.clone();
129 let mut next_model = self.model.clone();
130 next_transformer.update(features)?;
131 let transformed = next_transformer.transform(features)?;
132 next_model.learn_weighted(&transformed, target, weight)?;
133 self.transformer = next_transformer;
134 self.model = next_model;
135 Ok(())
136 }
137}
138
139#[derive(Debug, Clone)]
141#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
142pub struct ClassificationPipeline<T, M> {
143 transformer: T,
144 model: M,
145}
146
147impl<T, M> ClassificationPipeline<T, M>
148where
149 T: Transformer,
150 M: OnlineBinaryClassifier,
151{
152 pub fn new(transformer: T, model: M) -> Result<Self, RillError> {
154 if transformer.output_dim() != model.feature_count() {
155 return Err(RillError::DimensionMismatch {
156 expected: model.feature_count(),
157 actual: transformer.output_dim(),
158 });
159 }
160 Ok(Self { transformer, model })
161 }
162
163 pub fn transformer(&self) -> &T {
165 &self.transformer
166 }
167
168 pub fn model(&self) -> &M {
170 &self.model
171 }
172
173 pub fn learn_transactional(&mut self, features: &[f64], target: bool) -> Result<(), RillError>
175 where
176 T: Clone,
177 M: Clone,
178 {
179 let mut next_transformer = self.transformer.clone();
180 let mut next_model = self.model.clone();
181 next_transformer.update(features)?;
182 let transformed = next_transformer.transform(features)?;
183 next_model.learn(&transformed, target)?;
184 self.transformer = next_transformer;
185 self.model = next_model;
186 Ok(())
187 }
188}
189
190impl<T, M> OnlineBinaryClassifier for ClassificationPipeline<T, M>
191where
192 T: Transformer,
193 M: OnlineBinaryClassifier,
194{
195 fn feature_count(&self) -> usize {
196 self.transformer.input_dim()
197 }
198
199 fn samples_seen(&self) -> u64 {
200 self.transformer.samples_seen()
201 }
202
203 fn predict_proba(&self, features: &[f64]) -> Result<f64, RillError> {
204 let transformed = self.transformer.transform(features)?;
205 self.model.predict_proba(&transformed)
206 }
207
208 fn learn(&mut self, features: &[f64], target: bool) -> Result<(), RillError> {
209 self.transformer.update(features)?;
210 let transformed = self.transformer.transform(features)?;
211 self.model.learn(&transformed, target)
212 }
213
214 fn reset(&mut self) {
215 self.transformer.reset();
216 self.model.reset();
217 }
218}
219
220impl<T, M> WeightedOnlineBinaryClassifier for ClassificationPipeline<T, M>
223where
224 T: Transformer + Clone,
225 M: WeightedOnlineBinaryClassifier + Clone,
226{
227 fn learn_weighted(
228 &mut self,
229 features: &[f64],
230 target: bool,
231 weight: f64,
232 ) -> Result<(), RillError> {
233 validate_weight(weight)?;
234 if weight == 0.0 {
235 let _ = self.transformer.transform(features)?;
236 return Ok(());
237 }
238 let mut next_transformer = self.transformer.clone();
239 let mut next_model = self.model.clone();
240 next_transformer.update(features)?;
241 let transformed = next_transformer.transform(features)?;
242 next_model.learn_weighted(&transformed, target, weight)?;
243 self.transformer = next_transformer;
244 self.model = next_model;
245 Ok(())
246 }
247}
248
249#[cfg(feature = "serde")]
250impl<T, M> ValidateState for RegressionPipeline<T, M>
251where
252 T: Transformer + ValidateState,
253 M: OnlineRegressor + ValidateState,
254{
255 fn validate_state(&self) -> Result<(), RillError> {
256 self.transformer.validate_state()?;
257 self.model.validate_state()?;
258 if self.transformer.output_dim() != self.model.feature_count() {
259 return Err(RillError::InvalidState(format!(
260 "regression pipeline dimension mismatch: transformer output_dim {} != model feature_count {}",
261 self.transformer.output_dim(),
262 self.model.feature_count()
263 )));
264 }
265 Ok(())
266 }
267}
268
269#[cfg(feature = "serde")]
270impl<T, M> ValidateState for ClassificationPipeline<T, M>
271where
272 T: Transformer + ValidateState,
273 M: OnlineBinaryClassifier + ValidateState,
274{
275 fn validate_state(&self) -> Result<(), RillError> {
276 self.transformer.validate_state()?;
277 self.model.validate_state()?;
278 if self.transformer.output_dim() != self.model.feature_count() {
279 return Err(RillError::InvalidState(format!(
280 "classification pipeline dimension mismatch: transformer output_dim {} != model feature_count {}",
281 self.transformer.output_dim(),
282 self.model.feature_count()
283 )));
284 }
285 Ok(())
286 }
287}
288
289#[cfg(test)]
290mod tests {
291 use super::*;
292 use crate::metrics::Mae;
293 use crate::models::{LinearRegression, LinearRegressionConfig};
294 use crate::optim::{Optimizer, SgdConfig};
295 use crate::preprocessing::StandardScaler;
296 use crate::traits::Metric;
297 use rand::SeedableRng;
298
299 #[test]
300 fn pipeline_predict_does_not_update_transformer() {
301 let d = 2;
302 let scaler = StandardScaler::new(d).unwrap();
303 let model = LinearRegression::new(
304 d,
305 LinearRegressionConfig {
306 optimizer: Optimizer::sgd(d, SgdConfig::default()).unwrap(),
307 loss: Default::default(),
308 },
309 )
310 .unwrap();
311 let mut pipe = RegressionPipeline::new(scaler, model).unwrap();
312
313 let _ = pipe.predict(&[1.0, 2.0]).unwrap();
314 assert_eq!(pipe.transformer().samples_seen(), 0);
315
316 pipe.learn(&[1.0, 2.0], 3.0).unwrap();
317 assert_eq!(pipe.transformer().samples_seen(), 1);
318 }
319
320 #[test]
321 fn failed_pipeline_learn_does_not_mutate_either_stage() {
322 let scaler = StandardScaler::new(1).unwrap();
323 let model = LinearRegression::new(
324 1,
325 LinearRegressionConfig {
326 optimizer: Optimizer::sgd(1, SgdConfig::default()).unwrap(),
327 loss: Default::default(),
328 },
329 )
330 .unwrap();
331 let mut pipe = RegressionPipeline::new(scaler, model).unwrap();
332
333 assert!(pipe.learn_transactional(&[1.0], f64::NAN).is_err());
334 assert_eq!(pipe.transformer().samples_seen(), 0);
335 assert_eq!(pipe.model().samples_seen(), 0);
336 }
337
338 #[test]
339 fn weighted_pipeline_is_atomic_and_zero_weight_is_noop() {
340 let scaler = StandardScaler::new(1).unwrap();
341 let model = LinearRegression::new(
342 1,
343 LinearRegressionConfig {
344 optimizer: Optimizer::sgd(1, SgdConfig::default()).unwrap(),
345 loss: Default::default(),
346 },
347 )
348 .unwrap();
349 let mut pipe = RegressionPipeline::new(scaler, model).unwrap();
350 pipe.learn_weighted(&[2.0], 3.0, 0.0).unwrap();
351 assert_eq!(pipe.transformer().samples_seen(), 0);
352 assert_eq!(pipe.model().samples_seen(), 0);
353 pipe.learn_weighted(&[2.0], 3.0, 2.0).unwrap();
354 assert_eq!(pipe.transformer().samples_seen(), 1);
355 assert_eq!(pipe.model().samples_seen(), 1);
356 let before_transformer = pipe.transformer().clone();
357 let before_model = pipe.model().clone();
358 assert!(pipe.learn_weighted(&[f64::NAN], 3.0, 1.0).is_err());
359 assert_eq!(
360 pipe.transformer().samples_seen(),
361 before_transformer.samples_seen()
362 );
363 assert_eq!(pipe.model().weights(), before_model.weights());
364 }
365
366 #[test]
367 fn pipeline_dimension_mismatch_rejected() {
368 let scaler = StandardScaler::new(3).unwrap();
369 let model = LinearRegression::new(
370 2,
371 LinearRegressionConfig {
372 optimizer: Optimizer::sgd(2, SgdConfig::default()).unwrap(),
373 loss: Default::default(),
374 },
375 )
376 .unwrap();
377 assert!(RegressionPipeline::new(scaler, model).is_err());
378 }
379
380 #[test]
381 fn pipeline_learns_linear_relation() {
382 let d = 2;
383 let scaler = StandardScaler::new(d).unwrap();
384 let model = LinearRegression::new(
385 d,
386 LinearRegressionConfig {
387 optimizer: Optimizer::sgd(
388 d,
389 SgdConfig {
390 learning_rate: 0.05,
391 l2: 0.0,
392 },
393 )
394 .unwrap(),
395 loss: Default::default(),
396 },
397 )
398 .unwrap();
399 let mut pipe = RegressionPipeline::new(scaler, model).unwrap();
400 let mut mae = Mae::default();
401
402 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(11);
403 for _ in 0..500 {
404 let x1 = rand::Rng::gen_range(&mut rng, 0.0..1.0);
405 let x2 = rand::Rng::gen_range(&mut rng, 0.0..1.0);
406 let y = 3.0 * x1 + 2.0 * x2;
407 let pred = pipe.predict(&[x1, x2]).unwrap();
408 mae.update(y, pred).unwrap();
409 pipe.learn(&[x1, x2], y).unwrap();
410 }
411 let final_mae = mae.value().unwrap();
412 assert!(final_mae < 1.0, "final MAE too high: {final_mae}");
413 }
414}