Skip to main content

rill_ml/
pipeline.rs

1//! Static two-segment pipelines: transformer + model.
2//!
3//! The learning contract is fixed:
4//! - `predict(x)`: `transform(x)` → `model.predict()`. No state updates.
5//! - `learn(x, y)`: `transformer.update(x)` → `transform(x)` → `model.learn()`.
6//! - `learn_transactional` is the failure-atomic variant: neither stage is
7//!   committed unless all three operations succeed.
8//!
9//! The transformer never sees the target `y`, so there is no label leakage in
10//! the progressive-evaluation sense (the prediction for the current sample is
11//! produced *before* any state update).
12
13use 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/// A pipeline combining a transformer and a regressor.
20#[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    /// Create a new pipeline.
33    ///
34    /// Returns an error if the transformer's output dimension does not match
35    /// the model's feature count.
36    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    /// Borrow the transformer.
47    pub fn transformer(&self) -> &T {
48        &self.transformer
49    }
50
51    /// Borrow the model.
52    pub fn model(&self) -> &M {
53        &self.model
54    }
55
56    /// Learn one sample with all-or-nothing state changes.
57    ///
58    /// This clones both stages, applies the update to the clones, and commits
59    /// them only after every operation succeeds. Prefer this at reliability
60    /// boundaries; use [`OnlineRegressor::learn`] when avoiding the clone cost
61    /// is more important and both stages already provide atomic updates.
62    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
109/// Weighted labels flow to the model while the transformer observes each
110/// positive-weight event once. The whole transition is transactional.
111impl<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/// A pipeline combining a transformer and a binary classifier.
140#[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    /// Create a new classification pipeline.
153    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    /// Borrow the transformer.
164    pub fn transformer(&self) -> &T {
165        &self.transformer
166    }
167
168    /// Borrow the model.
169    pub fn model(&self) -> &M {
170        &self.model
171    }
172
173    /// Learn one classification sample with all-or-nothing state changes.
174    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
220/// Weighted labels flow to the classifier while the transformer observes each
221/// positive-weight event once. The whole transition is transactional.
222impl<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}