use crate::error::{ForecastError, Result};
use crate::postprocess::PredictionIntervals;
#[derive(Debug, Clone)]
pub struct CqrResult {
scores: Vec<f64>,
quantile_value: f64,
coverage: f64,
tau_lo: f64,
tau_hi: f64,
}
impl CqrResult {
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 tau_lo(&self) -> f64 {
self.tau_lo
}
pub fn tau_hi(&self) -> f64 {
self.tau_hi
}
pub fn predict_intervals(&self, q_lo: &[f64], q_hi: &[f64]) -> Result<PredictionIntervals> {
if q_lo.len() != q_hi.len() {
return Err(ForecastError::DimensionMismatch {
expected: q_lo.len(),
got: q_hi.len(),
});
}
let lower: Vec<f64> = q_lo.iter().map(|&q| q - self.quantile_value).collect();
let upper: Vec<f64> = q_hi.iter().map(|&q| q + self.quantile_value).collect();
PredictionIntervals::from_bounds(lower, upper, self.coverage)
}
}
#[derive(Debug, Clone)]
pub struct CqrPredictor {
coverage: f64,
}
impl CqrPredictor {
pub fn new(coverage: f64) -> Self {
assert!(
coverage > 0.0 && coverage < 1.0,
"coverage must be in (0, 1)"
);
Self { coverage }
}
pub fn coverage(&self) -> f64 {
self.coverage
}
pub fn tau_lo(&self) -> f64 {
(1.0 - self.coverage) / 2.0
}
pub fn tau_hi(&self) -> f64 {
1.0 - (1.0 - self.coverage) / 2.0
}
pub fn fit(
&self,
q_lo_calib: &[f64],
q_hi_calib: &[f64],
actuals: &[f64],
) -> Result<CqrResult> {
let n = actuals.len();
if n == 0 {
return Err(ForecastError::EmptyData);
}
if q_lo_calib.len() != n || q_hi_calib.len() != n {
return Err(ForecastError::DimensionMismatch {
expected: n,
got: q_lo_calib.len().min(q_hi_calib.len()),
});
}
let mut scores: Vec<f64> = (0..n)
.map(|i| (q_lo_calib[i] - actuals[i]).max(actuals[i] - q_hi_calib[i]))
.collect();
scores.sort_by(|a, b| a.partial_cmp(b).unwrap());
let level = self.coverage * (n as f64 + 1.0) / n as f64;
let level = level.min(1.0);
let idx = ((n as f64) * level).ceil() as usize;
let idx = idx.saturating_sub(1).min(n - 1);
let quantile_value = scores[idx];
Ok(CqrResult {
scores,
quantile_value,
coverage: self.coverage,
tau_lo: self.tau_lo(),
tau_hi: self.tau_hi(),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn calibrated_base_gives_small_adjustment() {
let n = 200;
let actuals: Vec<f64> = (0..n).map(|i| (i as f64 * 0.1).sin()).collect();
let q_lo: Vec<f64> = actuals.iter().map(|&y| y - 0.5).collect();
let q_hi: Vec<f64> = actuals.iter().map(|&y| y + 0.5).collect();
let cqr = CqrPredictor::new(0.90);
let result = cqr.fit(&q_lo, &q_hi, &actuals).unwrap();
assert!(
result.quantile_value() <= 0.0,
"well-calibrated CQR should not need to widen, got Q={}",
result.quantile_value()
);
}
#[test]
fn under_covering_base_widens_intervals() {
let n = 200;
let actuals: Vec<f64> = (0..n).map(|i| (i as f64 * 0.1).sin()).collect();
let q_lo: Vec<f64> = actuals.iter().map(|&y| y - 0.05).collect();
let q_hi: Vec<f64> = actuals.iter().map(|&y| y + 0.05).collect();
let mut shifted_actuals = actuals.clone();
for (i, a) in shifted_actuals.iter_mut().enumerate() {
if i % 2 == 0 {
*a += 0.6;
}
}
let cqr = CqrPredictor::new(0.90);
let result = cqr.fit(&q_lo, &q_hi, &shifted_actuals).unwrap();
assert!(
result.quantile_value() > 0.3,
"under-covering CQR should widen significantly, got Q={}",
result.quantile_value()
);
}
#[test]
fn predict_intervals_applies_symmetric_adjustment() {
let actuals = vec![0.0; 100];
let q_lo = vec![-1.0; 100];
let q_hi = vec![1.0; 100];
let cqr = CqrPredictor::new(0.80);
let result = cqr.fit(&q_lo, &q_hi, &actuals).unwrap();
let test_lo = vec![5.0, 10.0];
let test_hi = vec![7.0, 12.0];
let intervals = result.predict_intervals(&test_lo, &test_hi).unwrap();
let q = result.quantile_value();
assert!((intervals.lower()[0] - (5.0 - q)).abs() < 1e-12);
assert!((intervals.upper()[0] - (7.0 + q)).abs() < 1e-12);
assert_eq!(intervals.coverage(), 0.80);
}
#[test]
fn empty_calibration_errors() {
let cqr = CqrPredictor::new(0.90);
let err = cqr.fit(&[], &[], &[]).unwrap_err();
assert!(matches!(err, ForecastError::EmptyData));
}
#[test]
fn dimension_mismatch_errors() {
let cqr = CqrPredictor::new(0.90);
let err = cqr.fit(&[1.0, 2.0], &[3.0], &[4.0, 5.0]).unwrap_err();
assert!(matches!(err, ForecastError::DimensionMismatch { .. }));
}
#[test]
fn tau_lo_tau_hi_match_coverage() {
let cqr = CqrPredictor::new(0.90);
assert!((cqr.tau_lo() - 0.05).abs() < 1e-12);
assert!((cqr.tau_hi() - 0.95).abs() < 1e-12);
}
#[test]
fn coverage_actually_achieved_on_holdout() {
let n_per_set = 500;
let synth = |n: usize, salt: u64| -> Vec<f64> {
(0..n)
.map(|i| {
let h = (i as u64).wrapping_mul(2654435761).wrapping_add(salt);
((h % 1000) as f64 / 1000.0 - 0.5) * 2.0
})
.collect()
};
let eps_calib = synth(n_per_set, 1);
let eps_test = synth(n_per_set, 2);
let q_lo_calib = vec![-0.3; n_per_set];
let q_hi_calib = vec![0.3; n_per_set];
let q_lo_test = vec![-0.3; n_per_set];
let q_hi_test = vec![0.3; n_per_set];
let cqr = CqrPredictor::new(0.90);
let result = cqr.fit(&q_lo_calib, &q_hi_calib, &eps_calib).unwrap();
let intervals = result.predict_intervals(&q_lo_test, &q_hi_test).unwrap();
let cov = intervals.empirical_coverage(&eps_test).unwrap();
assert!(
(0.85..0.95).contains(&cov),
"CQR empirical coverage {:.3} should be near 0.90 (stationary case)",
cov
);
}
}