#![allow(clippy::cast_precision_loss)]
use std::sync::Arc;
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum OverlapPolicy {
ExplicitOverride,
RequireDiagnostics {
clip: Option<f64>,
trim: Option<f64>,
},
}
impl OverlapPolicy {
#[must_use]
pub const fn require_diagnostics() -> Self {
Self::RequireDiagnostics { clip: None, trim: None }
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct PropensityInterval {
pub low: f64,
pub high: f64,
}
#[derive(Clone, Debug, PartialEq)]
pub struct ClipSensitivity {
pub thresholds: Arc<[f64]>,
pub ess: Arc<[f64]>,
pub extreme_weight_counts: Arc<[u32]>,
}
#[derive(Clone, Debug, PartialEq)]
pub struct OverlapReport {
pub propensity_min: f64,
pub propensity_max: f64,
pub ess: Option<f64>,
pub extreme_weight_count: u32,
pub excluded_fraction: f64,
pub target_population_support: f64,
pub excluded_regions: Arc<[PropensityInterval]>,
pub clip: Option<f64>,
pub trim: Option<f64>,
pub clip_sensitivity: Option<ClipSensitivity>,
}
impl OverlapReport {
#[must_use]
pub fn from_propensities(
propensities: &[f64],
weights: Option<&[f64]>,
policy: OverlapPolicy,
) -> Self {
let (clip, trim) = match policy {
OverlapPolicy::ExplicitOverride => (None, None),
OverlapPolicy::RequireDiagnostics { clip, trim } => (clip, trim),
};
let mut min_p = f64::INFINITY;
let mut max_p = f64::NEG_INFINITY;
let mut excluded = 0u32;
let mut in_support = 0u32;
let support_lo = clip.or(trim).unwrap_or(0.0);
let support_hi = 1.0 - support_lo;
for &p in propensities {
min_p = min_p.min(p);
max_p = max_p.max(p);
if let Some(t) = trim {
if p < t || p > 1.0 - t {
excluded = excluded.saturating_add(1);
}
}
if p >= support_lo && p <= support_hi {
in_support = in_support.saturating_add(1);
}
}
if propensities.is_empty() {
min_p = f64::NAN;
max_p = f64::NAN;
}
let n = propensities.len().max(1) as f64;
let excluded_fraction = f64::from(excluded) / n;
let target_population_support =
if propensities.is_empty() { f64::NAN } else { f64::from(in_support) / n };
let excluded_regions: Arc<[PropensityInterval]> = match trim {
Some(t) if t > 0.0 => Arc::from([
PropensityInterval { low: 0.0, high: t },
PropensityInterval { low: 1.0 - t, high: 1.0 },
]),
_ => Arc::from([]),
};
let (ess, extreme_weight_count) = match weights {
Some(w) if !w.is_empty() => {
let (e, c) = weight_summary(w);
(Some(e), c)
}
_ => (None, 0),
};
let clip_sensitivity = clip.map(|c| clip_sensitivity_grid(propensities, c));
Self {
propensity_min: min_p,
propensity_max: max_p,
ess,
extreme_weight_count,
excluded_fraction,
target_population_support,
excluded_regions,
clip,
trim,
clip_sensitivity,
}
}
}
fn weight_summary(weights: &[f64]) -> (f64, u32) {
let sum: f64 = weights.iter().sum();
let sum_sq: f64 = weights.iter().map(|x| x * x).sum();
let ess = if sum_sq > 0.0 { (sum * sum) / sum_sq } else { 0.0 };
let extreme = weights.iter().filter(|&&x| x > 10.0).count();
(ess, u32::try_from(extreme).unwrap_or(u32::MAX))
}
fn ate_ipw_weights_from_propensity(propensities: &[f64], clip: f64) -> Vec<f64> {
let lo = clip.clamp(1e-6, 0.49);
let hi = 1.0 - lo;
propensities
.iter()
.map(|&p_raw| {
let p = p_raw.clamp(lo, hi);
1.0 / p + 1.0 / (1.0 - p)
})
.collect()
}
fn clip_sensitivity_grid(propensities: &[f64], clip: f64) -> ClipSensitivity {
let c = clip.clamp(1e-6, 0.49);
let candidates = [c * 0.5, c, (c * 2.0).min(0.49)];
let mut thresholds = Vec::with_capacity(3);
let mut ess_vals = Vec::with_capacity(3);
let mut extreme_counts = Vec::with_capacity(3);
for &thr in &candidates {
if thresholds.last().is_some_and(|&prev: &f64| (prev - thr).abs() < 1e-15) {
continue;
}
let rebuilt = ate_ipw_weights_from_propensity(propensities, thr);
let (ess, extreme) = weight_summary(&rebuilt);
thresholds.push(thr);
ess_vals.push(ess);
extreme_counts.push(extreme);
}
ClipSensitivity {
thresholds: Arc::from(thresholds),
ess: Arc::from(ess_vals),
extreme_weight_counts: Arc::from(extreme_counts),
}
}