use crate::linalg::LinalgInverse as _;
use crate::GreenersError;
use ndarray::{Array1, Array2};
use statrs::distribution::{ContinuousCDF, Normal};
use std::fmt;
#[derive(Debug)]
pub struct CausalImpactResult {
pub y: Array1<f64>,
pub counterfactual: Array1<f64>,
pub counterfactual_sd: Array1<f64>,
pub pointwise_effect: Array1<f64>,
pub cumulative_effect: Array1<f64>,
pub avg_effect: f64,
pub avg_effect_sd: f64,
pub avg_effect_ci: [f64; 2],
pub p_effect_positive: f64,
pub total_effect: f64,
pub total_effect_sd: f64,
pub total_effect_ci: [f64; 2],
pub n_pre: usize,
pub n_post: usize,
pub coefficients: Array1<f64>,
pub control_names: Vec<String>,
pub pre_r_squared: f64,
}
impl fmt::Display for CausalImpactResult {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
writeln!(f, "\n{:=^78}", " Causal Impact ")?;
writeln!(f, "Brodersen, Gallus, Henderson & Orban (2015)")?;
writeln!(f, "Bayesian structural time series")?;
writeln!(f, "{:<20} {:>12}", "Pre-treatment:", self.n_pre)?;
writeln!(f, "{:<20} {:>12}", "Post-treatment:", self.n_post)?;
writeln!(f, "{:<20} {:>12.6}", "Pre-R²:", self.pre_r_squared)?;
writeln!(f, "\n{:-^78}", "")?;
writeln!(f, " Average effect (post-treatment):")?;
writeln!(f, " {:<20} {:>12.6}", "Posterior mean:", self.avg_effect)?;
writeln!(f, " {:<20} {:>12.6}", "SD:", self.avg_effect_sd)?;
writeln!(
f,
" {:<20} [{:.4}, {:.4}]",
"95% CI:", self.avg_effect_ci[0], self.avg_effect_ci[1]
)?;
writeln!(
f,
" {:<20} {:>12.4}",
"P(effect > 0):", self.p_effect_positive
)?;
writeln!(f, "\n Cumulative effect:")?;
writeln!(f, " {:<20} {:>12.6}", "Total:", self.total_effect)?;
writeln!(f, " {:<20} {:>12.6}", "SD:", self.total_effect_sd)?;
writeln!(
f,
" {:<20} [{:.4}, {:.4}]",
"95% CI:", self.total_effect_ci[0], self.total_effect_ci[1]
)?;
writeln!(f, "\n Control coefficients:")?;
writeln!(f, " {:<14} {:>12}", "Variable", "Coef")?;
writeln!(f, "{:-^78}", "")?;
writeln!(f, " {:<14} {:>12.6}", "Intercept", self.coefficients[0])?;
for (j, name) in self.control_names.iter().enumerate() {
if j + 1 < self.coefficients.len() {
writeln!(f, " {:<14} {:>12.6}", name, self.coefficients[j + 1])?;
}
}
writeln!(f, "\n Post-treatment pointwise effects:")?;
let n_show = self.n_post.min(5);
writeln!(
f,
" {:<6} {:>12} {:>12} {:>12}",
"t", "Observed", "Predicted", "Effect"
)?;
writeln!(f, "{:-^78}", "")?;
for i in 0..n_show {
let idx = self.n_pre + i;
writeln!(
f,
" {:<6} {:>12.4} {:>12.4} {:>12.4}",
idx + 1,
self.y[idx],
self.counterfactual[idx],
self.pointwise_effect[idx]
)?;
}
write!(f, "{:=^78}", "")
}
}
pub struct CausalImpact;
impl CausalImpact {
pub fn fit(
y: &Array1<f64>,
controls: &Array2<f64>,
treatment_period: usize,
control_names: Option<Vec<String>>,
) -> Result<CausalImpactResult, GreenersError> {
let n = y.len();
let k = controls.ncols();
if controls.nrows() != n {
return Err(GreenersError::ShapeMismatch(
"CausalImpact: y and controls must have same n".into(),
));
}
if treatment_period >= n - 2 {
return Err(GreenersError::InvalidOperation(
"CausalImpact: treatment_period must leave at least 2 post-period obs".into(),
));
}
if treatment_period < k + 2 {
return Err(GreenersError::InvalidOperation(
"CausalImpact: pre-treatment period too short for controls".into(),
));
}
let names = control_names.unwrap_or_else(|| (0..k).map(|i| format!("c{}", i)).collect());
let n_pre = treatment_period;
let n_post = n - treatment_period;
let mut x_pre = Array2::zeros((n_pre, k + 1));
let mut y_pre = Array1::zeros(n_pre);
for i in 0..n_pre {
x_pre[(i, 0)] = 1.0;
for j in 0..k {
x_pre[(i, j + 1)] = controls[(i, j)];
}
y_pre[i] = y[i];
}
let xt = x_pre.t();
let xtx = xt.dot(&x_pre);
let xtx_inv = (&xtx + Array2::<f64>::eye(k + 1) * 1e-8).inv()?;
let xty = xt.dot(&y_pre);
let beta: Array1<f64> = xtx_inv.dot(&xty);
let y_pre_hat = x_pre.dot(&beta);
let residuals_pre = &y_pre - &y_pre_hat;
let sigma2 = residuals_pre.mapv(|r| r * r).sum() / (n_pre - k - 1) as f64;
let _sigma = sigma2.sqrt();
let y_pre_mean = y_pre.mean().unwrap_or(0.0);
let tss = y_pre.mapv(|v| (v - y_pre_mean).powi(2)).sum();
let sse = residuals_pre.mapv(|r| r * r).sum();
let pre_r_squared = if tss > 1e-15 { 1.0 - sse / tss } else { 0.0 };
let mut counterfactual = Array1::zeros(n);
let mut counterfactual_sd = Array1::zeros(n);
for i in 0..n {
let mut x_i = Array1::zeros(k + 1);
x_i[0] = 1.0;
for j in 0..k {
x_i[j + 1] = controls[(i, j)];
}
counterfactual[i] = beta.dot(&x_i);
let pred_var = x_i.dot(&xtx_inv).dot(&x_i) * sigma2 + sigma2;
counterfactual_sd[i] = pred_var.sqrt();
}
let level_var = sigma2 * 0.1;
let mut level = 0.0;
for i in 0..n {
if i < n_pre {
let resid = y[i] - counterfactual[i];
level = level * 0.9 + resid * 0.1;
}
counterfactual[i] += level;
counterfactual_sd[i] = (counterfactual_sd[i].powi(2) + level_var).sqrt();
}
let pointwise_effect = y - &counterfactual;
let mut cumulative_effect = Array1::zeros(n);
let mut cumsum = 0.0;
for i in 0..n {
if i >= n_pre {
cumsum += pointwise_effect[i];
}
cumulative_effect[i] = cumsum;
}
let post_effects: Vec<f64> = (n_pre..n).map(|i| pointwise_effect[i]).collect();
let avg_effect: f64 = post_effects.iter().sum::<f64>() / n_post as f64;
let post_sds: Vec<f64> = (n_pre..n).map(|i| counterfactual_sd[i]).collect();
let avg_var: f64 = post_sds.iter().map(|s| s * s).sum::<f64>() / (n_post * n_post) as f64;
let avg_effect_sd = avg_var.sqrt();
let total_effect: f64 = post_effects.iter().sum();
let total_var: f64 = post_sds.iter().map(|s| s * s).sum();
let total_effect_sd = total_var.sqrt();
let z = 1.959964;
let avg_effect_ci = [
avg_effect - z * avg_effect_sd,
avg_effect + z * avg_effect_sd,
];
let total_effect_ci = [
total_effect - z * total_effect_sd,
total_effect + z * total_effect_sd,
];
let normal =
Normal::new(0.0, 1.0).map_err(|e| GreenersError::InvalidOperation(e.to_string()))?;
let p_effect_positive = 1.0 - normal.cdf(avg_effect / avg_effect_sd.max(1e-10));
Ok(CausalImpactResult {
y: y.clone(),
counterfactual,
counterfactual_sd,
pointwise_effect,
cumulative_effect,
avg_effect,
avg_effect_sd,
avg_effect_ci,
p_effect_positive,
total_effect,
total_effect_sd,
total_effect_ci,
n_pre,
n_post,
coefficients: beta,
control_names: names,
pre_r_squared,
})
}
}