rill_ml/models/
baseline.rs1use 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#[derive(Debug, Clone)]
15#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
16#[non_exhaustive]
17pub struct BaselineConfig {
18 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#[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 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 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#[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 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 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#[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 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 pub const fn alpha(&self) -> f64 {
176 self.ew.alpha()
177 }
178
179 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}