use anofox_regression::solvers::{FittedIsotonic, IsotonicRegressor};
use faer::Col;
use crate::error::{ForecastError, Result};
use crate::postprocess::{PointForecasts, QuantileForecasts};
#[derive(Debug, Clone)]
pub struct IDRResult {
fitted_models: Vec<FittedIsotonic>,
quantiles: Vec<f64>,
x_grid: Vec<f64>,
y_grid: Vec<Vec<f64>>,
}
impl IDRResult {
pub fn quantiles(&self) -> &[f64] {
&self.quantiles
}
pub fn n_models(&self) -> usize {
self.fitted_models.len()
}
pub fn x_grid(&self) -> &[f64] {
&self.x_grid
}
pub fn y_grid(&self) -> &[Vec<f64>] {
&self.y_grid
}
}
#[derive(Debug, Clone)]
pub struct IDRPredictor {
quantiles: Vec<f64>,
}
impl IDRPredictor {
pub fn new(quantiles: Vec<f64>) -> Self {
for &q in &quantiles {
assert!(q > 0.0 && q < 1.0, "quantiles must be in (0, 1)");
}
for w in quantiles.windows(2) {
assert!(w[0] < w[1], "quantiles must be sorted in ascending order");
}
Self { quantiles }
}
pub fn quantiles(&self) -> &[f64] {
&self.quantiles
}
pub fn fit(&self, forecasts: &[f64], actuals: &[f64]) -> Result<IDRResult> {
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);
}
if n < 2 {
return Err(ForecastError::InsufficientData {
needed: 2,
got: n,
hint: None,
});
}
let mut x_grid: Vec<f64> = forecasts.to_vec();
x_grid.sort_by(|a, b| a.partial_cmp(b).unwrap());
x_grid.dedup();
let x_col = Col::from_fn(n, |i| forecasts[i]);
let y_col = Col::from_fn(n, |i| actuals[i]);
let iso = IsotonicRegressor::new();
let fitted = iso.fit_1d(&x_col, &y_col).map_err(|e| {
ForecastError::ConvergenceFailure(format!("Isotonic regression failed: {:?}", e))
})?;
let fitted_values = fitted.fitted_values();
let residuals: Vec<f64> = (0..n).map(|i| actuals[i] - fitted_values[i]).collect();
let mut sorted_residuals = residuals.clone();
sorted_residuals.sort_by(|a, b| a.partial_cmp(b).unwrap());
let residual_quantiles: Vec<f64> = self
.quantiles
.iter()
.map(|&q| {
let idx = ((n as f64) * q).floor() as usize;
let idx = idx.min(n - 1);
sorted_residuals[idx]
})
.collect();
let x_grid_col = Col::from_fn(x_grid.len(), |i| x_grid[i]);
let grid_predictions = fitted.predict_1d(&x_grid_col);
let y_grid: Vec<Vec<f64>> = (0..x_grid.len())
.map(|i| {
let base = grid_predictions[i];
residual_quantiles.iter().map(|&q| base + q).collect()
})
.collect();
Ok(IDRResult {
fitted_models: vec![fitted],
quantiles: self.quantiles.clone(),
x_grid,
y_grid,
})
}
pub fn predict(
&self,
result: &IDRResult,
point_forecasts: &PointForecasts,
) -> Result<QuantileForecasts> {
let values = point_forecasts.values();
if result.fitted_models.is_empty() {
return Err(ForecastError::FitRequired { model: None });
}
let fitted = &result.fitted_models[0];
let x_col = Col::from_fn(values.len(), |i| values[i]);
let base_preds = fitted.predict_1d(&x_col);
let residual_quantiles: Vec<f64> = if !result.y_grid.is_empty() && !result.x_grid.is_empty()
{
let first_base = result.y_grid[0]
.iter()
.zip(result.quantiles.iter())
.map(|(&y, _)| y)
.collect::<Vec<_>>();
let base_at_first = if !result.x_grid.is_empty() {
fitted.predict_single(result.x_grid[0])
} else {
0.0
};
first_base.iter().map(|&y| y - base_at_first).collect()
} else {
vec![0.0; self.quantiles.len()]
};
let forecast_values: Vec<Vec<f64>> = (0..values.len())
.map(|i| {
let base = base_preds[i];
residual_quantiles.iter().map(|&q| base + q).collect()
})
.collect();
QuantileForecasts::new(
point_forecasts.timestamps().to_vec(),
self.quantiles.clone(),
forecast_values,
)
}
pub fn predict_values(&self, result: &IDRResult, values: &[f64]) -> Result<QuantileForecasts> {
let forecasts = PointForecasts::from_values(values.to_vec());
self.predict(result, &forecasts)
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_relative_eq;
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()
}
mod construction {
use super::*;
#[test]
fn new_creates_predictor() {
let pred = IDRPredictor::new(vec![0.1, 0.5, 0.9]);
assert_eq!(pred.quantiles(), &[0.1, 0.5, 0.9]);
}
#[test]
fn new_single_quantile() {
let pred = IDRPredictor::new(vec![0.5]);
assert_eq!(pred.quantiles(), &[0.5]);
}
#[test]
fn new_many_quantiles() {
let quantiles: Vec<f64> = (1..20).map(|i| i as f64 / 20.0).collect();
let pred = IDRPredictor::new(quantiles.clone());
assert_eq!(pred.quantiles(), &quantiles);
}
#[test]
#[should_panic(expected = "quantiles must be in (0, 1)")]
fn new_panics_on_zero_quantile() {
IDRPredictor::new(vec![0.0, 0.5, 0.9]);
}
#[test]
#[should_panic(expected = "quantiles must be in (0, 1)")]
fn new_panics_on_one_quantile() {
IDRPredictor::new(vec![0.1, 0.5, 1.0]);
}
#[test]
#[should_panic(expected = "quantiles must be in (0, 1)")]
fn new_panics_on_negative_quantile() {
IDRPredictor::new(vec![-0.1, 0.5, 0.9]);
}
#[test]
#[should_panic(expected = "quantiles must be in (0, 1)")]
fn new_panics_on_quantile_greater_than_one() {
IDRPredictor::new(vec![0.1, 0.5, 1.5]);
}
#[test]
#[should_panic(expected = "quantiles must be sorted")]
fn new_panics_on_unsorted_quantiles() {
IDRPredictor::new(vec![0.9, 0.5, 0.1]);
}
#[test]
#[should_panic(expected = "quantiles must be sorted")]
fn new_panics_on_duplicate_quantiles() {
IDRPredictor::new(vec![0.5, 0.5, 0.9]);
}
#[test]
fn predictor_is_clonable() {
let pred = IDRPredictor::new(vec![0.1, 0.5, 0.9]);
let cloned = pred.clone();
assert_eq!(pred.quantiles(), cloned.quantiles());
}
#[test]
fn predictor_is_debuggable() {
let pred = IDRPredictor::new(vec![0.1, 0.5, 0.9]);
let debug_str = format!("{:?}", pred);
assert!(debug_str.contains("IDRPredictor"));
}
}
mod fit {
use super::*;
#[test]
fn fit_returns_result() {
let pred = IDRPredictor::new(vec![0.1, 0.5, 0.9]);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0];
let actuals = vec![10.5, 11.5, 12.5, 13.5, 14.5];
let result = pred.fit(&forecasts, &actuals).unwrap();
assert_eq!(result.quantiles(), &[0.1, 0.5, 0.9]);
assert_eq!(result.n_models(), 1);
}
#[test]
fn fit_with_two_data_points() {
let pred = IDRPredictor::new(vec![0.5]);
let forecasts = vec![10.0, 12.0];
let actuals = vec![11.0, 13.0];
let result = pred.fit(&forecasts, &actuals);
assert!(result.is_ok());
}
#[test]
fn fit_fails_on_length_mismatch() {
let pred = IDRPredictor::new(vec![0.1, 0.5, 0.9]);
let forecasts = vec![10.0, 11.0, 12.0];
let actuals = vec![10.5, 11.5];
let err = pred.fit(&forecasts, &actuals).unwrap_err();
assert!(
matches!(err, ForecastError::DimensionMismatch { .. }),
"Expected DimensionMismatch, got {:?}",
err
);
}
#[test]
fn fit_fails_on_empty_data() {
let pred = IDRPredictor::new(vec![0.1, 0.5, 0.9]);
let forecasts: Vec<f64> = vec![];
let actuals: Vec<f64> = vec![];
let err = pred.fit(&forecasts, &actuals).unwrap_err();
assert!(
matches!(err, ForecastError::EmptyData),
"Expected EmptyData, got {:?}",
err
);
}
#[test]
fn fit_fails_on_single_point() {
let pred = IDRPredictor::new(vec![0.5]);
let forecasts = vec![10.0];
let actuals = vec![10.5];
let err = pred.fit(&forecasts, &actuals).unwrap_err();
assert!(
matches!(
err,
ForecastError::InsufficientData {
needed: 2,
got: 1,
..
}
),
"Expected InsufficientData, got {:?}",
err
);
}
#[test]
fn x_grid_is_sorted() {
let pred = IDRPredictor::new(vec![0.5]);
let forecasts = vec![14.0, 10.0, 12.0, 11.0, 13.0];
let actuals = vec![14.5, 10.5, 12.5, 11.5, 13.5];
let result = pred.fit(&forecasts, &actuals).unwrap();
let grid = result.x_grid();
for i in 1..grid.len() {
assert!(grid[i] > grid[i - 1], "X grid should be sorted");
}
}
#[test]
fn x_grid_is_unique() {
let pred = IDRPredictor::new(vec![0.5]);
let forecasts = vec![10.0, 10.0, 12.0, 12.0, 14.0];
let actuals = vec![10.5, 10.3, 12.5, 12.7, 14.5];
let result = pred.fit(&forecasts, &actuals).unwrap();
let grid = result.x_grid();
assert_eq!(grid.len(), 3);
}
#[test]
fn x_grid_contains_all_unique_forecasts() {
let pred = IDRPredictor::new(vec![0.5]);
let forecasts = vec![10.0, 20.0, 30.0, 40.0, 50.0];
let actuals = vec![10.5, 20.5, 30.5, 40.5, 50.5];
let result = pred.fit(&forecasts, &actuals).unwrap();
let grid = result.x_grid();
assert_eq!(grid.len(), 5);
assert_relative_eq!(grid[0], 10.0, epsilon = 1e-10);
assert_relative_eq!(grid[4], 50.0, epsilon = 1e-10);
}
#[test]
fn y_grid_has_correct_dimensions() {
let pred = IDRPredictor::new(vec![0.1, 0.5, 0.9]);
let forecasts = vec![10.0, 20.0, 30.0, 40.0, 50.0];
let actuals = vec![10.5, 20.5, 30.5, 40.5, 50.5];
let result = pred.fit(&forecasts, &actuals).unwrap();
assert_eq!(result.y_grid().len(), result.x_grid().len());
for row in result.y_grid() {
assert_eq!(row.len(), 3);
}
}
#[test]
fn fit_with_negative_values() {
let pred = IDRPredictor::new(vec![0.5]);
let forecasts = vec![-5.0, -3.0, -1.0, 1.0, 3.0];
let actuals = vec![-4.5, -2.5, -0.5, 1.5, 3.5];
let result = pred.fit(&forecasts, &actuals);
assert!(result.is_ok());
}
#[test]
fn fit_with_constant_forecasts() {
let pred = IDRPredictor::new(vec![0.5]);
let forecasts = vec![10.0, 10.0, 10.0, 10.0, 10.0];
let actuals = vec![9.0, 10.0, 11.0, 12.0, 8.0];
let result = pred.fit(&forecasts, &actuals);
assert!(result.is_ok());
assert_eq!(result.unwrap().x_grid().len(), 1);
}
}
mod predict {
use super::*;
#[test]
fn predict_returns_quantile_forecasts() {
let pred = IDRPredictor::new(vec![0.1, 0.5, 0.9]);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0];
let actuals = vec![10.5, 11.5, 12.5, 13.5, 14.5];
let result = pred.fit(&forecasts, &actuals).unwrap();
let new_forecasts = PointForecasts::from_values(vec![20.0, 21.0]);
let quantiles = pred.predict(&result, &new_forecasts).unwrap();
assert_eq!(quantiles.n_times(), 2);
assert_eq!(quantiles.n_quantiles(), 3);
}
#[test]
fn predict_with_timestamps() {
let pred = IDRPredictor::new(vec![0.1, 0.5, 0.9]);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0];
let actuals = vec![10.5, 11.5, 12.5, 13.5, 14.5];
let result = pred.fit(&forecasts, &actuals).unwrap();
let timestamps = make_timestamps(2);
let new_forecasts = PointForecasts::new(timestamps.clone(), vec![20.0, 21.0]).unwrap();
let quantiles = pred.predict(&result, &new_forecasts).unwrap();
assert!(quantiles.has_timestamps());
assert_eq!(quantiles.timestamps(), ×tamps);
}
#[test]
fn predict_values_works() {
let pred = IDRPredictor::new(vec![0.1, 0.5, 0.9]);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0];
let actuals = vec![10.5, 11.5, 12.5, 13.5, 14.5];
let result = pred.fit(&forecasts, &actuals).unwrap();
let quantiles = pred.predict_values(&result, &[20.0, 21.0]).unwrap();
assert_eq!(quantiles.n_times(), 2);
assert!(!quantiles.has_timestamps());
}
#[test]
fn predict_single_value() {
let pred = IDRPredictor::new(vec![0.5]);
let forecasts: Vec<f64> = (0..10).map(|i| i as f64).collect();
let actuals: Vec<f64> = forecasts.iter().map(|&f| f + 1.0).collect();
let result = pred.fit(&forecasts, &actuals).unwrap();
let quantiles = pred.predict_values(&result, &[5.0]).unwrap();
assert_eq!(quantiles.n_times(), 1);
assert_eq!(quantiles.n_quantiles(), 1);
}
#[test]
fn quantile_values_are_monotonic() {
let pred = IDRPredictor::new(vec![0.1, 0.5, 0.9]);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0];
let actuals = vec![10.5, 10.0, 12.5, 13.5, 14.0];
let result = pred.fit(&forecasts, &actuals).unwrap();
let new_forecasts = PointForecasts::from_values(vec![20.0]);
let quantiles = pred.predict(&result, &new_forecasts).unwrap();
let row = quantiles.at_time(0).unwrap();
assert!(
row[0] <= row[1] && row[1] <= row[2],
"Quantiles should be monotonic: {:?}",
row
);
}
#[test]
fn higher_forecast_gives_higher_quantiles() {
let pred = IDRPredictor::new(vec![0.5]);
let forecasts: Vec<f64> = (0..20).map(|i| i as f64).collect();
let actuals: Vec<f64> = forecasts.iter().map(|&f| f + 1.0).collect();
let result = pred.fit(&forecasts, &actuals).unwrap();
let q_low = pred.predict_values(&result, &[5.0]).unwrap();
let q_high = pred.predict_values(&result, &[15.0]).unwrap();
assert!(
q_high.at_time(0).unwrap()[0] > q_low.at_time(0).unwrap()[0],
"Higher forecast should give higher quantile prediction"
);
}
#[test]
fn predict_at_training_points() {
let pred = IDRPredictor::new(vec![0.5]);
let forecasts: Vec<f64> = (0..20).map(|i| i as f64).collect();
let actuals: Vec<f64> = forecasts.iter().map(|&f| f + 1.0).collect();
let result = pred.fit(&forecasts, &actuals).unwrap();
let quantiles = pred.predict_values(&result, &forecasts).unwrap();
assert_eq!(quantiles.n_times(), 20);
}
#[test]
fn predict_extrapolation() {
let pred = IDRPredictor::new(vec![0.5]);
let forecasts: Vec<f64> = (0..10).map(|i| i as f64).collect();
let actuals: Vec<f64> = forecasts.iter().map(|&f| f + 1.0).collect();
let result = pred.fit(&forecasts, &actuals).unwrap();
let quantiles = pred.predict_values(&result, &[100.0]).unwrap();
assert_eq!(quantiles.n_times(), 1);
assert!(quantiles.at_time(0).unwrap()[0].is_finite());
}
#[test]
fn predict_returns_finite_values() {
let pred = IDRPredictor::new(vec![0.1, 0.5, 0.9]);
let forecasts: Vec<f64> = (0..20).map(|i| i as f64 * 10.0).collect();
let actuals: Vec<f64> = forecasts
.iter()
.enumerate()
.map(|(i, &f)| f + (i as f64).sin() * 5.0)
.collect();
let result = pred.fit(&forecasts, &actuals).unwrap();
let quantiles = pred.predict_values(&result, &[50.0, 100.0, 150.0]).unwrap();
for t in 0..3 {
let row = quantiles.at_time(t).unwrap();
for &val in row {
assert!(val.is_finite(), "All predicted quantiles must be finite");
}
}
}
}
mod idr_result {
use super::*;
#[test]
fn accessors_return_correct_values() {
let pred = IDRPredictor::new(vec![0.1, 0.5, 0.9]);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0];
let actuals = vec![10.5, 11.5, 12.5, 13.5, 14.5];
let result = pred.fit(&forecasts, &actuals).unwrap();
assert_eq!(result.quantiles(), &[0.1, 0.5, 0.9]);
assert_eq!(result.n_models(), 1);
assert!(!result.x_grid().is_empty());
assert!(!result.y_grid().is_empty());
}
#[test]
fn result_is_clonable() {
let pred = IDRPredictor::new(vec![0.1, 0.5, 0.9]);
let forecasts = vec![10.0, 11.0, 12.0, 13.0, 14.0];
let actuals = vec![10.5, 11.5, 12.5, 13.5, 14.5];
let result = pred.fit(&forecasts, &actuals).unwrap();
let cloned = result.clone();
assert_eq!(result.quantiles(), cloned.quantiles());
assert_eq!(result.x_grid(), cloned.x_grid());
}
#[test]
fn result_is_debuggable() {
let pred = IDRPredictor::new(vec![0.5]);
let forecasts = vec![10.0, 12.0, 14.0];
let actuals = vec![10.5, 12.5, 14.5];
let result = pred.fit(&forecasts, &actuals).unwrap();
let debug_str = format!("{:?}", result);
assert!(debug_str.contains("IDRResult"));
}
#[test]
fn y_grid_quantiles_are_monotonic() {
let pred = IDRPredictor::new(vec![0.1, 0.5, 0.9]);
let forecasts: Vec<f64> = (0..20).map(|i| i as f64).collect();
let actuals: Vec<f64> = forecasts
.iter()
.enumerate()
.map(|(i, &f)| f + ((i % 5) as f64 - 2.0))
.collect();
let result = pred.fit(&forecasts, &actuals).unwrap();
for row in result.y_grid() {
assert!(
row[0] <= row[1] && row[1] <= row[2],
"y_grid quantiles should be monotonic: {:?}",
row
);
}
}
}
mod calibration {
use super::*;
#[test]
fn calibration_on_linear_data() {
let pred = IDRPredictor::new(vec![0.1, 0.5, 0.9]);
let n = 100;
let forecasts: Vec<f64> = (0..n).map(|i| i as f64).collect();
let actuals: Vec<f64> = forecasts
.iter()
.enumerate()
.map(|(i, &f)| f + ((i % 5) as f64 - 2.0))
.collect();
let result = pred.fit(&forecasts, &actuals).unwrap();
let quantiles = pred.predict_values(&result, &forecasts).unwrap();
assert_eq!(quantiles.n_times(), n);
assert_eq!(quantiles.n_quantiles(), 3);
for t in 0..n {
let row = quantiles.at_time(t).unwrap();
assert!(row[0] <= row[1], "q0.1 <= q0.5 at t={}", t);
assert!(row[1] <= row[2], "q0.5 <= q0.9 at t={}", t);
}
}
#[test]
fn constant_bias_captured() {
let pred = IDRPredictor::new(vec![0.5]);
let forecasts: Vec<f64> = (0..50).map(|i| i as f64).collect();
let actuals: Vec<f64> = forecasts.iter().map(|&f| f + 1.0).collect();
let result = pred.fit(&forecasts, &actuals).unwrap();
let quantiles = pred.predict_values(&result, &[25.0]).unwrap();
let median = quantiles.at_time(0).unwrap()[0];
assert_relative_eq!(median, 26.0, epsilon = 1.0);
}
#[test]
fn larger_dataset_produces_more_grid_points() {
let pred = IDRPredictor::new(vec![0.5]);
let small_forecasts: Vec<f64> = (0..5).map(|i| i as f64 * 10.0).collect();
let small_actuals: Vec<f64> = small_forecasts.iter().map(|&f| f + 1.0).collect();
let small_result = pred.fit(&small_forecasts, &small_actuals).unwrap();
let large_forecasts: Vec<f64> = (0..50).map(|i| i as f64).collect();
let large_actuals: Vec<f64> = large_forecasts.iter().map(|&f| f + 1.0).collect();
let large_result = pred.fit(&large_forecasts, &large_actuals).unwrap();
assert!(
large_result.x_grid().len() > small_result.x_grid().len(),
"More unique training points should yield a larger grid"
);
}
}
mod edge_cases {
use super::*;
#[test]
fn fit_with_identical_values() {
let pred = IDRPredictor::new(vec![0.5]);
let forecasts = vec![5.0, 5.0, 5.0, 5.0, 5.0];
let actuals = vec![5.0, 5.0, 5.0, 5.0, 5.0];
let result = pred.fit(&forecasts, &actuals).unwrap();
let quantiles = pred.predict_values(&result, &[5.0]).unwrap();
let val = quantiles.at_time(0).unwrap()[0];
assert_relative_eq!(val, 5.0, epsilon = 1e-6);
}
#[test]
fn fit_with_large_values() {
let pred = IDRPredictor::new(vec![0.5]);
let forecasts = vec![1e10, 2e10, 3e10, 4e10, 5e10];
let actuals = vec![1.1e10, 2.1e10, 3.1e10, 4.1e10, 5.1e10];
let result = pred.fit(&forecasts, &actuals);
assert!(result.is_ok());
}
#[test]
fn fit_with_very_small_values() {
let pred = IDRPredictor::new(vec![0.5]);
let forecasts = vec![1e-10, 2e-10, 3e-10, 4e-10, 5e-10];
let actuals = vec![1.1e-10, 2.1e-10, 3.1e-10, 4.1e-10, 5.1e-10];
let result = pred.fit(&forecasts, &actuals);
assert!(result.is_ok());
}
#[test]
fn predict_with_many_quantiles() {
let quantiles: Vec<f64> = (1..10).map(|i| i as f64 / 10.0).collect();
let pred = IDRPredictor::new(quantiles);
let forecasts: Vec<f64> = (0..20).map(|i| i as f64).collect();
let actuals: Vec<f64> = forecasts
.iter()
.enumerate()
.map(|(i, &f)| f + (i as f64 * 0.3).sin())
.collect();
let result = pred.fit(&forecasts, &actuals).unwrap();
let quantile_forecasts = pred.predict_values(&result, &[10.0]).unwrap();
assert_eq!(quantile_forecasts.n_quantiles(), 9);
let row = quantile_forecasts.at_time(0).unwrap();
for i in 1..row.len() {
assert!(
row[i] >= row[i - 1],
"Quantiles should be monotonic: q[{}]={} < q[{}]={}",
i - 1,
row[i - 1],
i,
row[i]
);
}
}
}
}