use std::collections::HashMap;
use crate::core::Forecast;
use crate::error::{ForecastError, Result};
use crate::models::Forecaster;
pub const DEFAULT_RIDGE_LAMBDA: f64 = 0.1;
pub fn residual_ridge_shim<F: Forecaster + ?Sized>(
model: &F,
horizon: usize,
future_regressors: &HashMap<String, Vec<f64>>,
ridge_lambda: f64,
) -> Result<Forecast> {
let base = model.predict(horizon)?;
if future_regressors.is_empty() {
return Ok(base);
}
let mut names: Vec<&String> = future_regressors.keys().collect();
names.sort();
for &name in &names {
let v = &future_regressors[name];
if v.len() != horizon {
return Err(ForecastError::DimensionMismatch {
expected: horizon,
got: v.len(),
});
}
}
let residuals_vec = model.residual_component()?;
let train_regs = model.training_regressors().ok_or_else(|| {
ForecastError::InvalidParameter(format!(
"{} does not retain training regressors — residual-Ridge shim unavailable",
model.name()
))
})?;
for &name in &names {
if !train_regs.contains_key(name) {
return Err(ForecastError::InvalidParameter(format!(
"residual_ridge_shim: future regressor '{}' was not present at fit time",
name
)));
}
}
let residual_len = residuals_vec.len();
let p = names.len();
let mut x_rows: Vec<f64> = Vec::with_capacity(residual_len * p);
let mut y_rows: Vec<f64> = Vec::with_capacity(residual_len);
for (i, &r) in residuals_vec.iter().enumerate() {
if !r.is_finite() {
continue;
}
let mut row = Vec::with_capacity(p);
let mut row_finite = true;
for &name in &names {
let col = &train_regs[name];
if col.len() < residual_len {
return Err(ForecastError::DimensionMismatch {
expected: residual_len,
got: col.len(),
});
}
let offset = col.len() - residual_len + i;
let v = col[offset];
if !v.is_finite() {
row_finite = false;
break;
}
row.push(v);
}
if !row_finite {
continue;
}
x_rows.extend(row);
y_rows.push(r);
}
let n = y_rows.len();
if n == 0 {
return Err(ForecastError::InvalidParameter(
"residual_ridge_shim: no finite residuals to fit Ridge on".into(),
));
}
if n < p {
return Err(ForecastError::InsufficientData {
needed: p,
got: n,
hint: Some(format!(
"residual_ridge_shim: fewer finite residuals ({}) than regressors ({}); reduce regressor count or increase training window",
n, p
)),
});
}
let beta = solve_ridge(&x_rows, &y_rows, n, p, ridge_lambda)?;
let mut adjustment = vec![0.0_f64; horizon];
for (j, &name) in names.iter().enumerate() {
let fut = &future_regressors[name];
for (h, &v) in fut.iter().enumerate() {
adjustment[h] += beta[j] * v;
}
}
apply_adjustment(base, &adjustment)
}
fn apply_adjustment(mut base: Forecast, adjustment: &[f64]) -> Result<Forecast> {
let primary = base.primary_mut();
if primary.len() != adjustment.len() {
return Err(ForecastError::DimensionMismatch {
expected: primary.len(),
got: adjustment.len(),
});
}
for (p, a) in primary.iter_mut().zip(adjustment) {
*p += a;
}
Ok(base)
}
fn solve_ridge(x: &[f64], y: &[f64], n: usize, p: usize, lambda: f64) -> Result<Vec<f64>> {
let mut xtx = vec![0.0_f64; p * p];
for i in 0..p {
for j in 0..p {
let mut s = 0.0_f64;
for k in 0..n {
s += x[k * p + i] * x[k * p + j];
}
xtx[i * p + j] = s;
}
xtx[i * p + i] += lambda;
}
let mut xty = vec![0.0_f64; p];
for j in 0..p {
let mut s = 0.0_f64;
for k in 0..n {
s += x[k * p + j] * y[k];
}
xty[j] = s;
}
cholesky_solve(&xtx, &xty, p).ok_or_else(|| {
ForecastError::SingularMatrix(format!(
"residual_ridge_shim: X'X + {}·I is not positive-definite",
lambda
))
})
}
fn cholesky_solve(a: &[f64], b: &[f64], n: usize) -> Option<Vec<f64>> {
let mut l = vec![0.0_f64; n * n];
for i in 0..n {
for j in 0..=i {
let mut sum = a[i * n + j];
for k in 0..j {
sum -= l[i * n + k] * l[j * n + k];
}
if i == j {
if sum <= 0.0 {
return None;
}
l[i * n + j] = sum.sqrt();
} else {
if l[j * n + j] == 0.0 {
return None;
}
l[i * n + j] = sum / l[j * n + j];
}
}
}
let mut y = vec![0.0_f64; n];
for i in 0..n {
let mut sum = b[i];
for j in 0..i {
sum -= l[i * n + j] * y[j];
}
y[i] = sum / l[i * n + i];
}
let mut x = vec![0.0_f64; n];
for i in (0..n).rev() {
let mut sum = y[i];
for j in (i + 1)..n {
sum -= l[j * n + i] * x[j];
}
x[i] = sum / l[i * n + i];
}
Some(x)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::TimeSeries;
#[derive(Debug)]
struct MockDecomposable {
residuals: Vec<f64>,
train_regs: Option<HashMap<String, Vec<f64>>>,
base_forecast: Forecast,
}
impl Forecaster for MockDecomposable {
fn fit(&mut self, _series: &TimeSeries) -> Result<()> {
Ok(())
}
fn predict(&self, _horizon: usize) -> Result<Forecast> {
Ok(self.base_forecast.clone())
}
fn fitted_values(&self) -> Option<&[f64]> {
None
}
fn residuals(&self) -> Option<&[f64]> {
Some(&self.residuals)
}
fn training_regressors(&self) -> Option<&HashMap<String, Vec<f64>>> {
self.train_regs.as_ref()
}
fn name(&self) -> &str {
"MockDecomposable"
}
}
#[test]
fn shim_recovers_known_coefficient() {
let n = 60;
let x: Vec<f64> = (0..n).map(|i| ((i as f64) * 0.1).sin()).collect();
let residuals: Vec<f64> = (0..n)
.map(|i| 5.0 * x[i] + ((i % 7) as f64 - 3.0) * 0.001)
.collect();
let mut train_regs = HashMap::new();
train_regs.insert("x".to_string(), x.clone());
let horizon = 5;
let model = MockDecomposable {
residuals,
train_regs: Some(train_regs),
base_forecast: Forecast::from_values(vec![10.0; horizon]),
};
let future_x: Vec<f64> = (n..n + horizon).map(|i| ((i as f64) * 0.1).sin()).collect();
let mut future_regs = HashMap::new();
future_regs.insert("x".to_string(), future_x.clone());
let adjusted =
residual_ridge_shim(&model, horizon, &future_regs, DEFAULT_RIDGE_LAMBDA).unwrap();
for h in 0..horizon {
let expected = 10.0 + 5.0 * future_x[h];
let got = adjusted.primary()[h];
assert!(
(got - expected).abs() < 0.5,
"h={}: expected ≈ {}, got {}",
h,
expected,
got
);
}
}
#[test]
fn shim_empty_future_regressors_is_base_forecast() {
let n = 30;
let residuals: Vec<f64> = (0..n).map(|i| (i as f64) * 0.01).collect();
let horizon = 3;
let model = MockDecomposable {
residuals,
train_regs: Some(HashMap::new()),
base_forecast: Forecast::from_values(vec![1.0, 2.0, 3.0]),
};
let adjusted =
residual_ridge_shim(&model, horizon, &HashMap::new(), DEFAULT_RIDGE_LAMBDA).unwrap();
assert_eq!(adjusted.primary(), &[1.0, 2.0, 3.0]);
}
#[test]
fn shim_rejects_horizon_mismatch_on_future_regs() {
let n = 20;
let residuals = vec![0.0; n];
let mut train_regs = HashMap::new();
train_regs.insert("x".to_string(), vec![1.0; n]);
let horizon = 3;
let model = MockDecomposable {
residuals,
train_regs: Some(train_regs),
base_forecast: Forecast::from_values(vec![1.0; horizon]),
};
let mut future_regs = HashMap::new();
future_regs.insert("x".to_string(), vec![1.0; 5]); let err =
residual_ridge_shim(&model, horizon, &future_regs, DEFAULT_RIDGE_LAMBDA).unwrap_err();
assert!(matches!(err, ForecastError::DimensionMismatch { .. }));
}
#[test]
fn shim_rejects_unknown_regressor_name() {
let n = 20;
let residuals = vec![0.0; n];
let mut train_regs = HashMap::new();
train_regs.insert("x".to_string(), vec![1.0; n]);
let horizon = 3;
let model = MockDecomposable {
residuals,
train_regs: Some(train_regs),
base_forecast: Forecast::from_values(vec![1.0; horizon]),
};
let mut future_regs = HashMap::new();
future_regs.insert("z".to_string(), vec![1.0; horizon]); let err =
residual_ridge_shim(&model, horizon, &future_regs, DEFAULT_RIDGE_LAMBDA).unwrap_err();
assert!(matches!(err, ForecastError::InvalidParameter(_)));
}
#[test]
fn shim_errors_when_model_has_no_training_regressors() {
let horizon = 3;
let model = MockDecomposable {
residuals: vec![0.0; 20],
train_regs: None, base_forecast: Forecast::from_values(vec![1.0; horizon]),
};
let mut future_regs = HashMap::new();
future_regs.insert("anything".to_string(), vec![1.0; horizon]);
let err =
residual_ridge_shim(&model, horizon, &future_regs, DEFAULT_RIDGE_LAMBDA).unwrap_err();
assert!(matches!(err, ForecastError::InvalidParameter(_)));
}
}