Skip to main content

rill_ml/diagnostics/
prediction_report.rs

1//! Unified prediction report.
2//!
3//! Combines [`ResidualInterval`], [`WarmupTracker`], and [`TrainingSummary`]
4//! into a single diagnostic wrapper that produces an immutable
5//! [`PredictionReport`] for each prediction. This keeps the base model API
6//! clean: a model returns a plain prediction, and the caller can wrap it with
7//! [`PredictionReporter`] to obtain intervals, confidence levels, and
8//! warmup/baseline comparisons.
9//!
10//! Space complexity: `O(1)`.
11
12use crate::diagnostics::prediction_interval::{ResidualInterval, ResidualIntervalConfig};
13use crate::diagnostics::training_summary::{TrainingSummary, TrainingSummaryConfig};
14use crate::diagnostics::warmup::{WarmupConfig, WarmupState, WarmupTracker};
15use crate::error::{RillError, ensure_finite};
16
17/// Coarse confidence level derived from the warmup state and baseline
18/// comparison.
19#[derive(Debug, Clone, Copy, PartialEq, Eq)]
20#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
21#[non_exhaustive]
22pub enum Confidence {
23    /// No data has been observed yet.
24    None,
25    /// The model is still warming up or has degraded.
26    Low,
27    /// The model is usable but not yet stable.
28    Medium,
29    /// The model is stable and beating the baseline.
30    High,
31}
32
33impl Confidence {
34    /// Returns a short, stable string identifier.
35    ///
36    /// Possible return values: `"none"`, `"low"`, `"medium"`, `"high"`.
37    pub const fn as_str(&self) -> &'static str {
38        match self {
39            Confidence::None => "none",
40            Confidence::Low => "low",
41            Confidence::Medium => "medium",
42            Confidence::High => "high",
43        }
44    }
45}
46
47/// An immutable snapshot of diagnostics for a single prediction.
48#[derive(Debug, Clone)]
49#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
50pub struct PredictionReport {
51    prediction: f64,
52    lower_bound: Option<f64>,
53    upper_bound: Option<f64>,
54    confidence: Confidence,
55    samples_seen: u64,
56    recent_error: Option<f64>,
57    baseline_error: Option<f64>,
58    warmup_state: WarmupState,
59    beats_baseline: Option<bool>,
60}
61
62impl PredictionReport {
63    /// The prediction that this report was generated for.
64    pub const fn prediction(&self) -> f64 {
65        self.prediction
66    }
67
68    /// Lower bound of the prediction interval, or `None` if insufficient data.
69    pub const fn lower_bound(&self) -> Option<f64> {
70        self.lower_bound
71    }
72
73    /// Upper bound of the prediction interval, or `None` if insufficient data.
74    pub const fn upper_bound(&self) -> Option<f64> {
75        self.upper_bound
76    }
77
78    /// Coarse confidence level for this prediction.
79    pub const fn confidence(&self) -> Confidence {
80        self.confidence
81    }
82
83    /// Total number of samples observed so far.
84    pub const fn samples_seen(&self) -> u64 {
85        self.samples_seen
86    }
87
88    /// Recent (EW mean) absolute error, or `None` if no errors recorded.
89    pub const fn recent_error(&self) -> Option<f64> {
90        self.recent_error
91    }
92
93    /// Baseline error for comparison, or `None` if not set.
94    pub const fn baseline_error(&self) -> Option<f64> {
95        self.baseline_error
96    }
97
98    /// Current warmup state of the model.
99    pub const fn warmup_state(&self) -> WarmupState {
100        self.warmup_state
101    }
102
103    /// Whether the model is currently beating the baseline.
104    ///
105    /// Returns `None` if either recent error or baseline error is unavailable.
106    pub const fn beats_baseline(&self) -> Option<bool> {
107        self.beats_baseline
108    }
109}
110
111/// Diagnostic wrapper that integrates interval estimation, warmup tracking,
112/// and training summary statistics.
113///
114/// Produces a [`PredictionReport`] for each prediction without storing raw
115/// samples. The underlying model API is not affected: callers feed
116/// `(prediction, truth)` pairs via [`observe`](Self::observe) and request a
117/// report via [`report`](Self::report) when needed.
118///
119/// # Examples
120///
121/// ```
122/// use rill_ml::diagnostics::PredictionReporter;
123///
124/// let mut reporter = PredictionReporter::default();
125/// reporter.observe(10.0, 11.0).unwrap();
126/// reporter.observe(10.0, 9.0).unwrap();
127///
128/// let report = reporter.report(10.0).unwrap();
129/// assert_eq!(report.prediction(), 10.0);
130/// assert!(report.lower_bound().is_some());
131/// assert_eq!(report.samples_seen(), 2);
132/// ```
133#[derive(Debug, Clone)]
134#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
135pub struct PredictionReporter {
136    interval: ResidualInterval,
137    warmup: WarmupTracker,
138    summary: TrainingSummary,
139}
140
141impl PredictionReporter {
142    /// Create a new reporter with the given configurations.
143    ///
144    /// Each sub-component is constructed independently; configuration errors
145    /// are propagated as [`RillError`].
146    pub fn new(
147        interval_config: ResidualIntervalConfig,
148        warmup_config: WarmupConfig,
149        summary_config: TrainingSummaryConfig,
150    ) -> Result<Self, RillError> {
151        Ok(Self {
152            interval: ResidualInterval::new(interval_config)?,
153            warmup: WarmupTracker::new(warmup_config)?,
154            summary: TrainingSummary::new(summary_config)?,
155        })
156    }
157
158    /// Observe a prediction and its ground truth.
159    ///
160    /// Updates the interval estimator, warmup tracker, and training summary.
161    /// Non-finite inputs are rejected before any state is mutated.
162    pub fn observe(&mut self, prediction: f64, truth: f64) -> Result<(), RillError> {
163        self.interval.observe(prediction, truth)?;
164        let error = (truth - prediction).abs();
165        self.warmup.observe_sample(Some(error))?;
166        self.summary.record_error(error)?;
167        self.summary.record_sample()?;
168        Ok(())
169    }
170
171    /// Set the baseline error for comparison.
172    ///
173    /// Propagates to both the warmup tracker and the training summary.
174    pub fn set_baseline(&mut self, baseline: f64) -> Result<(), RillError> {
175        self.warmup.set_baseline(baseline)?;
176        self.summary.set_baseline_error(baseline)?;
177        Ok(())
178    }
179
180    /// Build an immutable report for the given prediction.
181    ///
182    /// If the interval estimator has insufficient data, the bounds are set to
183    /// `None` and no error is returned. Non-finite `prediction` values are
184    /// rejected.
185    pub fn report(&self, prediction: f64) -> Result<PredictionReport, RillError> {
186        ensure_finite("prediction", prediction)?;
187
188        let (lower_bound, upper_bound) = match self.interval.interval(prediction) {
189            Ok(iv) => (Some(iv.lower()), Some(iv.upper())),
190            Err(RillError::InsufficientData) => (None, None),
191            Err(e) => return Err(e),
192        };
193
194        let warmup_state = self.warmup.state();
195        let beats_baseline = self.summary.beats_baseline();
196        let samples_seen = self.summary.total_samples();
197        let recent_error = self.summary.recent_error();
198        let baseline_error = self.summary.baseline_error();
199
200        let confidence = match warmup_state {
201            WarmupState::NoData => Confidence::None,
202            WarmupState::WarmingUp | WarmupState::Degraded => Confidence::Low,
203            WarmupState::Usable => Confidence::Medium,
204            WarmupState::Stable => {
205                if matches!(beats_baseline, Some(true)) {
206                    Confidence::High
207                } else {
208                    Confidence::Medium
209                }
210            }
211        };
212
213        Ok(PredictionReport {
214            prediction,
215            lower_bound,
216            upper_bound,
217            confidence,
218            samples_seen,
219            recent_error,
220            baseline_error,
221            warmup_state,
222            beats_baseline,
223        })
224    }
225
226    /// Borrow the underlying training summary.
227    pub fn summary(&self) -> &TrainingSummary {
228        &self.summary
229    }
230
231    /// Current warmup state.
232    pub fn warmup_state(&self) -> WarmupState {
233        self.warmup.state()
234    }
235
236    /// Recent (EW mean) absolute error, or `None` if no errors recorded.
237    pub fn recent_error(&self) -> Option<f64> {
238        self.summary.recent_error()
239    }
240
241    /// Reset all three sub-components to their initial state.
242    pub fn reset(&mut self) {
243        self.interval.reset();
244        self.warmup.reset();
245        self.summary.reset();
246    }
247}
248
249impl Default for PredictionReporter {
250    fn default() -> Self {
251        Self::new(
252            ResidualIntervalConfig::default(),
253            WarmupConfig::default(),
254            TrainingSummaryConfig::default(),
255        )
256        .expect("default configs are valid")
257    }
258}
259
260#[cfg(test)]
261mod tests {
262    use super::*;
263
264    #[test]
265    fn confidence_as_str() {
266        assert_eq!(Confidence::None.as_str(), "none");
267        assert_eq!(Confidence::Low.as_str(), "low");
268        assert_eq!(Confidence::Medium.as_str(), "medium");
269        assert_eq!(Confidence::High.as_str(), "high");
270    }
271
272    #[test]
273    fn default_reporter_no_data() {
274        let reporter = PredictionReporter::default();
275        let r = reporter.report(0.0).unwrap();
276        assert_eq!(r.prediction(), 0.0);
277        assert_eq!(r.lower_bound(), None);
278        assert_eq!(r.upper_bound(), None);
279        assert_eq!(r.confidence(), Confidence::None);
280        assert_eq!(r.warmup_state(), WarmupState::NoData);
281        assert_eq!(r.samples_seen(), 0);
282        assert_eq!(r.recent_error(), None);
283        assert_eq!(r.baseline_error(), None);
284        assert_eq!(r.beats_baseline(), None);
285    }
286
287    #[test]
288    fn observe_then_report() {
289        let mut reporter = PredictionReporter::default();
290        reporter.observe(10.0, 11.0).unwrap(); // |error| = 1.0
291        reporter.observe(10.0, 9.0).unwrap(); // |error| = 1.0
292        let r = reporter.report(10.0).unwrap();
293        assert_eq!(r.prediction(), 10.0);
294        assert!(r.lower_bound().is_some());
295        assert!(r.upper_bound().is_some());
296        assert!(r.lower_bound().unwrap() < 10.0);
297        assert!(r.upper_bound().unwrap() > 10.0);
298        assert_eq!(r.samples_seen(), 2);
299        assert!(r.recent_error().is_some());
300    }
301
302    #[test]
303    fn set_baseline_enables_comparison() {
304        let mut reporter = PredictionReporter::default();
305        reporter.observe(0.0, 1.0).unwrap();
306        let r = reporter.report(0.0).unwrap();
307        assert_eq!(r.beats_baseline(), None);
308        assert_eq!(r.baseline_error(), None);
309        reporter.set_baseline(2.0).unwrap();
310        let r = reporter.report(0.0).unwrap();
311        assert_eq!(r.baseline_error(), Some(2.0));
312        assert_eq!(r.beats_baseline(), Some(true)); // recent_error=1.0 < 2.0
313    }
314
315    #[test]
316    fn confidence_progression() {
317        let warmup_config = WarmupConfig {
318            warming_up_threshold: 2,
319            usable_threshold: 5,
320            stable_threshold: 10,
321            degraded_error_ratio: 2.0,
322        };
323        let summary_config = TrainingSummaryConfig { error_alpha: 1.0 };
324        let mut reporter = PredictionReporter::new(
325            ResidualIntervalConfig::default(),
326            warmup_config,
327            summary_config,
328        )
329        .unwrap();
330
331        // NoData: no observations yet.
332        let r = reporter.report(0.0).unwrap();
333        assert_eq!(r.warmup_state(), WarmupState::NoData);
334        assert_eq!(r.confidence(), Confidence::None);
335
336        // WarmingUp: 1 sample (< warming_up_threshold=2).
337        reporter.observe(0.0, 0.5).unwrap();
338        let r = reporter.report(0.0).unwrap();
339        assert_eq!(r.warmup_state(), WarmupState::WarmingUp);
340        assert_eq!(r.confidence(), Confidence::Low);
341
342        // Usable: 5 samples, no baseline.
343        for _ in 0..4 {
344            reporter.observe(0.0, 0.5).unwrap();
345        }
346        let r = reporter.report(0.0).unwrap();
347        assert_eq!(r.warmup_state(), WarmupState::Usable);
348        assert_eq!(r.confidence(), Confidence::Medium);
349
350        // Set baseline; still Usable because samples < stable_threshold.
351        reporter.set_baseline(1.0).unwrap();
352        let r = reporter.report(0.0).unwrap();
353        assert_eq!(r.warmup_state(), WarmupState::Usable);
354        assert_eq!(r.confidence(), Confidence::Medium);
355
356        // Stable: 10 samples and recent_error (0.5) <= baseline (1.0).
357        for _ in 0..5 {
358            reporter.observe(0.0, 0.5).unwrap();
359        }
360        let r = reporter.report(0.0).unwrap();
361        assert_eq!(r.warmup_state(), WarmupState::Stable);
362        assert_eq!(r.confidence(), Confidence::High);
363    }
364
365    #[test]
366    fn degraded_state() {
367        let mut reporter = PredictionReporter::default();
368        reporter.set_baseline(0.4).unwrap();
369        // error=1.0 > baseline(0.4) * ratio(2.0) = 0.8
370        for _ in 0..5 {
371            reporter.observe(0.0, 1.0).unwrap();
372        }
373        let r = reporter.report(0.0).unwrap();
374        assert_eq!(r.warmup_state(), WarmupState::Degraded);
375        assert_eq!(r.confidence(), Confidence::Low);
376    }
377
378    #[test]
379    fn reset_clears_all() {
380        let mut reporter = PredictionReporter::default();
381        reporter.observe(10.0, 12.0).unwrap();
382        reporter.set_baseline(2.0).unwrap();
383        reporter.reset();
384        let r = reporter.report(0.0).unwrap();
385        assert_eq!(r.lower_bound(), None);
386        assert_eq!(r.upper_bound(), None);
387        assert_eq!(r.confidence(), Confidence::None);
388        assert_eq!(r.warmup_state(), WarmupState::NoData);
389        assert_eq!(r.samples_seen(), 0);
390        assert_eq!(r.recent_error(), None);
391        assert_eq!(r.baseline_error(), None);
392        assert_eq!(r.beats_baseline(), None);
393    }
394
395    #[test]
396    fn report_with_non_finite_prediction_errors() {
397        let mut reporter = PredictionReporter::default();
398        reporter.observe(0.0, 1.0).unwrap();
399        assert!(reporter.report(f64::NAN).is_err());
400        assert!(reporter.report(f64::INFINITY).is_err());
401        assert!(reporter.report(f64::NEG_INFINITY).is_err());
402    }
403
404    #[test]
405    fn observe_with_non_finite_rejected() {
406        let mut reporter = PredictionReporter::default();
407        assert!(reporter.observe(0.0, f64::NAN).is_err());
408        assert!(reporter.observe(0.0, f64::INFINITY).is_err());
409        assert!(reporter.observe(f64::NAN, 0.0).is_err());
410        // No state should have been recorded.
411        let r = reporter.report(0.0).unwrap();
412        assert_eq!(r.samples_seen(), 0);
413        assert_eq!(r.warmup_state(), WarmupState::NoData);
414    }
415
416    #[test]
417    fn samples_seen_tracked() {
418        let mut reporter = PredictionReporter::default();
419        for i in 0..10 {
420            reporter.observe(0.0, i as f64).unwrap();
421        }
422        let r = reporter.report(0.0).unwrap();
423        assert_eq!(r.samples_seen(), 10);
424    }
425
426    #[cfg(feature = "serde")]
427    #[test]
428    fn serde_roundtrip() {
429        let mut reporter = PredictionReporter::default();
430        reporter.observe(10.0, 12.0).unwrap();
431        reporter.observe(10.0, 9.0).unwrap();
432        reporter.set_baseline(3.0).unwrap();
433
434        let json = serde_json::to_string(&reporter).unwrap();
435        let restored: PredictionReporter = serde_json::from_str(&json).unwrap();
436
437        let r = restored.report(10.0).unwrap();
438        assert_eq!(r.samples_seen(), 2);
439        assert_eq!(r.baseline_error(), Some(3.0));
440        assert!(r.beats_baseline().is_some());
441    }
442}