Skip to main content

rill_ml/models/
baseline.rs

1//! Baseline regressors.
2//!
3//! These simple models serve as comparison baselines. A complex online model
4//! should be compared against at least [`MeanRegressor`] and
5//! [`ExponentiallyWeightedMeanRegressor`] before being considered useful.
6
7use crate::error::{RillError, checked_increment, ensure_finite, ensure_finite_target};
8#[cfg(feature = "serde")]
9use crate::persistence::ValidateState;
10use crate::stats::{ExponentiallyWeightedMean, Mean};
11use crate::traits::{OnlineRegressor, OnlineStatistic};
12
13/// Configuration shared by baseline regressors.
14#[derive(Debug, Clone)]
15#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
16#[non_exhaustive]
17pub struct BaselineConfig {
18    /// Prediction returned before any target has been observed.
19    pub initial_prediction: f64,
20}
21
22impl Default for BaselineConfig {
23    fn default() -> Self {
24        Self {
25            initial_prediction: 0.0,
26        }
27    }
28}
29
30fn validate_baseline_config(config: &BaselineConfig) -> Result<(), RillError> {
31    ensure_finite("initial_prediction", config.initial_prediction)
32}
33
34/// A regressor that always predicts the running mean of observed targets.
35///
36/// This is the simplest meaningful online regression baseline.
37#[derive(Debug, Clone)]
38#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
39pub struct MeanRegressor {
40    config: BaselineConfig,
41    mean: Mean,
42}
43
44impl MeanRegressor {
45    /// Create a new mean regressor with the given configuration.
46    pub fn new(config: BaselineConfig) -> Result<Self, RillError> {
47        validate_baseline_config(&config)?;
48        Ok(Self {
49            config,
50            mean: Mean::new(),
51        })
52    }
53
54    /// The current running mean of targets.
55    pub const fn mean(&self) -> f64 {
56        self.mean.value()
57    }
58}
59
60impl OnlineRegressor for MeanRegressor {
61    fn feature_count(&self) -> usize {
62        0
63    }
64
65    fn samples_seen(&self) -> u64 {
66        self.mean.samples_seen()
67    }
68
69    fn predict(&self, _features: &[f64]) -> Result<f64, RillError> {
70        if self.mean.count() == 0 {
71            Ok(self.config.initial_prediction)
72        } else {
73            Ok(self.mean.value())
74        }
75    }
76
77    fn learn(&mut self, _features: &[f64], target: f64) -> Result<(), RillError> {
78        ensure_finite_target(target)?;
79        self.mean.update(target)
80    }
81
82    fn reset(&mut self) {
83        self.mean.reset();
84    }
85}
86
87impl Default for MeanRegressor {
88    fn default() -> Self {
89        Self::new(BaselineConfig::default()).expect("default config is valid")
90    }
91}
92
93/// A regressor that always predicts the last observed target.
94#[derive(Debug, Clone)]
95#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
96pub struct LastValueRegressor {
97    config: BaselineConfig,
98    last_value: Option<f64>,
99    count: u64,
100}
101
102impl LastValueRegressor {
103    /// Create a new last-value regressor.
104    pub fn new(config: BaselineConfig) -> Result<Self, RillError> {
105        validate_baseline_config(&config)?;
106        Ok(Self {
107            config,
108            last_value: None,
109            count: 0,
110        })
111    }
112
113    /// The last observed target, if any.
114    pub const fn last_value(&self) -> Option<f64> {
115        self.last_value
116    }
117}
118
119impl OnlineRegressor for LastValueRegressor {
120    fn feature_count(&self) -> usize {
121        0
122    }
123
124    fn samples_seen(&self) -> u64 {
125        self.count
126    }
127
128    fn predict(&self, _features: &[f64]) -> Result<f64, RillError> {
129        Ok(self.last_value.unwrap_or(self.config.initial_prediction))
130    }
131
132    fn learn(&mut self, _features: &[f64], target: f64) -> Result<(), RillError> {
133        ensure_finite_target(target)?;
134        let next_count = checked_increment(self.count, "last-value sample")?;
135        self.last_value = Some(target);
136        self.count = next_count;
137        Ok(())
138    }
139
140    fn reset(&mut self) {
141        self.last_value = None;
142        self.count = 0;
143    }
144}
145
146impl Default for LastValueRegressor {
147    fn default() -> Self {
148        Self::new(BaselineConfig::default()).expect("default config is valid")
149    }
150}
151
152/// A regressor that predicts an exponentially weighted mean of targets.
153///
154/// Suitable when recent observations are more relevant than older ones.
155#[derive(Debug, Clone)]
156#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
157pub struct ExponentiallyWeightedMeanRegressor {
158    config: BaselineConfig,
159    ew: ExponentiallyWeightedMean,
160}
161
162impl ExponentiallyWeightedMeanRegressor {
163    /// Create a new EW mean regressor.
164    ///
165    /// `alpha` must be in `(0, 1]`.
166    pub fn new(alpha: f64, config: BaselineConfig) -> Result<Self, RillError> {
167        validate_baseline_config(&config)?;
168        Ok(Self {
169            config,
170            ew: ExponentiallyWeightedMean::new(alpha)?,
171        })
172    }
173
174    /// The configured alpha.
175    pub const fn alpha(&self) -> f64 {
176        self.ew.alpha()
177    }
178
179    /// The current weighted mean.
180    pub const fn value(&self) -> f64 {
181        self.ew.value()
182    }
183}
184
185impl OnlineRegressor for ExponentiallyWeightedMeanRegressor {
186    fn feature_count(&self) -> usize {
187        0
188    }
189
190    fn samples_seen(&self) -> u64 {
191        self.ew.samples_seen()
192    }
193
194    fn predict(&self, _features: &[f64]) -> Result<f64, RillError> {
195        if self.ew.count() == 0 {
196            Ok(self.config.initial_prediction)
197        } else {
198            Ok(self.ew.value())
199        }
200    }
201
202    fn learn(&mut self, _features: &[f64], target: f64) -> Result<(), RillError> {
203        ensure_finite_target(target)?;
204        self.ew.update(target)
205    }
206
207    fn reset(&mut self) {
208        self.ew.reset();
209    }
210}
211
212#[cfg(feature = "serde")]
213impl ValidateState for MeanRegressor {
214    fn validate_state(&self) -> Result<(), RillError> {
215        validate_baseline_config(&self.config)?;
216        self.mean.validate_state()
217    }
218}
219
220#[cfg(feature = "serde")]
221impl ValidateState for LastValueRegressor {
222    fn validate_state(&self) -> Result<(), RillError> {
223        validate_baseline_config(&self.config)?;
224        if let Some(value) = self.last_value {
225            ensure_finite("last_value", value)?;
226        }
227        Ok(())
228    }
229}
230
231#[cfg(feature = "serde")]
232impl ValidateState for ExponentiallyWeightedMeanRegressor {
233    fn validate_state(&self) -> Result<(), RillError> {
234        validate_baseline_config(&self.config)?;
235        self.ew.validate_state()
236    }
237}
238
239#[cfg(test)]
240mod tests {
241    use super::*;
242
243    #[test]
244    fn mean_regressor_cold_start() {
245        let r = MeanRegressor::default();
246        assert_eq!(r.predict(&[]).unwrap(), 0.0);
247    }
248
249    #[test]
250    fn mean_regressor_predicts_running_mean() {
251        let mut r = MeanRegressor::default();
252        r.learn(&[], 10.0).unwrap();
253        r.learn(&[], 20.0).unwrap();
254        assert_eq!(r.predict(&[]).unwrap(), 15.0);
255    }
256
257    #[test]
258    fn last_value_regressor_cold_start() {
259        let r = LastValueRegressor::default();
260        assert_eq!(r.predict(&[]).unwrap(), 0.0);
261    }
262
263    #[test]
264    fn last_value_regressor_tracks_last() {
265        let mut r = LastValueRegressor::default();
266        r.learn(&[], 10.0).unwrap();
267        r.learn(&[], 20.0).unwrap();
268        assert_eq!(r.predict(&[]).unwrap(), 20.0);
269    }
270
271    #[test]
272    fn ew_mean_regressor_cold_start() {
273        let r = ExponentiallyWeightedMeanRegressor::new(0.5, BaselineConfig::default()).unwrap();
274        assert_eq!(r.predict(&[]).unwrap(), 0.0);
275    }
276
277    #[test]
278    fn ew_mean_regressor_weights_recent() {
279        let mut r =
280            ExponentiallyWeightedMeanRegressor::new(0.5, BaselineConfig::default()).unwrap();
281        r.learn(&[], 10.0).unwrap();
282        r.learn(&[], 20.0).unwrap();
283        assert!((r.predict(&[]).unwrap() - 15.0).abs() < 1e-12);
284    }
285
286    #[test]
287    fn initial_prediction_custom() {
288        let r = MeanRegressor::new(BaselineConfig {
289            initial_prediction: 42.0,
290        })
291        .unwrap();
292        assert_eq!(r.predict(&[]).unwrap(), 42.0);
293    }
294
295    #[test]
296    fn non_finite_target_rejected() {
297        let mut r = MeanRegressor::default();
298        assert!(r.learn(&[], f64::NAN).is_err());
299    }
300}