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};
17
18/// A pipeline combining a transformer and a regressor.
19#[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    /// Create a new pipeline.
32    ///
33    /// Returns an error if the transformer's output dimension does not match
34    /// the model's feature count.
35    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    /// Borrow the transformer.
46    pub fn transformer(&self) -> &T {
47        &self.transformer
48    }
49
50    /// Borrow the model.
51    pub fn model(&self) -> &M {
52        &self.model
53    }
54
55    /// Learn one sample with all-or-nothing state changes.
56    ///
57    /// This clones both stages, applies the update to the clones, and commits
58    /// them only after every operation succeeds. Prefer this at reliability
59    /// boundaries; use [`OnlineRegressor::learn`] when avoiding the clone cost
60    /// is more important and both stages already provide atomic updates.
61    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/// A pipeline combining a transformer and a binary classifier.
109#[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    /// Create a new classification pipeline.
122    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    /// Borrow the transformer.
133    pub fn transformer(&self) -> &T {
134        &self.transformer
135    }
136
137    /// Borrow the model.
138    pub fn model(&self) -> &M {
139        &self.model
140    }
141
142    /// Learn one classification sample with all-or-nothing state changes.
143    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}