#![allow(clippy::cast_precision_loss)]
use std::sync::Arc;
use antecedent_estimate::OverlapPolicy;
use antecedent_stats::GlmOptions;
use crate::common::{RefutationProblem, RefutationReport};
use crate::error::ValidationError;
#[derive(Clone, Debug)]
pub struct OverlapRefuter {
pub eps: f64,
pub min_ess_fraction: f64,
pub glm_options: GlmOptions,
}
impl Default for OverlapRefuter {
fn default() -> Self {
Self::new()
}
}
impl OverlapRefuter {
#[must_use]
pub fn new() -> Self {
Self { eps: 0.05, min_ess_fraction: 0.5, glm_options: GlmOptions::default() }
}
pub fn refute(
&self,
problem: &RefutationProblem<'_>,
) -> Result<RefutationReport, ValidationError> {
let mut local = antecedent_stats::PropensityWorkspace::default();
self.refute_with_propensity(problem, &mut local)
}
pub fn refute_with_propensity(
&self,
problem: &RefutationProblem<'_>,
propensity: &mut antecedent_stats::PropensityWorkspace,
) -> Result<RefutationReport, ValidationError> {
let (report, replicates) = match &problem.original.overlap_report {
Some(r) => (r.clone(), 0),
None => (
crate::common::diagnostic_overlap_report_with(
problem,
&self.glm_options,
OverlapPolicy::require_diagnostics(),
propensity,
)?,
1,
),
};
let nrows = estimation_row_count(problem)? as f64;
let Some(ess) = report.ess else {
return Ok(RefutationReport {
refuter: Arc::from("overlap.assessment"),
original_ate: problem.original.ate,
refuted_ate: problem.original.ate,
comparison: f64::NAN,
informative: false,
passed: false,
failure_condition: Some(Arc::from(
"overlap report has no weights; ESS is undefined",
)),
replicates,
});
};
let ess_fraction = if nrows > 0.0 { ess / nrows } else { 0.0 };
let bounds_ok =
report.propensity_min >= self.eps && report.propensity_max <= 1.0 - self.eps;
let ess_ok = ess_fraction >= self.min_ess_fraction;
let passed = bounds_ok && ess_ok;
let comparison = 1.0 - ess_fraction;
Ok(RefutationReport {
refuter: Arc::from("overlap.assessment"),
original_ate: problem.original.ate,
refuted_ate: problem.original.ate,
comparison,
informative: true,
passed,
failure_condition: if passed {
None
} else {
Some(Arc::from(format!(
"propensity range [{}, {}] or ess_fraction={ess_fraction} failed eps={} / \
min_ess_fraction={}",
report.propensity_min, report.propensity_max, self.eps, self.min_ess_fraction
)))
},
replicates,
})
}
}
fn estimation_row_count(problem: &RefutationProblem<'_>) -> Result<usize, ValidationError> {
let mut ids = vec![problem.treatment(), problem.outcome()];
ids.extend_from_slice(&problem.estimand.adjustment_set);
let mask = problem.data.complete_case_mask(&ids).map_err(ValidationError::from)?;
Ok(mask.iter().filter(|&&k| k).count())
}