#![allow(
clippy::cast_precision_loss,
clippy::cast_possible_truncation,
clippy::similar_names,
clippy::needless_range_loop
)]
use antecedent_core::{
AssumptionSet, AverageEffectQuery, ExecutionContext, PopulationRegistry, TargetPopulation,
};
use antecedent_data::TabularData;
use antecedent_expr::IdentifiedEstimand;
use antecedent_stats::{
DenseLinearAlgebra, FaerBackend, GlmOptions, LeastSquaresWorkspace, PropensityWorkspace,
};
use crate::adjustment::EffectEstimate;
use crate::error::EstimationError;
use crate::overlap::{IpwTarget, OverlapPolicy};
use crate::propensity::{
PreparedPropensityProblem, PropensityModel, clamp_scores, clip_of, default_propensity_overlap,
gather, gather_into, prepare_propensity_problem_with_registry, split_by_treatment, trim_of,
trim_retained_rows,
};
use crate::se::AnalyticSeKind;
use crate::util::{BootstrapSeResult, bootstrap_se, stats_err};
#[derive(Clone, Debug, Default)]
pub struct AipwWorkspace {
pub propensity: PropensityWorkspace,
pub outcome: LeastSquaresWorkspace,
treated_design: Vec<f64>,
treated_outcome: Vec<f64>,
control_design: Vec<f64>,
control_outcome: Vec<f64>,
mu0: Vec<f64>,
mu1: Vec<f64>,
psi: Vec<f64>,
}
#[derive(Clone, Debug)]
pub struct AipwAte {
pub backend: FaerBackend,
pub bootstrap_replicates: u32,
pub overlap: OverlapPolicy,
pub glm_options: GlmOptions,
pub se_kind: AnalyticSeKind,
pub cluster_ids: Option<Vec<u32>>,
pub population_registry: Option<PopulationRegistry>,
pub multiway_ids: Option<Vec<Vec<u32>>>,
pub panel_times: Option<Vec<i64>>,
}
impl Default for AipwAte {
fn default() -> Self {
Self::new()
}
}
impl AipwAte {
#[must_use]
pub fn new() -> Self {
Self {
backend: FaerBackend,
bootstrap_replicates: 200,
overlap: default_propensity_overlap(),
glm_options: GlmOptions::default(),
se_kind: AnalyticSeKind::Homoskedastic,
cluster_ids: None,
population_registry: None,
multiway_ids: None,
panel_times: None,
}
}
#[must_use]
pub const fn with_backend(mut self, backend: FaerBackend) -> Self {
self.backend = backend;
self
}
#[must_use]
pub const fn with_bootstrap_replicates(mut self, replicates: u32) -> Self {
self.bootstrap_replicates = replicates;
self
}
#[must_use]
pub const fn with_overlap(mut self, overlap: OverlapPolicy) -> Self {
self.overlap = overlap;
self
}
#[must_use]
pub const fn with_glm_options(mut self, glm_options: GlmOptions) -> Self {
self.glm_options = glm_options;
self
}
#[must_use]
pub const fn with_se_kind(mut self, se_kind: AnalyticSeKind) -> Self {
self.se_kind = se_kind;
self
}
#[must_use]
pub fn with_cluster_ids(mut self, cluster_ids: Vec<u32>) -> Self {
self.cluster_ids = Some(cluster_ids);
self
}
#[must_use]
pub fn with_population_registry(mut self, registry: PopulationRegistry) -> Self {
self.population_registry = Some(registry);
self
}
#[must_use]
pub fn with_multiway_ids(mut self, multiway_ids: Vec<Vec<u32>>) -> Self {
self.multiway_ids = Some(multiway_ids);
self
}
#[must_use]
pub fn with_panel_times(mut self, panel_times: Vec<i64>) -> Self {
self.panel_times = Some(panel_times);
self
}
pub fn prepare(
&self,
data: &TabularData,
estimand: &IdentifiedEstimand,
query: &AverageEffectQuery,
) -> Result<PreparedPropensityProblem, EstimationError> {
prepare_propensity_problem_with_registry(
data,
estimand,
query,
self.overlap,
self.population_registry.as_ref(),
)
}
#[allow(clippy::too_many_lines)]
pub fn fit(
&self,
problem: &PreparedPropensityProblem,
workspace: &mut AipwWorkspace,
ctx: &ExecutionContext,
assumptions: AssumptionSet,
) -> Result<EffectEstimate, EstimationError> {
if !matches!(
problem.target_population,
TargetPopulation::AllObserved
| TargetPopulation::Treated
| TargetPopulation::Untreated
| TargetPopulation::Predicate(_)
) {
return Err(EstimationError::unsupported(
"AIPW supports AllObserved, Treated, Untreated, or Predicate target populations",
));
}
let model = PropensityModel::fit(
problem,
&self.backend,
&mut workspace.propensity,
&self.glm_options,
)?;
let retained = trim_retained_rows(&model.fit.scores, trim_of(problem.overlap))?;
let ncols = problem.design_ncols;
let (design_used, t_used, y_used, e_used) = match &retained {
Some(idx) => {
let mut design = Vec::new();
select_rows_colmajor(
&problem.design_matrix,
problem.nrows,
ncols,
idx,
&mut design,
);
(
design,
gather(&problem.treatment, idx),
gather(&problem.outcome, idx),
gather(&model.clipped_scores, idx),
)
}
None => (
problem.design_matrix.to_vec(),
problem.treatment.to_vec(),
problem.outcome.to_vec(),
model.clipped_scores.clone(),
),
};
let nrows = t_used.len();
let (beta0, beta1) = fit_outcome_models(
&design_used,
nrows,
ncols,
&t_used,
&y_used,
self.backend,
workspace,
)?;
predict_colmajor(&design_used, nrows, ncols, &beta0, &mut workspace.mu0);
predict_colmajor(&design_used, nrows, ncols, &beta1, &mut workspace.mu1);
aipw_psi(
&t_used,
&y_used,
&e_used,
&workspace.mu0,
&workspace.mu1,
&problem.target_population,
&mut workspace.psi,
)?;
let ate = workspace.psi.iter().sum::<f64>() / workspace.psi.len() as f64;
residualize_aipw_psi(&mut workspace.psi, &t_used, &e_used, &design_used, ncols)?;
let se_analytic = crate::se::influence_se_kind(
self.se_kind,
&workspace.psi,
problem.nrows,
self.cluster_ids.as_deref(),
self.multiway_ids.as_deref(),
self.panel_times.as_deref(),
retained.as_deref(),
)?;
let boot = if self.bootstrap_replicates == 0 {
None
} else {
Some(self.bootstrap_se(problem, workspace, ctx)?)
};
let overlap_report = Some(crate::propensity::propensity_overlap_report(
problem,
&model.fit.scores,
None,
IpwTarget::from_population(&problem.target_population).ok(),
));
Ok(EffectEstimate::new(ate, se_analytic, assumptions, problem.overlap)
.with_overlap_report(overlap_report)
.with_bootstrap(boot))
}
fn bootstrap_se(
&self,
problem: &PreparedPropensityProblem,
workspace: &mut AipwWorkspace,
ctx: &ExecutionContext,
) -> Result<BootstrapSeResult, EstimationError> {
let clip = clip_of(problem.overlap);
let trim = trim_of(problem.overlap);
let n = problem.nrows;
let ncols = problem.design_ncols;
let mut x_boot = vec![0.0; n * ncols];
let mut t_boot = vec![0.0; n];
let mut y_boot = vec![0.0; n];
let mut e = vec![0.0; n];
let mut design_trim: Vec<f64> = Vec::new();
let mut t_trim: Vec<f64> = Vec::new();
let mut y_trim: Vec<f64> = Vec::new();
let mut e_trim: Vec<f64> = Vec::new();
bootstrap_se(self.bootstrap_replicates, ctx, 0xA1D0_u64, n, |idx| {
crate::util::gather_bootstrap_vector(&mut t_boot, &problem.treatment, idx);
crate::util::gather_bootstrap_vector(&mut y_boot, &problem.outcome, idx);
crate::util::gather_bootstrap_design(
&mut x_boot,
&problem.design_matrix,
n,
ncols,
idx,
);
if antecedent_stats::fit_propensity_in_place(
&x_boot,
n,
ncols,
&t_boot,
&self.backend,
&mut workspace.propensity,
&self.glm_options,
)
.is_err()
{
return Ok(None);
}
let raw = &workspace.propensity.scores[..n];
e.copy_from_slice(raw);
if let Some(c) = clip {
clamp_scores(&mut e, c);
}
let Ok(retained) = trim_retained_rows(raw, trim) else {
return Ok(None);
};
let (design_used, t_used, y_used, e_used): (&[f64], &[f64], &[f64], &[f64]) =
match &retained {
Some(rows) => {
select_rows_colmajor(&x_boot, n, ncols, rows, &mut design_trim);
gather_into(&mut t_trim, &t_boot, rows);
gather_into(&mut y_trim, &y_boot, rows);
gather_into(&mut e_trim, &e, rows);
(&design_trim, &t_trim, &y_trim, &e_trim)
}
None => (&x_boot, &t_boot, &y_boot, &e),
};
let nrows = t_used.len();
let Ok((beta0, beta1)) = fit_outcome_models(
design_used,
nrows,
ncols,
t_used,
y_used,
self.backend,
workspace,
) else {
return Ok(None);
};
predict_colmajor(design_used, nrows, ncols, &beta0, &mut workspace.mu0);
predict_colmajor(design_used, nrows, ncols, &beta1, &mut workspace.mu1);
if aipw_psi(
t_used,
y_used,
e_used,
&workspace.mu0,
&workspace.mu1,
&problem.target_population,
&mut workspace.psi,
)
.is_err()
{
return Ok(None);
}
let m = workspace.psi.len() as f64;
Ok(Some(workspace.psi.iter().sum::<f64>() / m))
})
}
}
fn select_rows_colmajor(
matrix: &[f64],
nrows: usize,
ncols: usize,
idx: &[usize],
out: &mut Vec<f64>,
) {
let m = idx.len();
out.clear();
out.resize(m * ncols, 0.0);
for c in 0..ncols {
let src_base = c * nrows;
let dst_base = c * m;
for (r, &i) in idx.iter().enumerate() {
out[dst_base + r] = matrix[src_base + i];
}
}
}
fn select_values(values: &[f64], idx: &[usize], out: &mut Vec<f64>) {
out.clear();
out.extend(idx.iter().map(|&i| values[i]));
}
fn fit_outcome_models(
design_matrix: &[f64],
nrows: usize,
ncols: usize,
treatment: &[f64],
outcome: &[f64],
backend: FaerBackend,
workspace: &mut AipwWorkspace,
) -> Result<(Vec<f64>, Vec<f64>), EstimationError> {
let (treated_idx, control_idx) = split_by_treatment(treatment);
if treated_idx.is_empty() || control_idx.is_empty() {
return Err(EstimationError::data_msg(
"AIPW outcome regression requires both treated and control rows",
));
}
select_rows_colmajor(design_matrix, nrows, ncols, &control_idx, &mut workspace.control_design);
select_values(outcome, &control_idx, &mut workspace.control_outcome);
let fit0 = backend
.least_squares(
&workspace.control_design,
control_idx.len(),
ncols,
&workspace.control_outcome,
&mut workspace.outcome,
)
.map_err(stats_err)?;
let beta0 = fit0.coefficients;
select_rows_colmajor(design_matrix, nrows, ncols, &treated_idx, &mut workspace.treated_design);
select_values(outcome, &treated_idx, &mut workspace.treated_outcome);
let fit1 = backend
.least_squares(
&workspace.treated_design,
treated_idx.len(),
ncols,
&workspace.treated_outcome,
&mut workspace.outcome,
)
.map_err(stats_err)?;
let beta1 = fit1.coefficients;
Ok((beta0, beta1))
}
fn predict_colmajor(
design_matrix: &[f64],
nrows: usize,
ncols: usize,
coef: &[f64],
out: &mut Vec<f64>,
) {
out.clear();
out.resize(nrows, 0.0);
for (r, pred) in out.iter_mut().enumerate() {
let mut s = 0.0;
for c in 0..ncols {
s += design_matrix[c * nrows + r] * coef[c];
}
*pred = s;
}
}
fn aipw_psi(
treatment: &[f64],
outcome: &[f64],
propensity: &[f64],
mu0: &[f64],
mu1: &[f64],
target: &TargetPopulation,
out: &mut Vec<f64>,
) -> Result<(), EstimationError> {
out.clear();
out.reserve(treatment.len());
match target {
TargetPopulation::AllObserved | TargetPopulation::Predicate(_) => {
for (((&t, &y), &e), (&m0, &m1)) in
treatment.iter().zip(outcome).zip(propensity).zip(mu0.iter().zip(mu1))
{
let augmented = (m1 - m0) + (t / e) * (y - m1) - ((1.0 - t) / (1.0 - e)) * (y - m0);
out.push(augmented);
}
}
TargetPopulation::Treated => {
let pi = treatment.iter().filter(|&&t| t > 0.5).count() as f64 / treatment.len() as f64;
if pi <= 0.0 {
return Err(EstimationError::data_msg("ATT requires treated units"));
}
for (((&t, &y), &e), (&m0, &m1)) in
treatment.iter().zip(outcome).zip(propensity).zip(mu0.iter().zip(mu1))
{
let aug = (t / pi) * (m1 - m0) + (t / pi) * (y - m1)
- ((1.0 - t) / pi) * (e / (1.0 - e)) * (y - m0);
out.push(aug);
}
}
TargetPopulation::Untreated => {
let pi0 =
treatment.iter().filter(|&&t| t <= 0.5).count() as f64 / treatment.len() as f64;
if pi0 <= 0.0 {
return Err(EstimationError::data_msg("ATC requires control units"));
}
for (((&t, &y), &e), (&m0, &m1)) in
treatment.iter().zip(outcome).zip(propensity).zip(mu0.iter().zip(mu1))
{
let aug = ((1.0 - t) / pi0) * (m1 - m0) + (t / pi0) * ((1.0 - e) / e) * (y - m1)
- ((1.0 - t) / pi0) * (y - m0);
out.push(aug);
}
}
_ => {
return Err(EstimationError::unsupported(
"AIPW unsupported target population for IF construction",
));
}
}
Ok(())
}
fn residualize_aipw_psi(
psi: &mut [f64],
treatment: &[f64],
propensity: &[f64],
design_colmajor: &[f64],
ncols: usize,
) -> Result<(), EstimationError> {
let n = psi.len();
if ncols == 0 || n < 2 || design_colmajor.len() < n * ncols {
return Ok(());
}
let k = ncols;
let mut scores = vec![0.0; n * k];
for i in 0..n {
let e_resid = treatment[i] - propensity[i];
for c in 0..ncols {
scores[c * n + i] = design_colmajor[c * n + i] * e_resid;
}
}
let nf = n as f64;
let mut gram = vec![0.0; k * k];
let mut rhs = vec![0.0; k];
for c in 0..k {
for i in 0..n {
rhs[c] += scores[c * n + i] * psi[i];
}
rhs[c] /= nf;
for d in 0..k {
let mut acc = 0.0;
for i in 0..n {
acc += scores[c * n + i] * scores[d * n + i];
}
gram[c * k + d] = acc / nf;
}
}
let Some(alpha) = crate::propensity::weighting::solve_symmetric_posdef(&mut gram, &mut rhs, k)
else {
return Err(EstimationError::stats_msg(
"singular AIPW nuisance-score Gram; refusing an uncorrected analytic SE",
));
};
for i in 0..n {
let mut adj = 0.0;
for c in 0..k {
adj += scores[c * n + i] * alpha[c];
}
psi[i] -= adj;
}
Ok(())
}
#[cfg(test)]
#[allow(clippy::many_single_char_names, clippy::float_cmp)]
mod tests {
use std::sync::Arc;
use antecedent_core::{
CausalSchemaBuilder, ExecutionContext, MeasurementSpec, RoleHint, SmallRoleSet,
TargetPopulation, ValueType, VariableId,
};
use antecedent_data::{
Float64Column, OwnedColumn, OwnedColumnarStorage, TabularData, ValidityBitmap,
};
use antecedent_expr::ExprId;
use antecedent_expr::IdentifiedEstimand;
use super::*;
use crate::overlap::OverlapPolicy;
use antecedent_kernels::standard_normal;
fn confounded_scm(n: usize, seed: u64) -> (TabularData, IdentifiedEstimand) {
let (t, y, z) = confounded_columns(n, seed);
build_dataset(t, y, z)
}
fn confounded_columns(n: usize, seed: u64) -> (Vec<f64>, Vec<f64>, Vec<f64>) {
let mut rng = ExecutionContext::for_tests(seed).rng.stream(0x1234_u64);
let mut z = vec![0.0; n];
let mut t = vec![0.0; n];
let mut y = vec![0.0; n];
for i in 0..n {
let zi = standard_normal(&mut rng);
let logit = -0.5 + zi;
let p = 1.0 / (1.0 + (-logit).exp());
let ti = if rng.next_f64() < p { 1.0 } else { 0.0 };
let noise = standard_normal(&mut rng) * 0.5;
z[i] = zi;
t[i] = ti;
y[i] = 2.0 * ti + zi + noise;
}
(t, y, z)
}
fn build_dataset(t: Vec<f64>, y: Vec<f64>, z: Vec<f64>) -> (TabularData, IdentifiedEstimand) {
let n = t.len();
let mut b = CausalSchemaBuilder::new();
b.add_variable(
"t",
ValueType::Continuous,
SmallRoleSet::from_hint(RoleHint::TreatmentCandidate),
None,
None,
MeasurementSpec::default(),
)
.unwrap();
b.add_variable(
"y",
ValueType::Continuous,
SmallRoleSet::from_hint(RoleHint::OutcomeCandidate),
None,
None,
MeasurementSpec::default(),
)
.unwrap();
b.add_variable(
"z",
ValueType::Continuous,
SmallRoleSet::from_hint(RoleHint::Context),
None,
None,
MeasurementSpec::default(),
)
.unwrap();
let schema = b.build().unwrap();
let cols = vec![
OwnedColumn::Float64(
Float64Column::new(
VariableId::from_raw(0),
Arc::from(t),
ValidityBitmap::all_valid(n),
)
.unwrap(),
),
OwnedColumn::Float64(
Float64Column::new(
VariableId::from_raw(1),
Arc::from(y),
ValidityBitmap::all_valid(n),
)
.unwrap(),
),
OwnedColumn::Float64(
Float64Column::new(
VariableId::from_raw(2),
Arc::from(z),
ValidityBitmap::all_valid(n),
)
.unwrap(),
),
];
let storage = OwnedColumnarStorage::try_new(schema, cols, None, None).unwrap();
let estimand = IdentifiedEstimand::backdoor(
"backdoor.adjustment",
Arc::from([VariableId::from_raw(2)]),
ExprId::from_raw(0),
);
(TabularData::new(storage), estimand)
}
fn ctx() -> ExecutionContext {
ExecutionContext::for_tests(7)
}
#[test]
fn aipw_recovers_ate_two() {
let (data, estimand) = confounded_scm(800, 1);
let query =
AverageEffectQuery::binary_ate(VariableId::from_raw(0), VariableId::from_raw(1));
let est = AipwAte { bootstrap_replicates: 30, ..AipwAte::new() };
let prep = est.prepare(&data, &estimand, &query).unwrap();
let mut ws = AipwWorkspace::default();
let effect = est.fit(&prep, &mut ws, &ctx(), AssumptionSet::new()).unwrap();
assert!((effect.ate - 2.0).abs() < 0.3, "ate={}", effect.ate);
assert!(effect.se_bootstrap.is_some());
assert!(effect.overlap_report.is_some());
}
#[test]
fn aipw_rejects_explicit_override() {
let (data, estimand) = confounded_scm(200, 2);
let query =
AverageEffectQuery::binary_ate(VariableId::from_raw(0), VariableId::from_raw(1));
let est = AipwAte { overlap: OverlapPolicy::ExplicitOverride, ..AipwAte::new() };
let err = est.prepare(&data, &estimand, &query).unwrap_err();
assert!(matches!(err, EstimationError::Overlap { .. }));
}
#[test]
fn aipw_recovers_att_two() {
let (data, estimand) = confounded_scm(1_200, 3);
let query =
AverageEffectQuery::binary_ate(VariableId::from_raw(0), VariableId::from_raw(1))
.with_target_population(TargetPopulation::Treated);
let est = AipwAte { bootstrap_replicates: 0, ..AipwAte::new() };
let prep = est.prepare(&data, &estimand, &query).unwrap();
let mut ws = AipwWorkspace::default();
let fit = est.fit(&prep, &mut ws, &ctx(), AssumptionSet::new()).unwrap();
assert!((fit.ate - 2.0).abs() < 0.35, "att={}", fit.ate);
}
#[test]
fn att_if_doubly_robust_under_mu1_misspecification() {
let n = 4_000usize;
let mut rng = ExecutionContext::for_tests(42).rng.stream(0xA11u64);
let mut t = vec![0.0; n];
let mut y = vec![0.0; n];
let mut e = vec![0.0; n];
let mut mu0 = vec![0.0; n];
let mut mu1 = vec![0.0; n];
for i in 0..n {
let z = standard_normal(&mut rng);
let logit = -0.5 + z;
let pi_i = 1.0 / (1.0 + (-logit).exp());
let ti = if rng.next_f64() < pi_i { 1.0 } else { 0.0 };
let noise = standard_normal(&mut rng) * 0.5;
t[i] = ti;
y[i] = 2.0 * ti + z + noise;
e[i] = pi_i;
mu0[i] = z;
mu1[i] = 0.0;
}
let mut psi = Vec::new();
aipw_psi(&t, &y, &e, &mu0, &mu1, &TargetPopulation::Treated, &mut psi).unwrap();
let att = psi.iter().sum::<f64>() / psi.len() as f64;
assert!(
(att - 2.0).abs() < 0.15,
"ATT IF mean under μ₁ misspecification should stay near 2; got {att}"
);
}
#[test]
fn atc_if_deterministic_counterexample() {
let t = vec![1.0, 1.0, 0.0, 0.0];
let y = vec![3.0, 3.0, 1.0, 1.0];
let e = vec![0.5, 0.5, 0.5, 0.5];
let mu0 = vec![0.0, 0.0, 0.0, 0.0];
let mu1 = vec![3.0, 3.0, 3.0, 3.0];
let mut psi = Vec::new();
aipw_psi(&t, &y, &e, &mu0, &mu1, &TargetPopulation::Untreated, &mut psi).unwrap();
let atc = psi.iter().sum::<f64>() / psi.len() as f64;
assert!((atc - 2.0).abs() < 1e-12, "deterministic ATC IF mean should be 2; got {atc}");
}
#[test]
fn atc_if_doubly_robust_under_mu0_misspecification() {
let n = 4_000usize;
let mut rng = ExecutionContext::for_tests(42).rng.stream(0xA7Cu64);
let mut t = vec![0.0; n];
let mut y = vec![0.0; n];
let mut e = vec![0.0; n];
let mut mu0 = vec![0.0; n];
let mut mu1 = vec![0.0; n];
for i in 0..n {
let z = standard_normal(&mut rng);
let logit = -0.5 + z;
let pi_i = 1.0 / (1.0 + (-logit).exp());
let ti = if rng.next_f64() < pi_i { 1.0 } else { 0.0 };
let noise = standard_normal(&mut rng) * 0.5;
t[i] = ti;
y[i] = 2.0 * ti + z + noise;
e[i] = pi_i;
mu0[i] = 0.0;
mu1[i] = 2.0 + z;
}
let mut psi = Vec::new();
aipw_psi(&t, &y, &e, &mu0, &mu1, &TargetPopulation::Untreated, &mut psi).unwrap();
let atc = psi.iter().sum::<f64>() / psi.len() as f64;
assert!(
(atc - 2.0).abs() < 0.15,
"ATC IF mean under μ₀ misspecification should stay near 2; got {atc}"
);
}
#[test]
fn atc_if_doubly_robust_under_propensity_misspecification() {
let n = 4_000usize;
let mut rng = ExecutionContext::for_tests(43).rng.stream(0xBEEFu64);
let mut t = vec![0.0; n];
let mut y = vec![0.0; n];
let mut e = vec![0.0; n];
let mut mu0 = vec![0.0; n];
let mut mu1 = vec![0.0; n];
for i in 0..n {
let z = standard_normal(&mut rng);
let logit = -0.5 + z;
let pi_i = 1.0 / (1.0 + (-logit).exp());
let ti = if rng.next_f64() < pi_i { 1.0 } else { 0.0 };
let noise = standard_normal(&mut rng) * 0.5;
t[i] = ti;
y[i] = 2.0 * ti + z + noise;
e[i] = 0.5;
mu0[i] = z;
mu1[i] = 2.0 + z;
}
let mut psi = Vec::new();
aipw_psi(&t, &y, &e, &mu0, &mu1, &TargetPopulation::Untreated, &mut psi).unwrap();
let atc = psi.iter().sum::<f64>() / psi.len() as f64;
assert!(
(atc - 2.0).abs() < 0.15,
"ATC IF mean under propensity misspecification should stay near 2; got {atc}"
);
}
#[test]
fn atc_if_doubly_robust_when_both_correct() {
let n = 4_000usize;
let mut rng = ExecutionContext::for_tests(44).rng.stream(0xCAFEu64);
let mut t = vec![0.0; n];
let mut y = vec![0.0; n];
let mut e = vec![0.0; n];
let mut mu0 = vec![0.0; n];
let mut mu1 = vec![0.0; n];
for i in 0..n {
let z = standard_normal(&mut rng);
let logit = -0.5 + z;
let pi_i = 1.0 / (1.0 + (-logit).exp());
let ti = if rng.next_f64() < pi_i { 1.0 } else { 0.0 };
let noise = standard_normal(&mut rng) * 0.5;
t[i] = ti;
y[i] = 2.0 * ti + z + noise;
e[i] = pi_i;
mu0[i] = z;
mu1[i] = 2.0 + z;
}
let mut psi = Vec::new();
aipw_psi(&t, &y, &e, &mu0, &mu1, &TargetPopulation::Untreated, &mut psi).unwrap();
let atc = psi.iter().sum::<f64>() / psi.len() as f64;
assert!(
(atc - 2.0).abs() < 0.15,
"ATC IF mean with both models correct should stay near 2; got {atc}"
);
}
#[test]
fn aipw_hc1_multiway_newey_west_finite_se() {
let (data, estimand) = confounded_scm(400, 8);
let query =
AverageEffectQuery::binary_ate(VariableId::from_raw(0), VariableId::from_raw(1));
let n = 400;
let dim_a: Vec<u32> = (0..n).map(|i| u32::try_from(i % 20).unwrap_or(0)).collect();
let dim_b: Vec<u32> = (0..n).map(|i| u32::try_from(i % 15).unwrap_or(0)).collect();
for kind in
[AnalyticSeKind::Hc1, AnalyticSeKind::Multiway, AnalyticSeKind::NeweyWest { lag: 2 }]
{
let est = AipwAte {
bootstrap_replicates: 0,
se_kind: kind,
multiway_ids: Some(vec![dim_a.clone(), dim_b.clone()]),
..AipwAte::new()
};
let prep = est.prepare(&data, &estimand, &query).unwrap();
let mut ws = AipwWorkspace::default();
let fit = est.fit(&prep, &mut ws, &ctx(), AssumptionSet::new()).unwrap();
assert!(fit.se_analytic.is_finite() && fit.se_analytic > 0.0, "kind={kind:?}");
}
}
#[test]
fn aipw_trim_excludes_extreme_propensity_unit() {
let (mut t, mut y, mut z) = confounded_columns(800, 5);
t.push(1.0);
y.push(1000.0);
z.push(-8.0);
let (data, estimand) = build_dataset(t, y, z);
let query =
AverageEffectQuery::binary_ate(VariableId::from_raw(0), VariableId::from_raw(1));
let untrimmed = AipwAte { bootstrap_replicates: 0, ..AipwAte::new() };
let trimmed = AipwAte {
overlap: OverlapPolicy::RequireDiagnostics { clip: Some(0.01), trim: Some(0.02) },
..untrimmed.clone()
};
let mut ws = AipwWorkspace::default();
let prep = untrimmed.prepare(&data, &estimand, &query).unwrap();
let raw = untrimmed.fit(&prep, &mut ws, &ctx(), AssumptionSet::new()).unwrap();
let prep = trimmed.prepare(&data, &estimand, &query).unwrap();
let clean = trimmed.fit(&prep, &mut ws, &ctx(), AssumptionSet::new()).unwrap();
assert!((raw.ate - 2.0).abs() > 1.0, "outlier should distort untrimmed ate={}", raw.ate);
assert!((clean.ate - 2.0).abs() < 0.35, "trimmed ate={}", clean.ate);
let report = clean.overlap_report.as_ref().unwrap();
assert!(report.excluded_fraction > 0.0, "trim must report exclusions");
}
#[test]
fn aipw_works_with_efficient_backdoor_estimand() {
let (data, mut estimand) = confounded_scm(800, 4);
estimand.method = Arc::from("backdoor.efficient");
let query =
AverageEffectQuery::binary_ate(VariableId::from_raw(0), VariableId::from_raw(1));
let est = AipwAte { bootstrap_replicates: 0, ..AipwAte::new() };
let prep = est.prepare(&data, &estimand, &query).unwrap();
let mut ws = AipwWorkspace::default();
let effect = est.fit(&prep, &mut ws, &ctx(), AssumptionSet::new()).unwrap();
assert!((effect.ate - 2.0).abs() < 0.3, "ate={}", effect.ate);
}
}