use crate::error::{ForecastError, Result};
use crate::postprocess::{PointForecasts, PredictionIntervals};
#[derive(Debug, Clone, PartialEq)]
pub enum ConformalMethod {
Split {
cal_fraction: f64,
},
CrossVal {
n_folds: usize,
},
JackknifePlus,
}
impl Default for ConformalMethod {
fn default() -> Self {
Self::Split { cal_fraction: 0.2 }
}
}
#[derive(Debug, Clone)]
pub struct ConformalResult {
scores: Vec<f64>,
quantile_value: f64,
coverage: f64,
method: ConformalMethod,
}
impl ConformalResult {
pub fn scores(&self) -> &[f64] {
&self.scores
}
pub fn quantile_value(&self) -> f64 {
self.quantile_value
}
pub fn coverage(&self) -> f64 {
self.coverage
}
pub fn method(&self) -> &ConformalMethod {
&self.method
}
}
#[derive(Debug, Clone)]
pub struct ConformalPredictor {
coverage: f64,
method: ConformalMethod,
}
impl ConformalPredictor {
pub fn new(coverage: f64, method: ConformalMethod) -> Self {
assert!(
coverage > 0.0 && coverage < 1.0,
"coverage must be in (0, 1)"
);
Self { coverage, method }
}
pub fn split(coverage: f64) -> Self {
Self::new(coverage, ConformalMethod::Split { cal_fraction: 0.2 })
}
pub fn cross_val(coverage: f64, n_folds: usize) -> Self {
Self::new(coverage, ConformalMethod::CrossVal { n_folds })
}
pub fn jackknife_plus(coverage: f64) -> Self {
Self::new(coverage, ConformalMethod::JackknifePlus)
}
pub fn coverage(&self) -> f64 {
self.coverage
}
pub fn method(&self) -> &ConformalMethod {
&self.method
}
pub fn fit(&self, forecasts: &[f64], actuals: &[f64]) -> Result<ConformalResult> {
if forecasts.len() != actuals.len() {
return Err(ForecastError::DimensionMismatch {
expected: forecasts.len(),
got: actuals.len(),
});
}
let n = forecasts.len();
if n == 0 {
return Err(ForecastError::EmptyData);
}
match &self.method {
ConformalMethod::Split { cal_fraction } => {
self.fit_split(forecasts, actuals, *cal_fraction)
}
ConformalMethod::CrossVal { n_folds } => {
self.fit_cross_val(forecasts, actuals, *n_folds)
}
ConformalMethod::JackknifePlus => self.fit_jackknife_plus(forecasts, actuals),
}
}
fn fit_split(
&self,
forecasts: &[f64],
actuals: &[f64],
cal_fraction: f64,
) -> Result<ConformalResult> {
let n = forecasts.len();
let cal_size = ((n as f64) * cal_fraction).ceil() as usize;
if cal_size < 1 {
return Err(ForecastError::InsufficientData {
needed: 1,
got: cal_size,
});
}
let cal_start = n - cal_size;
let mut scores: Vec<f64> = forecasts[cal_start..]
.iter()
.zip(actuals[cal_start..].iter())
.map(|(f, a)| (f - a).abs())
.collect();
scores.sort_by(|a, b| a.partial_cmp(b).unwrap());
let adjusted_level = ((cal_size as f64 + 1.0) * self.coverage / cal_size as f64).min(1.0);
let quantile_idx = ((cal_size as f64) * adjusted_level).ceil() as usize;
let quantile_idx = quantile_idx.saturating_sub(1).min(scores.len() - 1);
let quantile_value = scores[quantile_idx];
Ok(ConformalResult {
scores,
quantile_value,
coverage: self.coverage,
method: self.method.clone(),
})
}
fn fit_cross_val(
&self,
forecasts: &[f64],
actuals: &[f64],
n_folds: usize,
) -> Result<ConformalResult> {
let n = forecasts.len();
if n_folds < 2 {
return Err(ForecastError::InvalidParameter(
"n_folds must be at least 2".to_string(),
));
}
if n < n_folds {
return Err(ForecastError::InsufficientData {
needed: n_folds,
got: n,
});
}
let mut scores: Vec<f64> = forecasts
.iter()
.zip(actuals.iter())
.map(|(f, a)| (f - a).abs())
.collect();
scores.sort_by(|a, b| a.partial_cmp(b).unwrap());
let quantile_idx = ((n as f64) * self.coverage).ceil() as usize;
let quantile_idx = quantile_idx.saturating_sub(1).min(scores.len() - 1);
let quantile_value = scores[quantile_idx];
Ok(ConformalResult {
scores,
quantile_value,
coverage: self.coverage,
method: self.method.clone(),
})
}
fn fit_jackknife_plus(&self, forecasts: &[f64], actuals: &[f64]) -> Result<ConformalResult> {
let n = forecasts.len();
if n < 2 {
return Err(ForecastError::InsufficientData { needed: 2, got: n });
}
let mut scores: Vec<f64> = forecasts
.iter()
.zip(actuals.iter())
.map(|(f, a)| (f - a).abs())
.collect();
scores.sort_by(|a, b| a.partial_cmp(b).unwrap());
let adjusted_level = (((n + 1) as f64) * self.coverage / n as f64).min(1.0);
let quantile_idx = ((n as f64) * adjusted_level).ceil() as usize;
let quantile_idx = quantile_idx.saturating_sub(1).min(scores.len() - 1);
let quantile_value = scores[quantile_idx];
Ok(ConformalResult {
scores,
quantile_value,
coverage: self.coverage,
method: self.method.clone(),
})
}
pub fn predict(
&self,
result: &ConformalResult,
point_forecasts: &PointForecasts,
) -> PredictionIntervals {
let values = point_forecasts.values();
let q = result.quantile_value;
let lower: Vec<f64> = values.iter().map(|&v| v - q).collect();
let upper: Vec<f64> = values.iter().map(|&v| v + q).collect();
PredictionIntervals::new(
point_forecasts.timestamps().to_vec(),
lower,
upper,
self.coverage,
)
.expect("Valid prediction intervals")
}
pub fn predict_values(&self, result: &ConformalResult, values: &[f64]) -> PredictionIntervals {
let q = result.quantile_value;
let lower: Vec<f64> = values.iter().map(|&v| v - q).collect();
let upper: Vec<f64> = values.iter().map(|&v| v + q).collect();
PredictionIntervals::from_bounds(lower, upper, self.coverage)
.expect("Valid prediction intervals")
}
}
#[cfg(test)]
mod tests {
use super::*;
mod conformal_method {
use super::*;
#[test]
fn default_is_split_with_20_percent() {
let method = ConformalMethod::default();
match method {
ConformalMethod::Split { cal_fraction } => {
assert!((cal_fraction - 0.2).abs() < 1e-10);
}
_ => panic!("Expected Split method"),
}
}
#[test]
fn split_stores_cal_fraction() {
let method = ConformalMethod::Split { cal_fraction: 0.3 };
if let ConformalMethod::Split { cal_fraction } = method {
assert!((cal_fraction - 0.3).abs() < 1e-10);
} else {
panic!("Expected Split method");
}
}
#[test]
fn cross_val_stores_n_folds() {
let method = ConformalMethod::CrossVal { n_folds: 5 };
if let ConformalMethod::CrossVal { n_folds } = method {
assert_eq!(n_folds, 5);
} else {
panic!("Expected CrossVal method");
}
}
#[test]
fn jackknife_plus_variant_exists() {
let method = ConformalMethod::JackknifePlus;
assert_eq!(method, ConformalMethod::JackknifePlus);
}
#[test]
fn methods_are_clonable() {
let method = ConformalMethod::Split { cal_fraction: 0.25 };
let cloned = method.clone();
assert_eq!(method, cloned);
}
}
mod construction {
use super::*;
#[test]
fn new_creates_predictor() {
let predictor =
ConformalPredictor::new(0.90, ConformalMethod::Split { cal_fraction: 0.2 });
assert!((predictor.coverage() - 0.90).abs() < 1e-10);
}
#[test]
fn split_creates_split_predictor() {
let predictor = ConformalPredictor::split(0.95);
assert!((predictor.coverage() - 0.95).abs() < 1e-10);
match predictor.method() {
ConformalMethod::Split { cal_fraction } => {
assert!((cal_fraction - 0.2).abs() < 1e-10);
}
_ => panic!("Expected Split method"),
}
}
#[test]
fn cross_val_creates_cv_predictor() {
let predictor = ConformalPredictor::cross_val(0.90, 5);
assert!((predictor.coverage() - 0.90).abs() < 1e-10);
match predictor.method() {
ConformalMethod::CrossVal { n_folds } => {
assert_eq!(*n_folds, 5);
}
_ => panic!("Expected CrossVal method"),
}
}
#[test]
fn jackknife_plus_creates_jackknife_predictor() {
let predictor = ConformalPredictor::jackknife_plus(0.90);
assert!((predictor.coverage() - 0.90).abs() < 1e-10);
assert_eq!(predictor.method(), &ConformalMethod::JackknifePlus);
}
#[test]
#[should_panic(expected = "coverage must be in (0, 1)")]
fn new_panics_on_zero_coverage() {
ConformalPredictor::new(0.0, ConformalMethod::default());
}
#[test]
#[should_panic(expected = "coverage must be in (0, 1)")]
fn new_panics_on_one_coverage() {
ConformalPredictor::new(1.0, ConformalMethod::default());
}
#[test]
#[should_panic(expected = "coverage must be in (0, 1)")]
fn new_panics_on_negative_coverage() {
ConformalPredictor::new(-0.1, ConformalMethod::default());
}
#[test]
fn predictor_is_clonable() {
let predictor = ConformalPredictor::split(0.90);
let cloned = predictor.clone();
assert!((cloned.coverage() - 0.90).abs() < 1e-10);
}
}
mod fit_split {
use super::*;
#[test]
fn fit_returns_result() {
let predictor = ConformalPredictor::split(0.90);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0];
let actuals = vec![10.5, 10.5, 12.5, 12.5, 14.5];
let result = predictor.fit(&forecasts, &actuals).unwrap();
assert!((result.coverage() - 0.90).abs() < 1e-10);
assert!(!result.scores().is_empty());
}
#[test]
fn fit_fails_on_length_mismatch() {
let predictor = ConformalPredictor::split(0.90);
let forecasts = vec![10.0, 11.0, 12.0];
let actuals = vec![10.5, 10.5];
let result = predictor.fit(&forecasts, &actuals);
assert!(result.is_err());
}
#[test]
fn fit_fails_on_empty_data() {
let predictor = ConformalPredictor::split(0.90);
let forecasts: Vec<f64> = vec![];
let actuals: Vec<f64> = vec![];
let result = predictor.fit(&forecasts, &actuals);
assert!(result.is_err());
}
#[test]
fn scores_are_sorted() {
let predictor = ConformalPredictor::split(0.90);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0, 17.0, 18.0, 19.0];
let actuals = vec![10.5, 10.0, 12.5, 13.5, 14.0, 15.5, 16.0, 17.5, 18.0, 19.5];
let result = predictor.fit(&forecasts, &actuals).unwrap();
let scores = result.scores();
for i in 1..scores.len() {
assert!(scores[i] >= scores[i - 1], "Scores should be sorted");
}
}
#[test]
fn quantile_value_is_positive() {
let predictor = ConformalPredictor::split(0.90);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0];
let actuals = vec![10.5, 10.5, 12.5, 12.5, 14.5];
let result = predictor.fit(&forecasts, &actuals).unwrap();
assert!(result.quantile_value() >= 0.0);
}
#[test]
fn higher_coverage_gives_larger_quantile() {
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0, 17.0, 18.0, 19.0];
let actuals = vec![9.0, 12.0, 11.0, 14.0, 13.0, 16.0, 15.0, 18.0, 17.0, 20.0];
let predictor_90 = ConformalPredictor::split(0.50);
let predictor_95 = ConformalPredictor::split(0.90);
let result_90 = predictor_90.fit(&forecasts, &actuals).unwrap();
let result_95 = predictor_95.fit(&forecasts, &actuals).unwrap();
assert!(result_95.quantile_value() >= result_90.quantile_value());
}
}
mod fit_cross_val {
use super::*;
#[test]
fn fit_returns_result() {
let predictor = ConformalPredictor::cross_val(0.90, 5);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0];
let actuals = vec![10.5, 10.5, 12.5, 12.5, 14.5];
let result = predictor.fit(&forecasts, &actuals).unwrap();
assert!((result.coverage() - 0.90).abs() < 1e-10);
}
#[test]
fn fit_fails_on_insufficient_folds() {
let predictor = ConformalPredictor::cross_val(0.90, 1);
let forecasts = vec![10.0, 11.0, 12.0];
let actuals = vec![10.5, 10.5, 12.5];
let result = predictor.fit(&forecasts, &actuals);
assert!(result.is_err());
}
#[test]
fn fit_fails_when_n_less_than_folds() {
let predictor = ConformalPredictor::cross_val(0.90, 10);
let forecasts = vec![10.0, 11.0, 12.0];
let actuals = vec![10.5, 10.5, 12.5];
let result = predictor.fit(&forecasts, &actuals);
assert!(result.is_err());
}
#[test]
fn uses_all_data_for_scores() {
let predictor = ConformalPredictor::cross_val(0.90, 5);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0, 17.0, 18.0, 19.0];
let actuals = vec![10.5, 10.5, 12.5, 12.5, 14.5, 15.5, 15.5, 17.5, 18.5, 19.5];
let result = predictor.fit(&forecasts, &actuals).unwrap();
assert_eq!(result.scores().len(), 10);
}
}
mod fit_jackknife_plus {
use super::*;
#[test]
fn fit_returns_result() {
let predictor = ConformalPredictor::jackknife_plus(0.90);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0];
let actuals = vec![10.5, 10.5, 12.5, 12.5, 14.5];
let result = predictor.fit(&forecasts, &actuals).unwrap();
assert!((result.coverage() - 0.90).abs() < 1e-10);
}
#[test]
fn fit_fails_on_single_point() {
let predictor = ConformalPredictor::jackknife_plus(0.90);
let forecasts = vec![10.0];
let actuals = vec![10.5];
let result = predictor.fit(&forecasts, &actuals);
assert!(result.is_err());
}
#[test]
fn works_with_two_points() {
let predictor = ConformalPredictor::jackknife_plus(0.90);
let forecasts = vec![10.0, 11.0];
let actuals = vec![10.5, 10.5];
let result = predictor.fit(&forecasts, &actuals);
assert!(result.is_ok());
}
#[test]
fn uses_all_data_for_scores() {
let predictor = ConformalPredictor::jackknife_plus(0.90);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0];
let actuals = vec![10.5, 10.5, 12.5, 12.5, 14.5];
let result = predictor.fit(&forecasts, &actuals).unwrap();
assert_eq!(result.scores().len(), 5);
}
}
mod predict {
use super::*;
use chrono::{TimeZone, Utc};
fn make_timestamps(n: usize) -> Vec<chrono::DateTime<Utc>> {
(0..n)
.map(|i| {
Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap()
+ chrono::Duration::days(i as i64)
})
.collect()
}
#[test]
fn predict_returns_intervals() {
let predictor = ConformalPredictor::split(0.90);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0, 17.0, 18.0, 19.0];
let actuals = vec![10.5, 10.5, 12.5, 12.5, 14.5, 15.5, 15.5, 17.5, 18.5, 19.5];
let result = predictor.fit(&forecasts, &actuals).unwrap();
let new_forecasts = PointForecasts::from_values(vec![20.0, 21.0, 22.0]);
let intervals = predictor.predict(&result, &new_forecasts);
assert_eq!(intervals.len(), 3);
assert!((intervals.coverage() - 0.90).abs() < 1e-10);
}
#[test]
fn predict_with_timestamps() {
let predictor = ConformalPredictor::split(0.90);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0, 17.0, 18.0, 19.0];
let actuals = vec![10.5, 10.5, 12.5, 12.5, 14.5, 15.5, 15.5, 17.5, 18.5, 19.5];
let result = predictor.fit(&forecasts, &actuals).unwrap();
let timestamps = make_timestamps(3);
let new_forecasts =
PointForecasts::new(timestamps.clone(), vec![20.0, 21.0, 22.0]).unwrap();
let intervals = predictor.predict(&result, &new_forecasts);
assert!(intervals.has_timestamps());
assert_eq!(intervals.timestamps(), ×tamps);
}
#[test]
fn intervals_are_symmetric() {
let predictor = ConformalPredictor::split(0.90);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0, 17.0, 18.0, 19.0];
let actuals = vec![10.5, 10.5, 12.5, 12.5, 14.5, 15.5, 15.5, 17.5, 18.5, 19.5];
let result = predictor.fit(&forecasts, &actuals).unwrap();
let new_forecasts = PointForecasts::from_values(vec![20.0]);
let intervals = predictor.predict(&result, &new_forecasts);
let point = 20.0;
let lower = intervals.lower()[0];
let upper = intervals.upper()[0];
let lower_diff = point - lower;
let upper_diff = upper - point;
assert!(
(lower_diff - upper_diff).abs() < 1e-10,
"Intervals should be symmetric"
);
}
#[test]
fn predict_values_works() {
let predictor = ConformalPredictor::split(0.90);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0, 17.0, 18.0, 19.0];
let actuals = vec![10.5, 10.5, 12.5, 12.5, 14.5, 15.5, 15.5, 17.5, 18.5, 19.5];
let result = predictor.fit(&forecasts, &actuals).unwrap();
let intervals = predictor.predict_values(&result, &[20.0, 21.0]);
assert_eq!(intervals.len(), 2);
assert!(!intervals.has_timestamps());
}
#[test]
fn larger_errors_give_wider_intervals() {
let forecasts_small = vec![10.0, 11.0, 12.0, 13.0, 14.0];
let actuals_small = vec![10.1, 11.1, 12.1, 13.1, 14.1];
let forecasts_large = vec![10.0, 11.0, 12.0, 13.0, 14.0];
let actuals_large = vec![8.0, 13.0, 10.0, 15.0, 12.0];
let predictor = ConformalPredictor::split(0.90);
let result_small = predictor.fit(&forecasts_small, &actuals_small).unwrap();
let result_large = predictor.fit(&forecasts_large, &actuals_large).unwrap();
assert!(result_large.quantile_value() > result_small.quantile_value());
}
}
mod coverage_validation {
use super::*;
#[test]
fn empirical_coverage_approximately_matches_target() {
let n = 100;
let forecasts: Vec<f64> = (0..n).map(|i| i as f64).collect();
let errors: Vec<f64> = (0..n)
.map(|i| ((i * 7 + 3) % 21) as f64 / 10.0 - 1.0)
.collect();
let actuals: Vec<f64> = forecasts
.iter()
.zip(errors.iter())
.map(|(f, e)| f + e)
.collect();
let predictor = ConformalPredictor::split(0.90);
let result = predictor.fit(&forecasts, &actuals).unwrap();
let new_forecasts: Vec<f64> = (100..150).map(|i| i as f64).collect();
let new_errors: Vec<f64> = (100..150)
.map(|i| ((i * 7 + 3) % 21) as f64 / 10.0 - 1.0)
.collect();
let new_actuals: Vec<f64> = new_forecasts
.iter()
.zip(new_errors.iter())
.map(|(f, e)| f + e)
.collect();
let intervals = predictor.predict_values(&result, &new_forecasts);
let empirical = intervals.empirical_coverage(&new_actuals).unwrap();
assert!(
empirical >= 0.70,
"Empirical coverage {} should be reasonably high",
empirical
);
}
}
mod conformal_result {
use super::*;
#[test]
fn accessors_return_correct_values() {
let predictor = ConformalPredictor::split(0.90);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0];
let actuals = vec![10.5, 10.5, 12.5, 12.5, 14.5];
let result = predictor.fit(&forecasts, &actuals).unwrap();
assert!(!result.scores().is_empty());
assert!(result.quantile_value() >= 0.0);
assert!((result.coverage() - 0.90).abs() < 1e-10);
match result.method() {
ConformalMethod::Split { .. } => {}
_ => panic!("Expected Split method"),
}
}
#[test]
fn result_is_clonable() {
let predictor = ConformalPredictor::split(0.90);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0];
let actuals = vec![10.5, 10.5, 12.5, 12.5, 14.5];
let result = predictor.fit(&forecasts, &actuals).unwrap();
let cloned = result.clone();
assert_eq!(result.scores(), cloned.scores());
assert!((result.quantile_value() - cloned.quantile_value()).abs() < 1e-10);
}
}
}