1use crate::error::{RillError, ensure_finite_target};
14#[cfg(feature = "serde")]
15use crate::persistence::ValidateState;
16use crate::traits::{OnlineBinaryClassifier, OnlineRegressor, Transformer};
17
18#[derive(Debug, Clone)]
20#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
21pub struct RegressionPipeline<T, M> {
22 transformer: T,
23 model: M,
24}
25
26impl<T, M> RegressionPipeline<T, M>
27where
28 T: Transformer,
29 M: OnlineRegressor,
30{
31 pub fn new(transformer: T, model: M) -> Result<Self, RillError> {
36 if transformer.output_dim() != model.feature_count() {
37 return Err(RillError::DimensionMismatch {
38 expected: model.feature_count(),
39 actual: transformer.output_dim(),
40 });
41 }
42 Ok(Self { transformer, model })
43 }
44
45 pub fn transformer(&self) -> &T {
47 &self.transformer
48 }
49
50 pub fn model(&self) -> &M {
52 &self.model
53 }
54
55 pub fn learn_transactional(&mut self, features: &[f64], target: f64) -> Result<(), RillError>
62 where
63 T: Clone,
64 M: Clone,
65 {
66 let mut next_transformer = self.transformer.clone();
67 let mut next_model = self.model.clone();
68 next_transformer.update(features)?;
69 let transformed = next_transformer.transform(features)?;
70 next_model.learn(&transformed, target)?;
71 self.transformer = next_transformer;
72 self.model = next_model;
73 Ok(())
74 }
75}
76
77impl<T, M> OnlineRegressor for RegressionPipeline<T, M>
78where
79 T: Transformer,
80 M: OnlineRegressor,
81{
82 fn feature_count(&self) -> usize {
83 self.transformer.input_dim()
84 }
85
86 fn samples_seen(&self) -> u64 {
87 self.transformer.samples_seen()
88 }
89
90 fn predict(&self, features: &[f64]) -> Result<f64, RillError> {
91 let transformed = self.transformer.transform(features)?;
92 self.model.predict(&transformed)
93 }
94
95 fn learn(&mut self, features: &[f64], target: f64) -> Result<(), RillError> {
96 ensure_finite_target(target)?;
97 self.transformer.update(features)?;
98 let transformed = self.transformer.transform(features)?;
99 self.model.learn(&transformed, target)
100 }
101
102 fn reset(&mut self) {
103 self.transformer.reset();
104 self.model.reset();
105 }
106}
107
108#[derive(Debug, Clone)]
110#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
111pub struct ClassificationPipeline<T, M> {
112 transformer: T,
113 model: M,
114}
115
116impl<T, M> ClassificationPipeline<T, M>
117where
118 T: Transformer,
119 M: OnlineBinaryClassifier,
120{
121 pub fn new(transformer: T, model: M) -> Result<Self, RillError> {
123 if transformer.output_dim() != model.feature_count() {
124 return Err(RillError::DimensionMismatch {
125 expected: model.feature_count(),
126 actual: transformer.output_dim(),
127 });
128 }
129 Ok(Self { transformer, model })
130 }
131
132 pub fn transformer(&self) -> &T {
134 &self.transformer
135 }
136
137 pub fn model(&self) -> &M {
139 &self.model
140 }
141
142 pub fn learn_transactional(&mut self, features: &[f64], target: bool) -> Result<(), RillError>
144 where
145 T: Clone,
146 M: Clone,
147 {
148 let mut next_transformer = self.transformer.clone();
149 let mut next_model = self.model.clone();
150 next_transformer.update(features)?;
151 let transformed = next_transformer.transform(features)?;
152 next_model.learn(&transformed, target)?;
153 self.transformer = next_transformer;
154 self.model = next_model;
155 Ok(())
156 }
157}
158
159impl<T, M> OnlineBinaryClassifier for ClassificationPipeline<T, M>
160where
161 T: Transformer,
162 M: OnlineBinaryClassifier,
163{
164 fn feature_count(&self) -> usize {
165 self.transformer.input_dim()
166 }
167
168 fn samples_seen(&self) -> u64 {
169 self.transformer.samples_seen()
170 }
171
172 fn predict_proba(&self, features: &[f64]) -> Result<f64, RillError> {
173 let transformed = self.transformer.transform(features)?;
174 self.model.predict_proba(&transformed)
175 }
176
177 fn learn(&mut self, features: &[f64], target: bool) -> Result<(), RillError> {
178 self.transformer.update(features)?;
179 let transformed = self.transformer.transform(features)?;
180 self.model.learn(&transformed, target)
181 }
182
183 fn reset(&mut self) {
184 self.transformer.reset();
185 self.model.reset();
186 }
187}
188
189#[cfg(feature = "serde")]
190impl<T, M> ValidateState for RegressionPipeline<T, M>
191where
192 T: Transformer + ValidateState,
193 M: OnlineRegressor + ValidateState,
194{
195 fn validate_state(&self) -> Result<(), RillError> {
196 self.transformer.validate_state()?;
197 self.model.validate_state()?;
198 if self.transformer.output_dim() != self.model.feature_count() {
199 return Err(RillError::InvalidState(format!(
200 "regression pipeline dimension mismatch: transformer output_dim {} != model feature_count {}",
201 self.transformer.output_dim(),
202 self.model.feature_count()
203 )));
204 }
205 Ok(())
206 }
207}
208
209#[cfg(feature = "serde")]
210impl<T, M> ValidateState for ClassificationPipeline<T, M>
211where
212 T: Transformer + ValidateState,
213 M: OnlineBinaryClassifier + ValidateState,
214{
215 fn validate_state(&self) -> Result<(), RillError> {
216 self.transformer.validate_state()?;
217 self.model.validate_state()?;
218 if self.transformer.output_dim() != self.model.feature_count() {
219 return Err(RillError::InvalidState(format!(
220 "classification pipeline dimension mismatch: transformer output_dim {} != model feature_count {}",
221 self.transformer.output_dim(),
222 self.model.feature_count()
223 )));
224 }
225 Ok(())
226 }
227}
228
229#[cfg(test)]
230mod tests {
231 use super::*;
232 use crate::metrics::Mae;
233 use crate::models::{LinearRegression, LinearRegressionConfig};
234 use crate::optim::{Optimizer, SgdConfig};
235 use crate::preprocessing::StandardScaler;
236 use crate::traits::Metric;
237 use rand::SeedableRng;
238
239 #[test]
240 fn pipeline_predict_does_not_update_transformer() {
241 let d = 2;
242 let scaler = StandardScaler::new(d).unwrap();
243 let model = LinearRegression::new(
244 d,
245 LinearRegressionConfig {
246 optimizer: Optimizer::sgd(d, SgdConfig::default()).unwrap(),
247 loss: Default::default(),
248 },
249 )
250 .unwrap();
251 let mut pipe = RegressionPipeline::new(scaler, model).unwrap();
252
253 let _ = pipe.predict(&[1.0, 2.0]).unwrap();
254 assert_eq!(pipe.transformer().samples_seen(), 0);
255
256 pipe.learn(&[1.0, 2.0], 3.0).unwrap();
257 assert_eq!(pipe.transformer().samples_seen(), 1);
258 }
259
260 #[test]
261 fn failed_pipeline_learn_does_not_mutate_either_stage() {
262 let scaler = StandardScaler::new(1).unwrap();
263 let model = LinearRegression::new(
264 1,
265 LinearRegressionConfig {
266 optimizer: Optimizer::sgd(1, SgdConfig::default()).unwrap(),
267 loss: Default::default(),
268 },
269 )
270 .unwrap();
271 let mut pipe = RegressionPipeline::new(scaler, model).unwrap();
272
273 assert!(pipe.learn_transactional(&[1.0], f64::NAN).is_err());
274 assert_eq!(pipe.transformer().samples_seen(), 0);
275 assert_eq!(pipe.model().samples_seen(), 0);
276 }
277
278 #[test]
279 fn pipeline_dimension_mismatch_rejected() {
280 let scaler = StandardScaler::new(3).unwrap();
281 let model = LinearRegression::new(
282 2,
283 LinearRegressionConfig {
284 optimizer: Optimizer::sgd(2, SgdConfig::default()).unwrap(),
285 loss: Default::default(),
286 },
287 )
288 .unwrap();
289 assert!(RegressionPipeline::new(scaler, model).is_err());
290 }
291
292 #[test]
293 fn pipeline_learns_linear_relation() {
294 let d = 2;
295 let scaler = StandardScaler::new(d).unwrap();
296 let model = LinearRegression::new(
297 d,
298 LinearRegressionConfig {
299 optimizer: Optimizer::sgd(
300 d,
301 SgdConfig {
302 learning_rate: 0.05,
303 l2: 0.0,
304 },
305 )
306 .unwrap(),
307 loss: Default::default(),
308 },
309 )
310 .unwrap();
311 let mut pipe = RegressionPipeline::new(scaler, model).unwrap();
312 let mut mae = Mae::default();
313
314 let mut rng = rand_chacha::ChaCha8Rng::seed_from_u64(11);
315 for _ in 0..500 {
316 let x1 = rand::Rng::gen_range(&mut rng, 0.0..1.0);
317 let x2 = rand::Rng::gen_range(&mut rng, 0.0..1.0);
318 let y = 3.0 * x1 + 2.0 * x2;
319 let pred = pipe.predict(&[x1, x2]).unwrap();
320 mae.update(y, pred).unwrap();
321 pipe.learn(&[x1, x2], y).unwrap();
322 }
323 let final_mae = mae.value().unwrap();
324 assert!(final_mae < 1.0, "final MAE too high: {final_mae}");
325 }
326}