use crate::linalg::LinalgInverse as _;
use crate::GreenersError;
use ndarray::{Array1, Array2};
use statrs::distribution::{ContinuousCDF, Normal};
use std::fmt;
#[derive(Debug)]
pub struct SyntheticDidResult {
pub att: f64,
pub se: f64,
pub t_stat: f64,
pub p_value: f64,
pub unit_weights: Array1<f64>,
pub time_weights: Array1<f64>,
pub synthetic_control: Array1<f64>,
pub treated_avg: Array1<f64>,
pub n_treated: usize,
pub n_control: usize,
pub n_pre: usize,
pub n_post: usize,
pub n_periods: usize,
}
impl fmt::Display for SyntheticDidResult {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
writeln!(f, "\n{:=^78}", " Synthetic DiD ")?;
writeln!(f, "Arkhangelsky et al. (2021)")?;
writeln!(f, "{:<20} {:>12}", "Treated units:", self.n_treated)?;
writeln!(f, "{:<20} {:>12}", "Control units:", self.n_control)?;
writeln!(f, "{:<20} {:>12}", "Pre periods:", self.n_pre)?;
writeln!(f, "{:<20} {:>12}", "Post periods:", self.n_post)?;
writeln!(f, "{:<20} {:>12.6}", "ATT (Synthetic DiD):", self.att)?;
writeln!(f, "{:<20} {:>12.6}", "Std. Error:", self.se)?;
writeln!(f, "{:<20} {:>12.3}", "t-statistic:", self.t_stat)?;
writeln!(f, "{:<20} {:>12.4}", "p-value:", self.p_value)?;
writeln!(f, "\n{:-^78}", "")?;
writeln!(f, " Unit weights (synthetic control):")?;
for (j, &w) in self.unit_weights.iter().enumerate() {
if w.abs() > 1e-6 {
writeln!(f, " Unit {:<6} {:>12.6}", j + 1, w)?;
}
}
writeln!(f, "\n Time weights (synthetic pre-period):")?;
for (t, &w) in self.time_weights.iter().enumerate() {
writeln!(f, " Period {:<5} {:>12.6}", t + 1, w)?;
}
writeln!(f, "\n Outcome paths (selected periods):")?;
writeln!(
f,
" {:<8} {:>14} {:>14} {:>14}",
"Period", "Treated avg", "Synth. control", "Gap"
)?;
let n_show = 5.min(self.n_periods);
let indices: Vec<usize> = if self.n_periods <= n_show {
(0..self.n_periods).collect()
} else {
(0..n_show)
.map(|i| i * (self.n_periods - 1) / (n_show - 1).max(1))
.collect()
};
for &idx in &indices {
let gap = self.treated_avg[idx] - self.synthetic_control[idx];
writeln!(
f,
" {:<8} {:>14.6} {:>14.6} {:>14.6}",
idx + 1,
self.treated_avg[idx],
self.synthetic_control[idx],
gap
)?;
}
write!(f, "{:=^78}", "")
}
}
pub struct SyntheticDiD;
impl SyntheticDiD {
pub fn fit(
y: &Array2<f64>,
treated: &[bool],
treatment_period: usize,
) -> Result<SyntheticDidResult, GreenersError> {
let n = y.nrows();
let t = y.ncols();
if treated.len() != n {
return Err(GreenersError::ShapeMismatch(
"SyntheticDiD: treated length must match y rows".into(),
));
}
if treatment_period >= t || treatment_period == 0 {
return Err(GreenersError::InvalidOperation(
"SyntheticDiD: treatment_period must be in (0, T)".into(),
));
}
let treated_indices: Vec<usize> = (0..n).filter(|&i| treated[i]).collect();
let control_indices: Vec<usize> = (0..n).filter(|&i| !treated[i]).collect();
let n_treated = treated_indices.len();
let n_control = control_indices.len();
let n_pre = treatment_period;
let n_post = t - treatment_period;
if n_treated == 0 || n_control == 0 {
return Err(GreenersError::InvalidOperation(
"SyntheticDiD: need at least 1 treated and 1 control unit".into(),
));
}
let mut treated_avg = Array1::zeros(t);
for &i in &treated_indices {
for time in 0..t {
treated_avg[time] += y[(i, time)];
}
}
treated_avg /= n_treated as f64;
let y_treated_pre = treated_avg.slice(ndarray::s![0..n_pre]).to_owned();
let mut y_control_pre = Array2::zeros((n_pre, n_control));
for (j, &ci) in control_indices.iter().enumerate() {
for time in 0..n_pre {
y_control_pre[(time, j)] = y[(ci, time)];
}
}
let xt = y_control_pre.t();
let xtx = xt.dot(&y_control_pre);
let xtx_inv = (&xtx + Array2::<f64>::eye(n_control) * 1e-6).inv()?;
let xty = xt.dot(&y_treated_pre);
let mut unit_weights: Array1<f64> = xtx_inv.dot(&xty);
for w in unit_weights.iter_mut() {
if *w < 0.0 {
*w = 0.0;
}
}
let w_sum: f64 = unit_weights.sum();
if w_sum > 1e-10 {
unit_weights /= w_sum;
}
let y_control_pre_avg: Array1<f64> = (0..n_pre)
.map(|time| {
control_indices.iter().map(|&ci| y[(ci, time)]).sum::<f64>() / n_control as f64
})
.collect();
let mut time_weights = Array1::zeros(n_pre);
if n_pre > 1 {
let mut x_time = Array2::zeros((n_pre - 1, n_pre - 1));
let mut y_time = Array1::zeros(n_pre - 1);
for i in 0..n_pre - 1 {
y_time[i] = y_control_pre_avg[n_pre - 1];
for j in 0..n_pre - 1 {
x_time[(i, j)] = y_control_pre_avg[j];
}
}
for w in time_weights.iter_mut() {
*w = 1.0 / n_pre as f64;
}
} else {
time_weights[0] = 1.0;
}
let mut synthetic_control = Array1::zeros(t);
for time in 0..t {
let mut val = 0.0;
for (j, &ci) in control_indices.iter().enumerate() {
val += unit_weights[j] * y[(ci, time)];
}
synthetic_control[time] = val;
}
let treated_post: f64 = treated_avg
.slice(ndarray::s![treatment_period..t])
.mean()
.unwrap_or(0.0);
let synth_post: f64 = synthetic_control
.slice(ndarray::s![treatment_period..t])
.mean()
.unwrap_or(0.0);
let treated_pre: f64 = treated_avg
.slice(ndarray::s![0..n_pre])
.mean()
.unwrap_or(0.0);
let synth_pre: f64 = synthetic_control
.slice(ndarray::s![0..n_pre])
.mean()
.unwrap_or(0.0);
let att = (treated_post - synth_post) - (treated_pre - synth_pre);
let mut placebo_atts: Vec<f64> = Vec::new();
for &placebo_treated in &control_indices {
let mut placebo_treated_vec = vec![false; n];
placebo_treated_vec[placebo_treated] = true;
let placebo_controls: Vec<usize> = control_indices
.iter()
.copied()
.filter(|&i| i != placebo_treated)
.collect();
if placebo_controls.is_empty() {
continue;
}
let placebo_treated_avg = y.row(placebo_treated).to_owned();
let mut placebo_y_control_pre = Array2::zeros((n_pre, placebo_controls.len()));
for (j, &ci) in placebo_controls.iter().enumerate() {
for time in 0..n_pre {
placebo_y_control_pre[(time, j)] = y[(ci, time)];
}
}
let placebo_xt = placebo_y_control_pre.t();
let placebo_xtx = placebo_xt.dot(&placebo_y_control_pre);
let placebo_xtx_inv =
match (&placebo_xtx + Array2::<f64>::eye(placebo_controls.len()) * 1e-6).inv() {
Ok(v) => v,
Err(_) => continue,
};
let placebo_xty = placebo_xt.dot(&placebo_treated_avg.slice(ndarray::s![0..n_pre]));
let mut placebo_w: Array1<f64> = placebo_xtx_inv.dot(&placebo_xty);
for w in placebo_w.iter_mut() {
if *w < 0.0 {
*w = 0.0;
}
}
let pw_sum: f64 = placebo_w.sum();
if pw_sum > 1e-10 {
placebo_w /= pw_sum;
}
let mut placebo_synth = Array1::zeros(t);
for time in 0..t {
let mut val = 0.0;
for (j, &ci) in placebo_controls.iter().enumerate() {
val += placebo_w[j] * y[(ci, time)];
}
placebo_synth[time] = val;
}
let pt_post: f64 = placebo_treated_avg
.slice(ndarray::s![treatment_period..t])
.mean()
.unwrap_or(0.0);
let ps_post: f64 = placebo_synth
.slice(ndarray::s![treatment_period..t])
.mean()
.unwrap_or(0.0);
let pt_pre: f64 = placebo_treated_avg
.slice(ndarray::s![0..n_pre])
.mean()
.unwrap_or(0.0);
let ps_pre: f64 = placebo_synth
.slice(ndarray::s![0..n_pre])
.mean()
.unwrap_or(0.0);
placebo_atts.push((pt_post - ps_post) - (pt_pre - ps_pre));
}
let se = if placebo_atts.len() > 1 {
let mean_placebo = placebo_atts.iter().sum::<f64>() / placebo_atts.len() as f64;
let var = placebo_atts
.iter()
.map(|a| (a - mean_placebo).powi(2))
.sum::<f64>()
/ (placebo_atts.len() - 1) as f64;
var.sqrt()
} else {
let residuals: Vec<f64> = (0..n_pre)
.map(|time| treated_avg[time] - synthetic_control[time])
.collect();
let res_var = residuals.iter().map(|r| r * r).sum::<f64>() / n_pre as f64;
(res_var / n_pre as f64).sqrt() + 1e-10
};
let t_stat = if se > 1e-10 { att / se } else { 0.0 };
let normal =
Normal::new(0.0, 1.0).map_err(|e| GreenersError::InvalidOperation(e.to_string()))?;
let p_value = 2.0 * (1.0 - normal.cdf(t_stat.abs()));
Ok(SyntheticDidResult {
att,
se,
t_stat,
p_value,
unit_weights,
time_weights,
synthetic_control,
treated_avg,
n_treated,
n_control,
n_pre,
n_post,
n_periods: t,
})
}
}