#![allow(
clippy::similar_names,
clippy::too_many_lines,
clippy::doc_markdown,
clippy::too_many_arguments,
clippy::cast_precision_loss
)]
use std::sync::Arc;
use antecedent_core::{
AssumptionSet, AverageEffectQuery, BufferMaterialization, Diagnostic, DiagnosticKind,
DiagnosticSeverity, ExecutionContext, ExecutionPerformanceRecord, Intervention,
InterventionSequence, LogicalAnalysisPlanRecord, PhysicalExecutionPlanRecord, ProvenanceGraph,
ProvenanceNode, SequencedIntervention, VERSION, VariableId,
};
use antecedent_data::{
IdRemap, MultiEnvironmentData, TableView, TabularData, TimeSeriesData, dedupe_variable_ids,
};
use antecedent_estimate::{CausalPosterior, EffectEstimate, EstimationWorkspace, OverlapPolicy};
use antecedent_expr::{IdentifiedEstimand, RdDesignParams};
use antecedent_graph::{CpdagReview, TemporalCpdagReview, TemporalGraphReview};
use antecedent_validate::{RefutationProblem, RefutationReport, ValidationSuite};
use crate::discovery::{
DiscoverParams, StaticDiscoverParams, discover_fci, discover_ges, discover_jpcmci_plus,
discover_lingam, discover_lpcmci, discover_notears, discover_pc, discover_pcmci,
discover_pcmci_plus, discover_rfci, discover_rpcmci,
};
use crate::discovery_defaults::resolve_ci;
use crate::error::CausalError;
use crate::result::CausalAnalysisResult;
use antecedent_discovery::{MultiDatasetConstraints, RegimeAssignment};
use super::builder::RefuteSuite;
pub(crate) struct AssembleArgs<'a> {
pub(crate) logical: &'a LogicalAnalysisPlanRecord,
pub(crate) physical: &'a PhysicalExecutionPlanRecord,
pub(crate) identification: antecedent_identify::IdentificationResult,
pub(crate) estimand: IdentifiedEstimand,
pub(crate) estimate: EffectEstimate,
pub(crate) distribution: Option<antecedent_estimate::InterventionalDistributionEstimate>,
pub(crate) posterior: Option<antecedent_estimate::CausalPosterior>,
pub(crate) mediation: Option<antecedent_estimate::TemporalMediationEstimate>,
pub(crate) counterfactual: Option<crate::gcm::IteResult>,
pub(crate) anomaly: Option<Vec<antecedent_attribution::AnomalyScores>>,
pub(crate) change_attribution: Option<antecedent_attribution::ChangeAttributionResult>,
pub(crate) mechanism_change: Option<Vec<antecedent_attribution::MechanismChangeDetection>>,
pub(crate) unit_change: Option<antecedent_attribution::UnitChangeResult>,
pub(crate) refutations: Vec<RefutationReport>,
pub(crate) diagnostics: Vec<Diagnostic>,
pub(crate) provenance: ProvenanceGraph,
pub(crate) treatment: VariableId,
pub(crate) outcome: VariableId,
pub(crate) wall_time_ns: u64,
pub(crate) latency_mode: Option<Arc<str>>,
pub(crate) stage_timings_ns: Vec<(Arc<str>, u64)>,
pub(crate) bootstrap_replicates_requested: Option<u32>,
pub(crate) bootstrap_replicates_ok: Option<u32>,
pub(crate) n_draws: Option<u32>,
pub(crate) cancelled: bool,
pub(crate) early_stopped: bool,
}
pub(crate) fn assemble_result(args: AssembleArgs<'_>) -> CausalAnalysisResult {
let copy_count = args
.physical
.materializations
.iter()
.filter(|(_, m)| !matches!(m, BufferMaterialization::Borrowed))
.count() as u64;
CausalAnalysisResult {
logical_plan: args.logical.clone(),
physical_plan: args.physical.clone(),
identification: args.identification,
estimand: args.estimand,
estimate: args.estimate,
distribution: args.distribution,
posterior: args.posterior,
mediation: args.mediation,
counterfactual: args.counterfactual,
anomaly: args.anomaly,
change_attribution: args.change_attribution,
mechanism_change: args.mechanism_change,
unit_change: args.unit_change,
refutations: args.refutations,
predictive_checks: Vec::new(),
diagnostics: args.diagnostics,
provenance: args.provenance,
performance: ExecutionPerformanceRecord {
wall_time_ns: Some(args.wall_time_ns),
peak_rss_bytes: None,
copy_count,
scalar_fallback_count: 0,
latency_mode: args.latency_mode,
stage_timings_ns: args.stage_timings_ns,
bootstrap_replicates_requested: args.bootstrap_replicates_requested,
bootstrap_replicates_ok: args.bootstrap_replicates_ok,
n_draws: args.n_draws,
cancelled: args.cancelled,
early_stopped: args.early_stopped,
},
treatment: args.treatment,
outcome: args.outcome,
}
}
pub(crate) type ProvStep<'a> = (&'a str, &'a str, &'a [&'a str], &'a AssumptionSet);
pub(crate) fn provenance_pair(first: ProvStep<'_>, second: ProvStep<'_>) -> ProvenanceGraph {
let mut provenance = ProvenanceGraph::new();
for (artifact_id, operation, parents, assumptions) in [first, second] {
let parent_arcs: Arc<[Arc<str>]> =
parents.iter().map(|p| Arc::<str>::from(*p)).collect::<Vec<_>>().into();
provenance.push(ProvenanceNode {
artifact_id: Arc::from(artifact_id),
operation: Arc::from(operation),
parents: parent_arcs,
assumptions: assumptions.clone(),
library_version: Arc::from(VERSION),
config_digest: Some(Arc::from("temporal")),
});
}
provenance
}
pub(crate) fn run_pcmci_review(
data: &TimeSeriesData,
max_lag: u32,
alpha: f64,
fdr: Option<antecedent_stats::FdrAdjustment>,
ci: Arc<dyn antecedent_stats::ConditionalIndependence + Send + Sync>,
ctx: &ExecutionContext,
) -> Result<TemporalGraphReview, CausalError> {
let vars: Vec<VariableId> = data.schema().variables().iter().map(|v| v.id).collect();
let params = DiscoverParams {
max_lag,
alpha,
fdr,
ci,
multi_dataset: MultiDatasetConstraints::default(),
};
let result = discover_pcmci(data, &vars, ¶ms, ctx)?;
Ok(result.review)
}
pub(crate) fn run_pcmci_plus_review(
data: &TimeSeriesData,
max_lag: u32,
alpha: f64,
fdr: Option<antecedent_stats::FdrAdjustment>,
ci: Arc<dyn antecedent_stats::ConditionalIndependence + Send + Sync>,
ctx: &ExecutionContext,
) -> Result<TemporalCpdagReview, CausalError> {
let vars: Vec<VariableId> = data.schema().variables().iter().map(|v| v.id).collect();
let params = DiscoverParams {
max_lag,
alpha,
fdr,
ci,
multi_dataset: MultiDatasetConstraints::default(),
};
let result = discover_pcmci_plus(data, &vars, ¶ms, ctx)?;
Ok(result.review)
}
pub(crate) fn run_jpcmci_plus_review(
data: &MultiEnvironmentData,
max_lag: u32,
alpha: f64,
fdr: Option<antecedent_stats::FdrAdjustment>,
multi_dataset: &MultiDatasetConstraints,
ci: Arc<dyn antecedent_stats::ConditionalIndependence + Send + Sync>,
ctx: &ExecutionContext,
) -> Result<TemporalCpdagReview, CausalError> {
let vars: Vec<VariableId> = data.schema().variables().iter().map(|v| v.id).collect();
let system: Vec<VariableId> =
vars.into_iter().filter(|v| !multi_dataset.is_context(*v)).collect();
if system.is_empty() {
return Err(CausalError::Compile {
message: "jpcmci+ needs ≥1 system variable after excluding context_variables".into(),
});
}
let params = DiscoverParams { max_lag, alpha, fdr, ci, multi_dataset: multi_dataset.clone() };
let result = discover_jpcmci_plus(data, &system, ¶ms, ctx)?;
Ok(result.review)
}
pub(crate) fn run_rpcmci_discovery(
data: &TimeSeriesData,
max_lag: u32,
alpha: f64,
fdr: Option<antecedent_stats::FdrAdjustment>,
assignment: &RegimeAssignment,
ci: Arc<dyn antecedent_stats::ConditionalIndependence + Send + Sync>,
ctx: &ExecutionContext,
) -> Result<antecedent_discovery::RpcmciDiscoveryResult, CausalError> {
let vars: Vec<VariableId> = data.schema().variables().iter().map(|v| v.id).collect();
let params = DiscoverParams {
max_lag,
alpha,
fdr,
ci,
multi_dataset: MultiDatasetConstraints::default(),
};
if assignment.len() != data.row_count() {
return Err(CausalError::Compile {
message: format!(
"RPCMCI regime_assignment length {} != series length {}",
assignment.len(),
data.row_count()
),
});
}
discover_rpcmci(data, &vars, assignment, ¶ms, None, ctx)
}
pub(crate) fn run_lpcmci_review(
data: &TimeSeriesData,
max_lag: u32,
alpha: f64,
fdr: Option<antecedent_stats::FdrAdjustment>,
ci: Arc<dyn antecedent_stats::ConditionalIndependence + Send + Sync>,
ctx: &ExecutionContext,
) -> Result<antecedent_graph::TemporalPagReview, CausalError> {
let vars: Vec<VariableId> = data.schema().variables().iter().map(|v| v.id).collect();
let params = DiscoverParams {
max_lag,
alpha,
fdr,
ci,
multi_dataset: MultiDatasetConstraints::default(),
};
let result = discover_lpcmci(data, &vars, ¶ms, ctx)?;
Ok(result.review)
}
pub(crate) fn run_pc_review(
data: &TabularData,
alpha: f64,
max_cond_size: usize,
fdr: Option<antecedent_stats::FdrAdjustment>,
ci: Arc<dyn antecedent_stats::ConditionalIndependence + Send + Sync>,
ctx: &ExecutionContext,
) -> Result<CpdagReview, CausalError> {
let vars: Vec<VariableId> = data.schema().variables().iter().map(|v| v.id).collect();
let params =
StaticDiscoverParams { alpha, max_cond_size, fdr, ci, screen_pc: false, max_subset: None };
let result = discover_pc(data, &vars, ¶ms, ctx)?;
Ok(result.review)
}
pub(crate) fn run_ges_review(
data: &TabularData,
alpha: f64,
max_cond_size: usize,
fdr: Option<antecedent_stats::FdrAdjustment>,
ci: Arc<dyn antecedent_stats::ConditionalIndependence + Send + Sync>,
ctx: &ExecutionContext,
) -> Result<CpdagReview, CausalError> {
let vars: Vec<VariableId> = data.schema().variables().iter().map(|v| v.id).collect();
let params =
StaticDiscoverParams { alpha, max_cond_size, fdr, ci, screen_pc: false, max_subset: None };
let result = discover_ges(data, &vars, ¶ms, ctx)?;
Ok(result.review)
}
pub(crate) fn run_lingam_review(
data: &TabularData,
max_cond_size: usize,
prune_threshold: f64,
ctx: &ExecutionContext,
) -> Result<antecedent_graph::DagReview, CausalError> {
let vars: Vec<VariableId> = data.schema().variables().iter().map(|v| v.id).collect();
let params = StaticDiscoverParams {
alpha: 0.05,
max_cond_size,
fdr: None,
ci: resolve_ci("parcorr", None)?,
screen_pc: false,
max_subset: None,
};
let result = discover_lingam(data, &vars, ¶ms, prune_threshold, ctx)?;
Ok(result.review)
}
pub(crate) fn run_notears_review(
data: &TabularData,
max_cond_size: usize,
lambda: f64,
threshold: f64,
standardize: bool,
ctx: &ExecutionContext,
) -> Result<antecedent_graph::DagReview, CausalError> {
let vars: Vec<VariableId> = data.schema().variables().iter().map(|v| v.id).collect();
let params = StaticDiscoverParams {
alpha: 0.05,
max_cond_size,
fdr: None,
ci: resolve_ci("parcorr", None)?,
screen_pc: false,
max_subset: None,
};
let result = discover_notears(data, &vars, ¶ms, lambda, threshold, standardize, ctx)?;
Ok(result.discovery.review)
}
pub(crate) fn run_fci_review(
data: &TabularData,
alpha: f64,
max_cond_size: usize,
fdr: Option<antecedent_stats::FdrAdjustment>,
ci: Arc<dyn antecedent_stats::ConditionalIndependence + Send + Sync>,
ctx: &ExecutionContext,
) -> Result<antecedent_graph::PagReview, CausalError> {
let vars: Vec<VariableId> = data.schema().variables().iter().map(|v| v.id).collect();
let params =
StaticDiscoverParams { alpha, max_cond_size, fdr, ci, screen_pc: false, max_subset: None };
let result = discover_fci(data, &vars, ¶ms, ctx)?;
Ok(result.review)
}
pub(crate) fn run_rfci_review(
data: &TabularData,
alpha: f64,
max_cond_size: usize,
fdr: Option<antecedent_stats::FdrAdjustment>,
ci: Arc<dyn antecedent_stats::ConditionalIndependence + Send + Sync>,
ctx: &ExecutionContext,
) -> Result<antecedent_graph::PagReview, CausalError> {
let vars: Vec<VariableId> = data.schema().variables().iter().map(|v| v.id).collect();
let params =
StaticDiscoverParams { alpha, max_cond_size, fdr, ci, screen_pc: false, max_subset: None };
let result = discover_rfci(data, &vars, ¶ms, ctx)?;
Ok(result.review)
}
pub(crate) fn run_refuters(
data: &TabularData,
estimand: &IdentifiedEstimand,
query: &AverageEffectQuery,
estimate: &EffectEstimate,
workspace: &mut EstimationWorkspace,
propensity: Option<&mut antecedent_stats::PropensityWorkspace>,
ctx: &ExecutionContext,
suite: RefuteSuite,
estimator: &str,
custom: &[Arc<dyn antecedent_validate::CustomEffectValidator>],
temporal: Option<antecedent_validate::TemporalRefitContext<'_>>,
) -> Result<Vec<RefutationReport>, CausalError> {
let problem = RefutationProblem {
data,
estimand,
query,
original: estimate,
estimator: Some(estimator),
temporal,
};
let mut validation = match suite {
RefuteSuite::None => {
if custom.is_empty() {
return Ok(Vec::new());
}
ValidationSuite::new()
}
RefuteSuite::Cheap => ValidationSuite::overlap_and_evalue(),
RefuteSuite::PlaceboAndRcc => ValidationSuite::placebo_and_rcc(),
RefuteSuite::Full => ValidationSuite::full_effect(),
};
for v in custom {
validation = validation.with_custom(Arc::clone(v));
}
let outcomes = match propensity {
Some(pws) => validation
.run_with_propensity(&problem, workspace, pws, ctx)
.map_err(CausalError::from)?,
None => validation.run(&problem, workspace, ctx).map_err(CausalError::from)?,
};
Ok(ValidationSuite::reports_only(&outcomes))
}
pub(crate) fn resolve_analysis_ci(
discovery_ci: Option<&Arc<dyn antecedent_stats::ConditionalIndependence + Send + Sync>>,
) -> Result<Arc<dyn antecedent_stats::ConditionalIndependence + Send + Sync>, CausalError> {
match discovery_ci {
Some(ci) => Ok(Arc::clone(ci)),
None => resolve_ci("parcorr", None),
}
}
pub(crate) fn effect_from_posterior(
posterior: &CausalPosterior,
) -> Result<EffectEstimate, CausalError> {
let eq = posterior.effect_column().ok_or_else(|| CausalError::Compile {
message: "Bayesian posterior missing effect column".into(),
})?;
let ate = posterior.summaries.mean[eq];
let se = posterior.summaries.sd[eq];
Ok(EffectEstimate::new(ate, se, posterior.assumptions.clone(), OverlapPolicy::ExplicitOverride))
}
pub(crate) fn overlap_diagnostic(overlap: OverlapPolicy) -> Diagnostic {
match overlap {
OverlapPolicy::ExplicitOverride => Diagnostic::new(
"estimate.overlap.explicit_override",
DiagnosticKind::Scientific,
DiagnosticSeverity::Info,
"estimator used ExplicitOverride for positivity (not a propensity-based method)",
),
OverlapPolicy::RequireDiagnostics { .. } => Diagnostic::new(
"estimate.overlap.require_diagnostics",
DiagnosticKind::Scientific,
DiagnosticSeverity::Info,
"estimator used RequireDiagnostics for mandatory positivity diagnostics",
),
}
}
pub(crate) fn push_conflict_diagnostics(
diagnostics: &mut Vec<Diagnostic>,
summary: &antecedent_prob::ConflictSummary,
) {
for (i, id) in summary.source_ids.iter().enumerate() {
let req = summary.alphas_requested.get(i).copied().unwrap_or(f64::NAN);
let app = summary.alphas_applied.get(i).copied().unwrap_or(f64::NAN);
let p = summary
.p_values
.get(i)
.and_then(|x| *x)
.map_or_else(|| "none".to_string(), |v| format!("{v}"));
let kl = summary
.kl_values
.get(i)
.and_then(|x| *x)
.map_or_else(|| "none".to_string(), |v| format!("{v}"));
let mut d = Diagnostic::new(
"bayes.prior_bank.conflict",
DiagnosticKind::Scientific,
DiagnosticSeverity::Info,
format!(
"external prior {id}: alpha_requested={req}, alpha_applied={app}, p={p}, kl={kl}"
),
);
d.fields = Arc::from([
(Arc::from("source_id"), Arc::clone(id)),
(Arc::from("alpha_requested"), Arc::from(format!("{req}"))),
(Arc::from("alpha_applied"), Arc::from(format!("{app}"))),
]);
diagnostics.push(d);
}
}
pub(crate) fn columns_for_ate_estimand(
query: &AverageEffectQuery,
estimand: &IdentifiedEstimand,
) -> Vec<VariableId> {
dedupe_variable_ids(
std::iter::once(query.treatment)
.chain(std::iter::once(query.outcome))
.chain(query.effect_modifiers.iter().copied())
.chain(estimand.adjustment_set.iter().copied())
.chain(estimand.instruments.iter().copied())
.chain(estimand.mediators.iter().copied())
.chain(estimand.rd_design.map(|rd| rd.running_variable)),
)
}
pub(crate) fn project_for_ate_estimate(
data: &TabularData,
query: &AverageEffectQuery,
estimand: &IdentifiedEstimand,
) -> Result<(TabularData, AverageEffectQuery, IdentifiedEstimand), CausalError> {
let ids = columns_for_ate_estimand(query, estimand);
if ids.len() == data.schema().len() {
return Ok((data.clone(), query.clone(), estimand.clone()));
}
let (projected, remap) = data.project(&ids)?;
let query_p = remap_average_effect_query(query, &remap)?;
let estimand_p = remap_identified_estimand(estimand, &remap)?;
Ok((projected, query_p, estimand_p))
}
fn remap_variable_slice(
ids: &[VariableId],
remap: &IdRemap,
) -> Result<Arc<[VariableId]>, CausalError> {
let mapped: Result<Vec<_>, _> = ids.iter().map(|id| remap.map(*id)).collect();
Ok(Arc::from(mapped?))
}
fn remap_intervention(
intervention: &Intervention,
remap: &IdRemap,
) -> Result<Intervention, CausalError> {
match intervention {
Intervention::Set { variable, value } => {
Ok(Intervention::Set { variable: remap.map(*variable)?, value: value.clone() })
}
Intervention::Shift { variable, delta } => {
Ok(Intervention::Shift { variable: remap.map(*variable)?, delta: delta.clone() })
}
Intervention::Stochastic { variable, policy } => {
Ok(Intervention::Stochastic { variable: remap.map(*variable)?, policy: policy.clone() })
}
Intervention::Soft { variable, mechanism } => {
Ok(Intervention::Soft { variable: remap.map(*variable)?, mechanism: mechanism.clone() })
}
Intervention::Sequence(seq) => {
let steps: Result<Vec<_>, CausalError> = seq
.steps
.iter()
.map(|s| {
Ok(SequencedIntervention {
intervention: remap_intervention(&s.intervention, remap)?,
temporal: s.temporal.clone(),
})
})
.collect();
Ok(Intervention::Sequence(InterventionSequence::new(steps?)))
}
other => Err(CausalError::Compile {
message: format!("cannot remap unsupported intervention variant: {other:?}"),
}),
}
}
fn remap_average_effect_query(
query: &AverageEffectQuery,
remap: &IdRemap,
) -> Result<AverageEffectQuery, CausalError> {
Ok(AverageEffectQuery::new(
remap.map(query.treatment)?,
remap.map(query.outcome)?,
remap_variable_slice(&query.effect_modifiers, remap)?,
remap_intervention(&query.control, remap)?,
remap_intervention(&query.active, remap)?,
query.target_population.clone(),
))
}
fn remap_identified_estimand(
estimand: &IdentifiedEstimand,
remap: &IdRemap,
) -> Result<IdentifiedEstimand, CausalError> {
let rd_design = match &estimand.rd_design {
None => None,
Some(rd) => {
Some(RdDesignParams::new(remap.map(rd.running_variable)?, rd.cutoff, rd.bandwidth))
}
};
Ok(IdentifiedEstimand::new(
Arc::clone(&estimand.method),
remap_variable_slice(&estimand.adjustment_set, remap)?,
remap_variable_slice(&estimand.instruments, remap)?,
remap_variable_slice(&estimand.mediators, remap)?,
estimand.functional,
rd_design,
))
}
pub(crate) fn projection_diagnostic(full_cols: usize, projected_cols: usize) -> Option<Diagnostic> {
if projected_cols >= full_cols {
return None;
}
Some(Diagnostic::new(
"exec.project.columns",
DiagnosticKind::Execution,
DiagnosticSeverity::Info,
format!("projected {full_cols} → {projected_cols} columns after identification"),
))
}
pub(crate) fn evaluate_bayesian_prior_sensitivity(
cfg: &crate::inference::BayesianConfig,
est: &antecedent_estimate::BayesianGComputationAte,
prep: &antecedent_estimate::PreparedBayesianProblem,
status: antecedent_identify::IdentificationStatus,
posterior: &CausalPosterior,
ws: &mut antecedent_estimate::BayesianGCompWorkspace,
ctx: &ExecutionContext,
) -> Result<
(antecedent_prob::PriorSensitivitySummary, antecedent_validate::PriorSensitivity),
CausalError,
> {
use antecedent_validate::{ExternalAlphaSensitivity, PriorSensitivity};
if let Some(ext) = cfg.external_compose.as_ref() {
let alphas_applied: Arc<[f64]> = posterior.conflict_summary.as_ref().map_or_else(
|| Arc::clone(&ext.composed.alphas_applied),
|cs| Arc::clone(&cs.alphas_applied),
);
let sens = PriorSensitivity::standard_alpha_grid();
let (summary, _) = sens
.evaluate_external_alpha(
est,
prep,
status,
ws,
ctx,
ExternalAlphaSensitivity { sources: &ext.sources, alphas_applied: &alphas_applied },
)
.map_err(CausalError::from)?;
Ok((summary, sens))
} else {
let sens = PriorSensitivity::standard_grid();
let (summary, _) = sens.evaluate(est, prep, status, ws, ctx).map_err(CausalError::from)?;
Ok((summary, sens))
}
}