#![allow(
clippy::cast_precision_loss,
clippy::cast_possible_truncation,
clippy::needless_range_loop,
clippy::too_many_lines,
clippy::many_single_char_names
)]
use std::sync::Arc;
use antecedent_core::{CausalRng, ExecutionContext, KernelPolicy};
use antecedent_estimate::{
BayesianGCompWorkspace, BayesianGComputationAte, CausalPosterior, PreparedBayesianProblem,
};
use antecedent_identify::IdentificationStatus;
use antecedent_kernels::{PosteriorReduceOp, reduce_posterior_draws, standard_normal};
use antecedent_prob::{
ExternalPriorSource, HessianFactorization, PriorSensitivitySummary, PriorSet,
compose_external_priors_with_alphas,
};
use antecedent_stats::GlmFamily;
use crate::common::RefutationReport;
use crate::error::ValidationError;
#[derive(Clone, Debug)]
pub struct PredictiveCheckReport {
pub kind: PredictiveCheckKind,
pub observed: f64,
pub predictive_mean: f64,
pub predictive_sd: f64,
pub p_value: f64,
pub observed_dispersion: f64,
pub predictive_dispersion_mean: f64,
pub dispersion_p_value: f64,
pub n_sims: u32,
}
impl PredictiveCheckReport {
#[must_use]
pub fn to_refutation_report(&self, original_ate: f64, alpha: f64) -> RefutationReport {
let name = match self.kind {
PredictiveCheckKind::Prior => "prior_predictive",
PredictiveCheckKind::Posterior => "posterior_predictive",
};
let mean_ok = self.p_value.is_finite() && self.p_value >= alpha;
let dispersion_ok = self.dispersion_p_value.is_finite() && self.dispersion_p_value >= alpha;
let passed = mean_ok && dispersion_ok;
let comparison = self.p_value.min(self.dispersion_p_value);
RefutationReport {
refuter: Arc::from(name),
original_ate,
refuted_ate: self.predictive_mean,
comparison,
informative: true,
passed,
failure_condition: if passed {
None
} else if !mean_ok && !dispersion_ok {
Some(Arc::from(format!(
"predictive check failed on mean (p={} < alpha={alpha}) and dispersion \
(p={} < alpha={alpha})",
self.p_value, self.dispersion_p_value
)))
} else if !mean_ok {
Some(Arc::from(format!(
"predictive check failed (p={} < alpha={alpha})",
self.p_value
)))
} else {
Some(Arc::from(format!(
"predictive dispersion check failed (p={} < alpha={alpha}); predictive spread \
does not match observed spread even though the mean matches",
self.dispersion_p_value
)))
},
replicates: self.n_sims,
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
pub enum PredictiveCheckKind {
Prior,
Posterior,
}
#[derive(Clone, Debug)]
pub struct PriorPredictiveCheck {
pub n_sims: u32,
pub seed: u64,
pub family: GlmFamily,
}
impl Default for PriorPredictiveCheck {
fn default() -> Self {
Self::new()
}
}
impl PriorPredictiveCheck {
#[must_use]
pub fn new() -> Self {
Self { n_sims: 200, seed: 0, family: GlmFamily::GaussianIdentity }
}
pub fn check(
&self,
problem: &PreparedBayesianProblem,
ctx: &ExecutionContext,
) -> Result<PredictiveCheckReport, ValidationError> {
let p = problem.design.ncols;
let prior = PriorSet::weakly_informative(p);
self.check_with_prior(problem, &prior, ctx)
}
pub fn check_with_prior(
&self,
problem: &PreparedBayesianProblem,
prior: &PriorSet,
ctx: &ExecutionContext,
) -> Result<PredictiveCheckReport, ValidationError> {
let n = problem.design.nrows;
let p = problem.design.ncols;
if n == 0 || p == 0 {
return Err(ValidationError::estimation_msg("empty design for PPC"));
}
let coef_prior = prior.gaussian_coefficients().ok_or_else(|| {
ValidationError::estimation_msg("prior missing Gaussian coefficients for PPC")
})?;
if coef_prior.len() != p {
return Err(ValidationError::estimation_msg(
"prior coefficient dimension mismatch for PPC",
));
}
let mut rng = CausalRng::from_seed(self.seed);
let mut mean_summaries = Vec::with_capacity(self.n_sims as usize);
let mut disp_summaries = Vec::with_capacity(self.n_sims as usize);
let mut beta = vec![0.0; p];
let mut y_pred = vec![0.0; n];
for _ in 0..self.n_sims {
for c in 0..p {
beta[c] =
coef_prior.mean[c] + coef_prior.variance[c].sqrt() * standard_normal(&mut rng);
}
for r in 0..n {
let mut eta = 0.0;
for c in 0..p {
eta += problem.design.matrix[c * n + r] * beta[c];
}
y_pred[r] = self.family.mean_from_eta(eta);
}
push_mean_and_dispersion(
&y_pred,
ctx.kernel_policy,
&mut mean_summaries,
&mut disp_summaries,
);
}
Ok(summarize_predictive_check(
PredictiveCheckKind::Prior,
&problem.design.outcome,
ctx.kernel_policy,
&mean_summaries,
&disp_summaries,
self.n_sims,
))
}
}
#[derive(Clone, Debug)]
pub struct PosteriorPredictiveCheck {
pub n_sims: u32,
pub family: GlmFamily,
}
impl Default for PosteriorPredictiveCheck {
fn default() -> Self {
Self::new()
}
}
impl PosteriorPredictiveCheck {
#[must_use]
pub fn new() -> Self {
Self { n_sims: 200, family: GlmFamily::GaussianIdentity }
}
pub fn check(
&self,
problem: &PreparedBayesianProblem,
posterior: &CausalPosterior,
) -> Result<PredictiveCheckReport, ValidationError> {
let n = problem.design.nrows;
let p = problem.design.ncols;
let n_draws = posterior.draws.n_draws.min(self.n_sims as usize);
if n_draws == 0 {
return Err(ValidationError::estimation_msg("no posterior draws for PPC"));
}
let policy = KernelPolicy::default_policy();
let mut mean_summaries = Vec::with_capacity(n_draws);
let mut disp_summaries = Vec::with_capacity(n_draws);
let mut y_pred = vec![0.0; n];
for d in 0..n_draws {
for r in 0..n {
let mut eta = 0.0;
for c in 0..p {
let x = problem.design.matrix[c * n + r];
let b = posterior.draws.get(d, c).map_err(ValidationError::from)?;
eta += x * b;
}
y_pred[r] = self.family.mean_from_eta(eta);
}
push_mean_and_dispersion(&y_pred, policy, &mut mean_summaries, &mut disp_summaries);
}
Ok(summarize_predictive_check(
PredictiveCheckKind::Posterior,
&problem.design.outcome,
policy,
&mean_summaries,
&disp_summaries,
n_draws as u32,
))
}
}
pub const DEFAULT_MAX_RELATIVE_PRIOR_RANGE: f64 = 0.5;
#[derive(Clone, Debug)]
pub struct PriorSensitivity {
pub scales: Arc<[f64]>,
pub alphas: Arc<[f64]>,
pub max_relative_range: f64,
}
#[derive(Clone, Copy, Debug)]
pub struct ExternalAlphaSensitivity<'a> {
pub sources: &'a [ExternalPriorSource],
pub alphas_applied: &'a [f64],
}
impl Default for PriorSensitivity {
fn default() -> Self {
Self::standard_grid()
}
}
impl PriorSensitivity {
#[must_use]
pub fn standard_grid() -> Self {
Self {
scales: Arc::from(vec![0.5, 1.0, 2.0, 5.0, 10.0, 20.0]),
alphas: Arc::from([]),
max_relative_range: DEFAULT_MAX_RELATIVE_PRIOR_RANGE,
}
}
#[must_use]
pub fn standard_alpha_grid() -> Self {
Self {
scales: Arc::from([]),
alphas: Arc::from(vec![0.0, 0.25, 0.5, 0.75, 1.0]),
max_relative_range: DEFAULT_MAX_RELATIVE_PRIOR_RANGE,
}
}
fn grid_len(&self) -> usize {
if self.alphas.is_empty() { self.scales.len() } else { self.alphas.len() }
}
pub fn evaluate(
&self,
estimator: &BayesianGComputationAte,
problem: &PreparedBayesianProblem,
identification: IdentificationStatus,
workspace: &mut BayesianGCompWorkspace,
ctx: &ExecutionContext,
) -> Result<(PriorSensitivitySummary, Vec<CausalPosterior>), ValidationError> {
if self.scales.is_empty() {
return Err(ValidationError::estimation_msg(
"prior sensitivity scale grid is empty (use evaluate_external_alpha for α mode)",
));
}
let mut means = Vec::with_capacity(self.scales.len());
let mut sds = Vec::with_capacity(self.scales.len());
let mut posts = Vec::with_capacity(self.scales.len());
for &scale in self.scales.iter() {
let est = BayesianGComputationAte {
prior_scale: scale,
n_draws: estimator.n_draws.min(200),
seed: estimator.seed,
backend: estimator.backend,
likelihood: estimator.likelihood,
overlap: estimator.overlap,
prior: None,
};
let post = est.fit(problem, identification, workspace, ctx).map_err(|e| {
ValidationError::estimation_msg(format!("prior sensitivity fit failed: {e}"))
})?;
let eq = post.effect_column().ok_or_else(|| {
ValidationError::estimation_msg("missing effect column in sensitivity fit")
})?;
means.push(post.summaries.mean[eq]);
sds.push(post.summaries.sd[eq]);
posts.push(post);
}
Ok((
PriorSensitivitySummary {
prior_scales: Arc::clone(&self.scales),
alphas: Arc::from([]),
effect_means: Arc::from(means),
effect_sds: Arc::from(sds),
},
posts,
))
}
pub fn evaluate_external_alpha(
&self,
estimator: &BayesianGComputationAte,
problem: &PreparedBayesianProblem,
identification: IdentificationStatus,
workspace: &mut BayesianGCompWorkspace,
ctx: &ExecutionContext,
external: ExternalAlphaSensitivity<'_>,
) -> Result<(PriorSensitivitySummary, Vec<CausalPosterior>), ValidationError> {
if self.alphas.is_empty() {
return Err(ValidationError::estimation_msg("prior sensitivity alpha grid is empty"));
}
if external.sources.len() != external.alphas_applied.len() {
return Err(ValidationError::estimation_msg(
"evaluate_external_alpha: sources / alphas_applied length mismatch",
));
}
let n_coef = problem.design.ncols;
let baseline = PriorSet::weakly_informative(n_coef);
let requested: Vec<f64> = external.sources.iter().map(|s| s.weight.alpha).collect();
let mut means = Vec::with_capacity(self.alphas.len());
let mut sds = Vec::with_capacity(self.alphas.len());
let mut posts = Vec::with_capacity(self.alphas.len());
for &mult in self.alphas.iter() {
if !mult.is_finite() || !(0.0..=1.0).contains(&mult) {
return Err(ValidationError::estimation_msg(
"prior sensitivity alpha multiplier must be finite and in [0, 1]",
));
}
let scaled: Vec<f64> =
external.alphas_applied.iter().map(|&a| (a * mult).clamp(0.0, 1.0)).collect();
let composed = compose_external_priors_with_alphas(
external.sources,
&requested,
&scaled,
&baseline,
)
.map_err(|e| {
ValidationError::estimation_msg(format!("prior sensitivity compose failed: {e}"))
})?;
let est = BayesianGComputationAte {
prior_scale: estimator.prior_scale,
n_draws: estimator.n_draws.min(200),
seed: estimator.seed,
backend: estimator.backend,
likelihood: estimator.likelihood,
overlap: estimator.overlap,
prior: Some(composed.prior),
};
let post = est.fit(problem, identification, workspace, ctx).map_err(|e| {
ValidationError::estimation_msg(format!("prior sensitivity α fit failed: {e}"))
})?;
let eq = post.effect_column().ok_or_else(|| {
ValidationError::estimation_msg("missing effect column in α sensitivity fit")
})?;
means.push(post.summaries.mean[eq]);
sds.push(post.summaries.sd[eq]);
posts.push(post);
}
Ok((
PriorSensitivitySummary {
prior_scales: Arc::from([]),
alphas: Arc::clone(&self.alphas),
effect_means: Arc::from(means),
effect_sds: Arc::from(sds),
},
posts,
))
}
#[must_use]
pub fn to_report(
&self,
summary: &PriorSensitivitySummary,
original_ate: f64,
) -> RefutationReport {
let min = summary.effect_means.iter().copied().fold(f64::INFINITY, f64::min);
let max = summary.effect_means.iter().copied().fold(f64::NEG_INFINITY, f64::max);
let range = max - min;
let denom = summary
.effect_means
.iter()
.copied()
.map(f64::abs)
.fold(original_ate.abs(), f64::max)
.max(1e-8);
let relative = range / denom;
let passed = relative.is_finite() && relative <= self.max_relative_range;
let kind =
if summary.alphas.is_empty() { "prior_sensitivity" } else { "prior_sensitivity_alpha" };
RefutationReport {
refuter: Arc::from(kind),
original_ate,
refuted_ate: summary.effect_means.last().copied().unwrap_or(original_ate),
comparison: relative,
informative: true,
passed,
failure_condition: if passed {
None
} else {
Some(Arc::from(format!(
"prior sensitivity relative range {relative} exceeds max {}",
self.max_relative_range
)))
},
replicates: u32::try_from(self.grid_len()).unwrap_or(u32::MAX),
}
}
}
fn summarize_check(
kind: PredictiveCheckKind,
observed: f64,
summaries: &[f64],
n_sims: u32,
) -> PredictiveCheckReport {
let policy = KernelPolicy::default_policy();
let mean = reduce_posterior_draws(summaries, PosteriorReduceOp::Mean, &policy).unwrap_or(0.0);
let sd = reduce_posterior_draws(summaries, PosteriorReduceOp::Std, &policy).unwrap_or(0.0);
let n = summaries.len() as f64;
let below = summaries.iter().filter(|&&x| x <= observed).count() as f64;
let above = summaries.iter().filter(|&&x| x >= observed).count() as f64;
let p_lower = (1.0 + below) / (1.0 + n);
let p_upper = (1.0 + above) / (1.0 + n);
let p = (2.0 * p_lower.min(p_upper)).min(1.0);
PredictiveCheckReport {
kind,
observed,
predictive_mean: mean,
predictive_sd: sd,
p_value: p,
observed_dispersion: 0.0,
predictive_dispersion_mean: 0.0,
dispersion_p_value: 1.0,
n_sims,
}
}
fn push_mean_and_dispersion(
y_pred: &[f64],
policy: KernelPolicy,
mean_summaries: &mut Vec<f64>,
disp_summaries: &mut Vec<f64>,
) {
let n = y_pred.len().max(1) as f64;
let mean_y = y_pred.iter().sum::<f64>() / n;
let sd_y = reduce_posterior_draws(y_pred, PosteriorReduceOp::Std, &policy).unwrap_or(0.0);
mean_summaries.push(mean_y);
disp_summaries.push(sd_y);
}
fn summarize_predictive_check(
kind: PredictiveCheckKind,
outcome: &[f64],
policy: KernelPolicy,
mean_summaries: &[f64],
disp_summaries: &[f64],
n_sims: u32,
) -> PredictiveCheckReport {
let n = outcome.len().max(1) as f64;
let observed_mean = outcome.iter().sum::<f64>() / n;
let observed_dispersion =
reduce_posterior_draws(outcome, PosteriorReduceOp::Std, &policy).unwrap_or(0.0);
let mean_report = summarize_check(kind, observed_mean, mean_summaries, n_sims);
let disp_report = summarize_check(kind, observed_dispersion, disp_summaries, n_sims);
PredictiveCheckReport {
kind,
observed: mean_report.observed,
predictive_mean: mean_report.predictive_mean,
predictive_sd: mean_report.predictive_sd,
p_value: mean_report.p_value,
observed_dispersion,
predictive_dispersion_mean: disp_report.predictive_mean,
dispersion_p_value: disp_report.p_value,
n_sims,
}
}
#[must_use]
pub fn with_prior_sensitivity(
mut posterior: CausalPosterior,
summary: PriorSensitivitySummary,
) -> CausalPosterior {
posterior.prior_sensitivity = Some(summary);
posterior
}
#[cfg(test)]
mod tests {
use super::*;
use antecedent_core::{
AverageEffectQuery, CausalSchemaBuilder, MeasurementSpec, RoleHint, SmallRoleSet,
ValueType, VariableId,
};
use antecedent_data::{
Float64Column, OwnedColumn, OwnedColumnarStorage, TabularData, ValidityBitmap,
};
use antecedent_estimate::{BayesianBackendKind, BayesianGComputationAte};
use antecedent_expr::{ExprId, IdentifiedEstimand};
use antecedent_identify::IdentificationStatus;
use antecedent_prob::{ExternalPriorWeight, GaussianCoefficientPrior, PriorSpec};
fn toy() -> (TabularData, IdentifiedEstimand, AverageEffectQuery) {
let n = 60usize;
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 t: Vec<f64> = (0..n).map(|i| (i % 2) as f64).collect();
let z: Vec<f64> = (0..n).map(|i| i as f64 * 0.05).collect();
let y: Vec<f64> = (0..n).map(|i| 1.0 + 2.0 * t[i] + 0.3 * z[i]).collect();
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),
);
let query =
AverageEffectQuery::binary_ate(VariableId::from_raw(0), VariableId::from_raw(1));
(TabularData::new(storage), estimand, query)
}
#[test]
fn prior_and_posterior_ppc_run() {
let (data, estimand, query) = toy();
let bayes = BayesianGComputationAte {
backend: BayesianBackendKind::ConjugateGaussian,
n_draws: 100,
seed: 2,
prior_scale: 10.0,
..BayesianGComputationAte::new()
};
let prep = bayes.prepare(&data, &estimand, &query).unwrap();
let ctx = ExecutionContext::for_tests(1);
let prior_rep = PriorPredictiveCheck { n_sims: 50, seed: 3, ..PriorPredictiveCheck::new() }
.check(&prep, &ctx)
.unwrap();
assert_eq!(prior_rep.kind, PredictiveCheckKind::Prior);
assert!(prior_rep.p_value.is_finite());
let mut ws = BayesianGCompWorkspace::default();
let post = bayes
.fit(&prep, IdentificationStatus::NonparametricallyIdentified, &mut ws, &ctx)
.unwrap();
let post_rep = PosteriorPredictiveCheck { n_sims: 50, ..PosteriorPredictiveCheck::new() }
.check(&prep, &post)
.unwrap();
assert_eq!(post_rep.kind, PredictiveCheckKind::Posterior);
}
#[test]
fn summarize_check_observed_outside_range_never_reports_zero() {
let n = 200usize;
let summaries: Vec<f64> = (0..n).map(|i| i as f64 / n as f64).collect(); let min_p = 2.0 / (n as f64 + 1.0);
let low = summarize_check(PredictiveCheckKind::Posterior, -10.0, &summaries, n as u32);
assert!(low.p_value > 0.0, "p_value must be strictly positive, got {}", low.p_value);
assert!(low.p_value >= min_p, "p_value {} below the 2/(n+1) floor {min_p}", low.p_value);
let high = summarize_check(PredictiveCheckKind::Posterior, 10.0, &summaries, n as u32);
assert!(high.p_value > 0.0, "p_value must be strictly positive, got {}", high.p_value);
assert!(high.p_value >= min_p, "p_value {} below the 2/(n+1) floor {min_p}", high.p_value);
}
#[test]
fn summarize_check_observed_near_centre_gives_high_p_value() {
let n = 200usize;
let summaries: Vec<f64> = (0..n).map(|i| i as f64 / n as f64).collect(); let centre = summarize_check(PredictiveCheckKind::Posterior, 0.5, &summaries, n as u32);
assert!(centre.p_value > 0.9, "expected p_value near 1, got {}", centre.p_value);
}
#[test]
fn prior_sensitivity_grid() {
let (data, estimand, query) = toy();
let bayes = BayesianGComputationAte {
backend: BayesianBackendKind::ConjugateGaussian,
n_draws: 80,
seed: 4,
..BayesianGComputationAte::new()
};
let prep = bayes.prepare(&data, &estimand, &query).unwrap();
let mut ws = BayesianGCompWorkspace::default();
let ctx = ExecutionContext::for_tests(1);
let sens = PriorSensitivity {
scales: Arc::from(vec![1.0, 10.0, 50.0]),
alphas: Arc::from([]),
max_relative_range: DEFAULT_MAX_RELATIVE_PRIOR_RANGE,
};
let (summary, posts) = sens
.evaluate(
&bayes,
&prep,
IdentificationStatus::NonparametricallyIdentified,
&mut ws,
&ctx,
)
.unwrap();
assert_eq!(summary.prior_scales.len(), 3);
assert!(summary.alphas.is_empty());
assert_eq!(posts.len(), 3);
let rep =
sens.to_report(&summary, posts[0].summaries.mean[posts[0].effect_column().unwrap()]);
assert!(rep.passed);
}
#[test]
fn prior_sensitivity_external_alpha_pulls_toward_source() {
let fixture: serde_json::Value = serde_json::from_str(include_str!(
"../../../conformance/validate/bayesian_checks/expected.json"
))
.unwrap();
assert!(
fixture["contracts"]["prior_sensitivity_full_trust_moves_toward_source"]
.as_bool()
.unwrap()
);
let (data, estimand, query) = toy();
let bayes = BayesianGComputationAte {
backend: BayesianBackendKind::ConjugateGaussian,
n_draws: 120,
seed: 7,
..BayesianGComputationAte::new()
};
let prep = bayes.prepare(&data, &estimand, &query).unwrap();
let n = prep.design.ncols;
let t_col = prep.design.treatment_column().expect("treatment column");
let mut mean = vec![0.0; n];
mean[t_col] = 8.0;
let mut source_prior = PriorSet::new();
source_prior.push(PriorSpec::GaussianCoefficients(GaussianCoefficientPrior {
mean: Arc::from(mean),
variance: Arc::from(vec![0.05; n]),
}));
let sources = [ExternalPriorSource {
id: Arc::from("survey_a"),
prior: source_prior,
weight: ExternalPriorWeight::power(1.0).unwrap(),
ess: None,
}];
let alphas_applied = [1.0_f64];
let mut ws = BayesianGCompWorkspace::default();
let ctx = ExecutionContext::for_tests(1);
let sens = PriorSensitivity::standard_alpha_grid();
let (summary, _) = sens
.evaluate_external_alpha(
&bayes,
&prep,
IdentificationStatus::NonparametricallyIdentified,
&mut ws,
&ctx,
ExternalAlphaSensitivity { sources: &sources, alphas_applied: &alphas_applied },
)
.unwrap();
assert_eq!(summary.alphas.len(), 5);
assert!(summary.prior_scales.is_empty());
assert!(summary.effect_means.iter().all(|m| m.is_finite()));
let m0 = summary.effect_means[0];
let m1 = *summary.effect_means.last().unwrap();
assert!(
(m1 - 8.0).abs() < (m0 - 8.0).abs(),
"m=1 mean {m1} should be closer to 8 than m=0 mean {m0}"
);
let rep = sens.to_report(&summary, m1);
assert_eq!(rep.refuter.as_ref(), "prior_sensitivity_alpha");
assert!(rep.informative);
assert!(rep.comparison.is_finite() && rep.comparison > 0.0);
}
#[test]
fn ppc_catches_variance_misspecification_mean_ok() {
let n = 400usize;
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 mut rng = CausalRng::from_seed(99);
let t: Vec<f64> = (0..n).map(|_| 0.0).collect();
let z: Vec<f64> = (0..n).map(|_| 0.0).collect();
let y: Vec<f64> = (0..n).map(|_| 5.0 * standard_normal(&mut rng)).collect();
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),
);
let query =
AverageEffectQuery::binary_ate(VariableId::from_raw(0), VariableId::from_raw(1));
let data = TabularData::new(storage);
let bayes = BayesianGComputationAte {
backend: BayesianBackendKind::ConjugateGaussian,
n_draws: 300,
seed: 5,
prior_scale: 0.05,
..BayesianGComputationAte::new()
};
let prep = bayes.prepare(&data, &estimand, &query).unwrap();
let ctx = ExecutionContext::for_tests(1);
let rep = PriorPredictiveCheck { n_sims: 300, seed: 6, ..PriorPredictiveCheck::new() }
.check(&prep, &ctx)
.unwrap();
assert!(
rep.p_value >= 0.05,
"expected the mean axis to look fine (unbiased predictive mean), got p={}",
rep.p_value
);
assert!(
rep.dispersion_p_value < 0.05,
"expected the dispersion axis to catch the 5x variance mismatch, got \
dispersion_p_value={}",
rep.dispersion_p_value
);
let refuted = rep.to_refutation_report(0.0, 0.05);
assert!(
!refuted.passed,
"predictive check must fail overall when dispersion is badly wrong even \
though the mean matches"
);
}
#[test]
fn sbc_to_report_gates_on_uniformity_not_just_mean() {
let n_reps = 200u32;
let n_draws = 100usize;
let ranks: Vec<u32> =
(0..n_reps).map(|i| if i % 2 == 0 { 0 } else { n_draws as u32 }).collect();
let n_d = n_draws as f64;
let fracs: Vec<f64> = ranks.iter().map(|&r| f64::from(r) / n_d).collect();
let mean_rank_frac = fracs.iter().sum::<f64>() / fracs.len() as f64;
assert!(
(0.35..=0.65).contains(&mean_rank_frac),
"fixture sanity: mean rank frac {mean_rank_frac} should look fine on its own"
);
let bins = 10usize;
let mut counts = vec![0.0; bins];
for &r in &ranks {
let b = ((u64::from(r) * bins as u64) / (n_draws as u64)).min(bins as u64 - 1) as usize;
counts[b] += 1.0;
}
let expected = f64::from(n_reps) / bins as f64;
let mut chi2 = 0.0;
for c in counts {
let d = c - expected;
chi2 += d * d / expected.max(1.0);
}
assert!(
chi2 > SBC_CHI2_CRITICAL_9DF_P99,
"fixture sanity: U-shaped ranks should trip the χ² statistic, got {chi2}"
);
let report = SbcReport { ranks: Arc::from(ranks), mean_rank_frac, uniformity_stat: chi2 };
let sbc = SimulationBasedCalibration { n_reps, n_draws, seed: 0 };
let rep = sbc.to_report(&report, 2.0);
assert!(
!rep.passed,
"SBC must fail on a U-shaped (non-uniform) rank distribution even though \
mean_rank_frac={mean_rank_frac:.3} is in [0.35, 0.65]"
);
}
mod calibration_gate {
use super::*;
fn coverage_band(n_reps: u32, level: f64) -> (f64, f64) {
let se = (level * (1.0 - level) / f64::from(n_reps)).sqrt();
let lo = (level - 4.0 * se).max(0.5);
let hi = (level + 4.0 * se).min(1.0);
(lo, hi)
}
#[test]
#[ignore = "calibration: run via scripts/gate_calibration.sh"]
fn sbc_conjugate_gaussian_ranks_are_uniform() {
let (data, estimand, query) = toy();
let bayes = BayesianGComputationAte {
backend: BayesianBackendKind::ConjugateGaussian,
n_draws: 300,
seed: 11,
prior_scale: 5.0,
..BayesianGComputationAte::new()
};
let prep = bayes.prepare(&data, &estimand, &query).unwrap();
let mut ws = BayesianGCompWorkspace::default();
let ctx = ExecutionContext::for_tests(1);
let sbc = SimulationBasedCalibration { n_reps: 200, n_draws: 300, seed: 42 };
let report = sbc
.check(
&bayes,
&prep,
IdentificationStatus::NonparametricallyIdentified,
&mut ws,
&ctx,
)
.unwrap();
let rep = sbc.to_report(&report, 2.0);
assert!(
rep.passed,
"SBC should pass for a correctly specified conjugate Gaussian model: \
mean_rank_frac={:.3} chi2={:.3}",
report.mean_rank_frac, report.uniformity_stat
);
assert!(
(0.35..=0.65).contains(&report.mean_rank_frac),
"mean_rank_frac={:.3} outside [0.35, 0.65]",
report.mean_rank_frac
);
assert!(
report.uniformity_stat < SBC_CHI2_CRITICAL_9DF_P99,
"chi2={:.3} exceeds critical value {SBC_CHI2_CRITICAL_9DF_P99:.3}",
report.uniformity_stat
);
}
#[test]
#[ignore = "calibration: run via scripts/gate_calibration.sh"]
fn posterior_calibration_synthetic_scm_nominal_90_coverage() {
let (data, estimand, query) = toy();
let bayes = BayesianGComputationAte {
backend: BayesianBackendKind::ConjugateGaussian,
n_draws: 300,
seed: 21,
prior_scale: 5.0,
..BayesianGComputationAte::new()
};
let prep = bayes.prepare(&data, &estimand, &query).unwrap();
let mut ws = BayesianGCompWorkspace::default();
let ctx = ExecutionContext::for_tests(1);
let calib = PosteriorCalibrationOnSyntheticScm {
n_reps: 200,
n_draws: 300,
level: 0.9,
seed: 77,
};
let report = calib
.check(
&bayes,
&prep,
IdentificationStatus::NonparametricallyIdentified,
&mut ws,
&ctx,
)
.unwrap();
let (lo, hi) = coverage_band(report.n_reps, 0.9);
assert!(
report.coverage >= lo && report.coverage <= hi,
"nominal 90% credible-interval coverage={:.3} outside [{:.3}, {:.3}] \
({} reps); mean_abs_error={:.3}",
report.coverage,
lo,
hi,
report.n_reps,
report.mean_abs_error
);
}
}
}
#[derive(Clone, Copy, Debug)]
pub struct McmcDiagnosticsCheck {
pub max_rhat: f64,
pub min_ess: f64,
pub max_divergences: u32,
}
impl Default for McmcDiagnosticsCheck {
fn default() -> Self {
Self { max_rhat: 1.05, min_ess: 10.0, max_divergences: u32::MAX / 4 }
}
}
impl McmcDiagnosticsCheck {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn check(&self, posterior: &CausalPosterior) -> Option<RefutationReport> {
let d = &posterior.diagnostics;
if d.factorization != HessianFactorization::Mcmc {
return None;
}
let rhat = d.rhat_max.unwrap_or(f64::INFINITY);
let ess = d.ess_bulk_min.unwrap_or(0.0);
let divs = d.n_divergences.unwrap_or(u32::MAX);
let passed = rhat.is_finite()
&& rhat <= self.max_rhat
&& ess >= self.min_ess
&& divs <= self.max_divergences
&& d.allows_posterior();
let ate = posterior
.effect_column()
.and_then(|c| posterior.summaries.mean.get(c).copied())
.unwrap_or(f64::NAN);
Some(RefutationReport {
refuter: Arc::from("mcmc_diagnostics"),
original_ate: ate,
refuted_ate: ate,
comparison: rhat,
informative: true,
passed,
failure_condition: if passed {
None
} else {
Some(Arc::from(format!(
"MCMC diagnostics failed: rhat={rhat:.4} ess={ess:.1} divergences={divs}"
)))
},
replicates: d.n_chains.unwrap_or(0),
})
}
}
#[derive(Clone, Debug)]
pub struct SimulationBasedCalibration {
pub n_reps: u32,
pub n_draws: usize,
pub seed: u64,
}
impl Default for SimulationBasedCalibration {
fn default() -> Self {
Self { n_reps: 50, n_draws: 100, seed: 0 }
}
}
#[derive(Clone, Debug)]
pub struct SbcReport {
pub ranks: Arc<[u32]>,
pub mean_rank_frac: f64,
pub uniformity_stat: f64,
}
impl SimulationBasedCalibration {
#[must_use]
pub fn new(n_reps: u32) -> Self {
Self { n_reps: n_reps.max(1), ..Self::default() }
}
pub fn check(
&self,
estimator: &BayesianGComputationAte,
problem: &PreparedBayesianProblem,
identification: IdentificationStatus,
workspace: &mut BayesianGCompWorkspace,
ctx: &ExecutionContext,
) -> Result<SbcReport, ValidationError> {
let mut rng = CausalRng::from_seed(self.seed);
let n = problem.design.nrows;
let p = problem.design.ncols;
let t_col = problem
.design
.treatment_column()
.ok_or_else(|| ValidationError::estimation_msg("SBC: missing treatment column"))?;
let mut ranks = Vec::with_capacity(self.n_reps as usize);
let mut est = estimator.clone();
est.n_draws = self.n_draws;
let scale = estimator.prior_scale.max(1e-6);
for rep in 0..self.n_reps {
let mut beta = vec![0.0; p];
for c in 0..p {
beta[c] = scale * standard_normal(&mut rng);
}
let true_effect = (problem.active - problem.control) * beta[t_col];
let mut y_rep = vec![0.0; n];
for r in 0..n {
let mut eta = 0.0;
for c in 0..p {
eta += problem.design.matrix[c * n + r] * beta[c];
}
y_rep[r] = eta + standard_normal(&mut rng);
}
let mut sim_problem = problem.clone();
let mut design = sim_problem.design.clone();
design.outcome = Arc::from(y_rep);
sim_problem.design = design;
est.seed = self.seed ^ (u64::from(rep).wrapping_mul(0x9E37));
let post = est
.fit(&sim_problem, identification, workspace, ctx)
.map_err(|e| ValidationError::estimation_msg(format!("SBC refit failed: {e}")))?;
let col = post
.effect_column()
.ok_or_else(|| ValidationError::estimation_msg("SBC: no effect column"))?;
let draws = post
.draws
.column(col)
.map_err(|e| ValidationError::estimation_msg(format!("SBC draws: {e}")))?;
let mut rank = 0u32;
for &d in draws {
if d < true_effect {
rank += 1;
}
}
ranks.push(rank);
}
let n_d = self.n_draws.max(1) as f64;
let fracs: Vec<f64> = ranks.iter().map(|&r| f64::from(r) / n_d).collect();
let mean_rank_frac =
reduce_posterior_draws(&fracs, PosteriorReduceOp::Mean, &ctx.kernel_policy)
.unwrap_or(0.5);
let bins = 10usize;
let mut counts = vec![0.0; bins];
let n_draws_u = u64::try_from(self.n_draws.max(1)).unwrap_or(1);
let bins_u = u64::try_from(bins).unwrap_or(1);
for &r in &ranks {
let b = usize::try_from(u64::from(r) * bins_u / n_draws_u).unwrap_or(0).min(bins - 1);
counts[b] += 1.0;
}
let expected = f64::from(self.n_reps) / bins as f64;
let mut chi2 = 0.0;
for c in counts {
let d = c - expected;
chi2 += d * d / expected.max(1.0);
}
Ok(SbcReport { ranks: Arc::from(ranks), mean_rank_frac, uniformity_stat: chi2 })
}
#[must_use]
pub fn to_report(&self, report: &SbcReport, original_ate: f64) -> RefutationReport {
let mean_ok = (0.35..=0.65).contains(&report.mean_rank_frac);
let uniform_ok = report.uniformity_stat.is_finite()
&& report.uniformity_stat <= SBC_CHI2_CRITICAL_9DF_P99;
let passed = mean_ok && uniform_ok;
RefutationReport {
refuter: Arc::from("sbc"),
original_ate,
refuted_ate: report.mean_rank_frac,
comparison: report.uniformity_stat,
informative: true,
passed,
failure_condition: if passed {
None
} else if !mean_ok && !uniform_ok {
Some(Arc::from(format!(
"SBC mean rank frac {:.3} outside [0.35, 0.65] and χ²={:.3} exceeds critical \
value {SBC_CHI2_CRITICAL_9DF_P99:.3} (9 df, p=0.99)",
report.mean_rank_frac, report.uniformity_stat
)))
} else if !mean_ok {
Some(Arc::from(format!(
"SBC mean rank frac {:.3} outside [0.35, 0.65]",
report.mean_rank_frac
)))
} else {
Some(Arc::from(format!(
"SBC rank distribution non-uniform: χ²={:.3} exceeds critical value \
{SBC_CHI2_CRITICAL_9DF_P99:.3} (9 df, p=0.99); mean rank frac {:.3} looked \
fine but ranks are not uniformly distributed (U- or M-shaped)",
report.uniformity_stat, report.mean_rank_frac
)))
},
replicates: self.n_reps,
}
}
}
const SBC_CHI2_CRITICAL_9DF_P99: f64 = 21.666;
#[derive(Clone, Debug)]
pub struct PosteriorCalibrationOnSyntheticScm {
pub n_reps: u32,
pub n_draws: usize,
pub level: f64,
pub seed: u64,
}
impl Default for PosteriorCalibrationOnSyntheticScm {
fn default() -> Self {
Self { n_reps: 40, n_draws: 100, level: 0.9, seed: 0 }
}
}
#[derive(Clone, Debug)]
pub struct PosteriorCalibrationReport {
pub coverage: f64,
pub mean_abs_error: f64,
pub n_reps: u32,
}
impl PosteriorCalibrationOnSyntheticScm {
pub fn check(
&self,
estimator: &BayesianGComputationAte,
problem: &PreparedBayesianProblem,
identification: IdentificationStatus,
workspace: &mut BayesianGCompWorkspace,
ctx: &ExecutionContext,
) -> Result<PosteriorCalibrationReport, ValidationError> {
let mut rng = CausalRng::from_seed(self.seed);
let n = problem.design.nrows;
let p = problem.design.ncols;
let t_col = problem
.design
.treatment_column()
.ok_or_else(|| ValidationError::estimation_msg("calibration: missing treatment"))?;
let mut covered = 0u32;
let mut abs_err = 0.0;
let mut est = estimator.clone();
est.n_draws = self.n_draws;
let alpha = ((1.0 - self.level) / 2.0).clamp(0.0, 0.5);
for rep in 0..self.n_reps {
let true_ate = standard_normal(&mut rng);
let mut beta = vec![0.0; p];
let diff = problem.active - problem.control;
beta[t_col] = if diff.abs() > 1e-12 { true_ate / diff } else { true_ate };
for c in 0..p {
if c != t_col {
beta[c] = 0.5 * standard_normal(&mut rng);
}
}
let mut y = vec![0.0; n];
for r in 0..n {
let mut eta = 0.0;
for c in 0..p {
eta += problem.design.matrix[c * n + r] * beta[c];
}
y[r] = eta + standard_normal(&mut rng);
}
let mut sim = problem.clone();
let mut design = sim.design.clone();
design.outcome = Arc::from(y);
sim.design = design;
est.seed = self.seed ^ (u64::from(rep).wrapping_mul(0xC2B2));
let post = est
.fit(&sim, identification, workspace, ctx)
.map_err(|e| ValidationError::estimation_msg(format!("calibration refit: {e}")))?;
let col = post
.effect_column()
.ok_or_else(|| ValidationError::estimation_msg("calibration: no effect"))?;
let mut draws = post.draws.column(col).map_err(ValidationError::from)?.to_vec();
draws.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let lo = quantile_sorted(&draws, alpha);
let hi = quantile_sorted(&draws, 1.0 - alpha);
let mean = reduce_posterior_draws(&draws, PosteriorReduceOp::Mean, &ctx.kernel_policy)
.unwrap_or(0.0);
abs_err += (mean - true_ate).abs();
if true_ate >= lo && true_ate <= hi {
covered += 1;
}
}
Ok(PosteriorCalibrationReport {
coverage: f64::from(covered) / f64::from(self.n_reps.max(1)),
mean_abs_error: abs_err / f64::from(self.n_reps.max(1)),
n_reps: self.n_reps,
})
}
}
fn quantile_sorted(sorted: &[f64], q: f64) -> f64 {
if sorted.is_empty() {
return 0.0;
}
let max_idx = sorted.len() - 1;
let rank = (max_idx as f64 * q.clamp(0.0, 1.0)).round();
let idx = (0..=max_idx)
.min_by(|&a, &b| {
(a as f64 - rank)
.abs()
.partial_cmp(&(b as f64 - rank).abs())
.unwrap_or(std::cmp::Ordering::Equal)
})
.unwrap_or(0);
sorted[idx]
}