use std::collections::HashMap;
use crate::error::{ForecastError, Result};
use crate::utils::ols::{ols_fit, ols_residuals};
pub const VIF_WARN: f64 = 5.0;
pub const VIF_FAIL: f64 = 10.0;
pub const COND_WARN: f64 = 30.0;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Severity {
Ok,
Warn,
Fail,
}
#[derive(Debug, Clone)]
pub struct MulticollinearityReport {
pub columns: Vec<(String, f64, Severity)>,
pub condition_number: f64,
pub vif_warn: f64,
pub vif_fail: f64,
pub cond_warn: f64,
}
impl MulticollinearityReport {
pub fn failing(&self) -> Vec<&str> {
self.columns
.iter()
.filter(|(_, _, s)| *s == Severity::Fail)
.map(|(name, _, _)| name.as_str())
.collect()
}
pub fn warning(&self) -> Vec<&str> {
self.columns
.iter()
.filter(|(_, _, s)| *s == Severity::Warn)
.map(|(name, _, _)| name.as_str())
.collect()
}
pub fn is_ill_conditioned(&self) -> bool {
self.condition_number > self.cond_warn
}
}
pub fn variance_inflation_factors(columns: &[Vec<f64>]) -> Result<Vec<f64>> {
let p = columns.len();
if p == 0 {
return Err(ForecastError::EmptyData);
}
let n = columns[0].len();
if n < 2 {
return Err(ForecastError::InsufficientData {
needed: 2,
got: n,
hint: Some("VIF needs ≥ 2 observations".into()),
});
}
for c in columns.iter() {
if c.len() != n {
return Err(ForecastError::DimensionMismatch {
expected: n,
got: c.len(),
});
}
}
if p == 1 {
return Ok(vec![1.0]);
}
let mut vifs = Vec::with_capacity(p);
for j in 0..p {
let y = &columns[j];
let y_mean = y.iter().sum::<f64>() / n as f64;
let tss: f64 = y.iter().map(|v| (v - y_mean).powi(2)).sum();
if tss <= f64::EPSILON {
vifs.push(1.0);
continue;
}
let mut other_regressors: HashMap<String, Vec<f64>> = HashMap::new();
for (k, c) in columns.iter().enumerate() {
if k != j {
other_regressors.insert(format!("c{}", k), c.clone());
}
}
match ols_fit(y, &other_regressors) {
Ok(fit) => match ols_residuals(y, &fit, &other_regressors) {
Ok(res) => {
let rss: f64 = res.iter().map(|r| r * r).sum();
let r2 = 1.0 - rss / tss;
if r2 >= 1.0 - f64::EPSILON {
vifs.push(f64::INFINITY);
} else {
vifs.push(1.0 / (1.0 - r2));
}
}
Err(_) => vifs.push(f64::INFINITY),
},
Err(_) => vifs.push(f64::INFINITY),
}
}
Ok(vifs)
}
pub fn condition_number(columns: &[Vec<f64>]) -> Result<f64> {
let p = columns.len();
if p == 0 {
return Err(ForecastError::EmptyData);
}
let n = columns[0].len();
if n < 2 {
return Err(ForecastError::InsufficientData {
needed: 2,
got: n,
hint: Some("condition_number needs ≥ 2 observations".into()),
});
}
for c in columns.iter() {
if c.len() != n {
return Err(ForecastError::DimensionMismatch {
expected: n,
got: c.len(),
});
}
}
let mut xtx = vec![vec![0.0_f64; p]; p];
for i in 0..p {
for j in i..p {
let mut s = 0.0_f64;
for k in 0..n {
s += columns[i][k] * columns[j][k];
}
xtx[i][j] = s;
xtx[j][i] = s;
}
}
let lambda_max = power_iteration_max(&xtx, p);
let lambda_min = power_iteration_min(&xtx, p);
if !(lambda_max.is_finite() && lambda_min.is_finite()) || lambda_min <= 0.0 {
return Ok(f64::INFINITY);
}
Ok((lambda_max / lambda_min).sqrt())
}
pub fn multicollinearity_report(
columns: &[Vec<f64>],
names: &[String],
) -> Result<MulticollinearityReport> {
multicollinearity_report_with_thresholds(columns, names, VIF_WARN, VIF_FAIL, COND_WARN)
}
pub fn multicollinearity_report_with_thresholds(
columns: &[Vec<f64>],
names: &[String],
vif_warn: f64,
vif_fail: f64,
cond_warn: f64,
) -> Result<MulticollinearityReport> {
if names.len() != columns.len() {
return Err(ForecastError::DimensionMismatch {
expected: columns.len(),
got: names.len(),
});
}
let vifs = variance_inflation_factors(columns)?;
let cond = condition_number(columns)?;
let cols = names
.iter()
.zip(vifs.iter())
.map(|(name, &vif)| {
let severity = if vif > vif_fail {
Severity::Fail
} else if vif > vif_warn {
Severity::Warn
} else {
Severity::Ok
};
(name.clone(), vif, severity)
})
.collect();
Ok(MulticollinearityReport {
columns: cols,
condition_number: cond,
vif_warn,
vif_fail,
cond_warn,
})
}
fn power_iteration_max(a: &[Vec<f64>], p: usize) -> f64 {
if p == 0 {
return 0.0;
}
let mut v = vec![1.0_f64 / (p as f64).sqrt(); p];
let mut lambda = 0.0_f64;
for _ in 0..200 {
let av = mat_vec(a, &v, p);
let new_lambda = dot(&v, &av, p);
let norm = av.iter().map(|x| x * x).sum::<f64>().sqrt();
if norm == 0.0 {
return 0.0;
}
v = av.iter().map(|x| x / norm).collect();
if (new_lambda - lambda).abs() < 1e-10 * new_lambda.abs().max(1.0) {
return new_lambda;
}
lambda = new_lambda;
}
lambda
}
fn power_iteration_min(a: &[Vec<f64>], p: usize) -> f64 {
if p == 0 {
return 0.0;
}
let mut shifted = vec![vec![0.0_f64; p]; p];
for i in 0..p {
for j in 0..p {
shifted[i][j] = a[i][j];
}
shifted[i][i] += 1e-10;
}
let l = match cholesky(&shifted, p) {
Some(l) => l,
None => return 0.0,
};
let mut v = vec![1.0_f64 / (p as f64).sqrt(); p];
let mut lambda = f64::INFINITY;
for _ in 0..200 {
let y = forward_sub(&l, &v, p);
let w = back_sub(&l, &y, p);
let norm = w.iter().map(|x| x * x).sum::<f64>().sqrt();
if norm == 0.0 || !norm.is_finite() {
return 0.0;
}
let v_next: Vec<f64> = w.iter().map(|x| x / norm).collect();
let av = mat_vec(a, &v_next, p);
let new_lambda = dot(&v_next, &av, p);
if (new_lambda - lambda).abs() < 1e-10 * new_lambda.abs().max(1.0) {
return new_lambda.max(0.0);
}
lambda = new_lambda;
v = v_next;
}
lambda.max(0.0)
}
fn mat_vec(a: &[Vec<f64>], v: &[f64], p: usize) -> Vec<f64> {
let mut out = vec![0.0_f64; p];
for i in 0..p {
let mut s = 0.0;
for j in 0..p {
s += a[i][j] * v[j];
}
out[i] = s;
}
out
}
fn dot(a: &[f64], b: &[f64], p: usize) -> f64 {
let mut s = 0.0;
for i in 0..p {
s += a[i] * b[i];
}
s
}
fn cholesky(a: &[Vec<f64>], n: usize) -> Option<Vec<Vec<f64>>> {
let mut l = vec![vec![0.0_f64; n]; n];
for i in 0..n {
for j in 0..=i {
let mut sum = a[i][j];
for k in 0..j {
sum -= l[i][k] * l[j][k];
}
if i == j {
if sum <= 0.0 {
return None;
}
l[i][j] = sum.sqrt();
} else {
if l[j][j] == 0.0 {
return None;
}
l[i][j] = sum / l[j][j];
}
}
}
Some(l)
}
fn forward_sub(l: &[Vec<f64>], b: &[f64], n: usize) -> Vec<f64> {
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][j] * y[j];
}
if l[i][i] == 0.0 {
return vec![f64::INFINITY; n];
}
y[i] = sum / l[i][i];
}
y
}
fn back_sub(l: &[Vec<f64>], y: &[f64], n: usize) -> Vec<f64> {
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][i] * x[j];
}
if l[i][i] == 0.0 {
return vec![f64::INFINITY; n];
}
x[i] = sum / l[i][i];
}
x
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_relative_eq;
#[test]
fn vif_orthogonal_columns_near_one() {
let n = 200;
let c1: Vec<f64> = (0..n).map(|i| (i as f64 * 0.1).sin()).collect();
let c2: Vec<f64> = (0..n).map(|i| (i as f64 * 0.7).cos()).collect();
let c3: Vec<f64> = (0..n).map(|i| (i as f64 * 0.03).sin()).collect();
let vifs = variance_inflation_factors(&[c1, c2, c3]).unwrap();
for v in vifs {
assert!(
v < 2.0,
"orthogonal sinusoids should give VIF near 1, got {}",
v
);
}
}
#[test]
fn vif_perfect_collinearity_is_infinite() {
let c1: Vec<f64> = (0..50).map(|i| i as f64).collect();
let c2: Vec<f64> = c1.iter().map(|x| 2.0 * x).collect();
let vifs = variance_inflation_factors(&[c1, c2]).unwrap();
assert!(vifs[0].is_infinite() || vifs[0] > 1e6);
assert!(vifs[1].is_infinite() || vifs[1] > 1e6);
}
#[test]
fn vif_constant_column_is_one() {
let c1: Vec<f64> = (0..30).map(|i| i as f64).collect();
let c2: Vec<f64> = vec![5.0; 30];
let vifs = variance_inflation_factors(&[c1, c2]).unwrap();
assert_relative_eq!(vifs[1], 1.0, epsilon = 1e-12);
}
#[test]
fn vif_single_column_is_one() {
let c1: Vec<f64> = (0..20).map(|i| i as f64).collect();
let vifs = variance_inflation_factors(&[c1]).unwrap();
assert_eq!(vifs, vec![1.0]);
}
#[test]
fn condition_number_orthogonal_low() {
let n = 100;
let c1: Vec<f64> = (0..n).map(|i| (i as f64 * 0.1).sin()).collect();
let c2: Vec<f64> = (0..n).map(|i| (i as f64 * 0.1).cos()).collect();
let cond = condition_number(&[c1, c2]).unwrap();
assert!(cond < 5.0, "expected near-1 condition, got {}", cond);
}
#[test]
fn condition_number_collinear_huge() {
let c1: Vec<f64> = (0..40).map(|i| i as f64).collect();
let c2: Vec<f64> = c1.iter().map(|x| x + 1e-9).collect();
let cond = condition_number(&[c1, c2]).unwrap();
assert!(
cond > 1e3,
"expected huge condition number for near-duplicates, got {}",
cond
);
}
#[test]
fn report_flags_failing_columns() {
let c1: Vec<f64> = (0..50).map(|i| i as f64).collect();
let c2: Vec<f64> = c1.iter().map(|x| 2.0 * x + 1e-10).collect();
let c3: Vec<f64> = (0..50).map(|i| (i as f64 * 0.3).sin()).collect();
let names = vec!["x1".into(), "x2".into(), "x3".into()];
let report = multicollinearity_report(&[c1, c2, c3], &names).unwrap();
let failing = report.failing();
assert!(failing.contains(&"x1") || failing.contains(&"x2"));
assert!(!failing.contains(&"x3"));
assert!(report.is_ill_conditioned());
}
#[test]
fn report_thresholds_respected() {
let c1: Vec<f64> = (0..30).map(|i| i as f64).collect();
let c2: Vec<f64> = (0..30).map(|i| (i as f64 * 0.5).sin()).collect();
let names = vec!["a".into(), "b".into()];
let report = multicollinearity_report(&[c1, c2], &names).unwrap();
for (_, vif, sev) in &report.columns {
assert!(*vif < VIF_FAIL);
assert!(*sev == Severity::Ok || *sev == Severity::Warn);
}
}
#[test]
fn dimension_mismatch_errors() {
let c1 = vec![1.0, 2.0, 3.0];
let c2 = vec![1.0, 2.0];
let err = variance_inflation_factors(&[c1, c2]).unwrap_err();
assert!(matches!(err, ForecastError::DimensionMismatch { .. }));
}
#[test]
fn empty_columns_errors() {
let err = variance_inflation_factors(&[]).unwrap_err();
assert!(matches!(err, ForecastError::EmptyData));
}
#[test]
fn names_length_mismatch_errors() {
let c1: Vec<f64> = (0..20).map(|i| i as f64).collect();
let c2: Vec<f64> = (0..20).map(|i| (i as f64).sqrt()).collect();
let names = vec!["only_one".into()];
let err = multicollinearity_report(&[c1, c2], &names).unwrap_err();
assert!(matches!(err, ForecastError::DimensionMismatch { .. }));
}
}