use super::*;
impl super::CausalAnalysis {
pub(super) fn execute_counterfactual(
&self,
data: &TabularData,
graph: &Dag,
query: &antecedent_core::CounterfactualQuery,
physical: &PhysicalExecutionPlan,
ctx: &ExecutionContext,
) -> Result<CausalAnalysisResult, CausalError> {
let started = Instant::now();
query.validate().map_err(|e| CausalError::Compile { message: e.to_string() })?;
let fitted = fit_gcm(graph.clone(), data)?;
let outcome = *query.outcomes.first().ok_or_else(|| CausalError::Compile {
message: "counterfactual query missing outcome".into(),
})?;
let (treatment, active, control) = binary_cf_interventions(query)?;
let ite = counterfactual_ite(fitted.model, data, treatment, outcome, active, control, ctx)?;
let estimate = EffectEstimate::new(
ite.mean_ite,
f64::NAN,
antecedent_core::AssumptionSet::default(),
OverlapPolicy::ExplicitOverride,
);
let (identification, estimand) = parametric_scm_identification(
CausalQuery::Counterfactual(query.clone()),
treatment,
outcome,
);
let mut diagnostics = vec![Diagnostic::new(
"gcm.counterfactual",
DiagnosticKind::Scientific,
DiagnosticSeverity::Info,
format!("noise_inference={:?}", ite.noise_inference),
)];
let physical_record =
self.apply_callback_plan_marks(physical.record.clone(), &mut diagnostics);
Ok(assemble_result(AssembleArgs {
logical: &physical.logical.record,
physical: &physical_record,
identification,
estimand,
estimate,
distribution: None,
posterior: None,
mediation: None,
counterfactual: Some(ite),
anomaly: None,
change_attribution: None,
mechanism_change: None,
unit_change: None,
refutations: Vec::new(),
diagnostics,
provenance: ProvenanceGraph::new(),
treatment,
outcome,
wall_time_ns: u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX),
latency_mode: self.latency_mode.map(|m| Arc::from(m.as_str())),
stage_timings_ns: Vec::new(),
bootstrap_replicates_requested: Some(self.bootstrap_replicates),
bootstrap_replicates_ok: None,
n_draws: None,
cancelled: false,
early_stopped: false,
}))
}
pub(super) fn execute_anomaly(
&self,
data: &TabularData,
graph: &Dag,
query: &antecedent_core::AnomalyAttributionQuery,
physical: &PhysicalExecutionPlan,
ctx: &ExecutionContext,
) -> Result<CausalAnalysisResult, CausalError> {
let _ = ctx;
let started = Instant::now();
query.validate().map_err(|e| CausalError::Compile { message: e.to_string() })?;
let fitted = fit_gcm(graph.clone(), data)?;
let scores = anomaly_attribution(
&fitted.model,
data,
query.targets.iter().copied(),
query.max_units,
)?;
let outcome = *query.targets.first().unwrap_or(&VariableId::from_raw(0));
let treatment = outcome;
let estimate = nan_effect();
let (identification, estimand) = parametric_scm_identification(
CausalQuery::AnomalyAttribution(query.clone()),
treatment,
outcome,
);
let mut diagnostics = Vec::new();
let physical_record =
self.apply_callback_plan_marks(physical.record.clone(), &mut diagnostics);
Ok(assemble_result(AssembleArgs {
logical: &physical.logical.record,
physical: &physical_record,
identification,
estimand,
estimate,
distribution: None,
posterior: None,
mediation: None,
counterfactual: None,
anomaly: Some(scores),
change_attribution: None,
mechanism_change: None,
unit_change: None,
refutations: Vec::new(),
diagnostics,
provenance: ProvenanceGraph::new(),
treatment,
outcome,
wall_time_ns: u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX),
latency_mode: self.latency_mode.map(|m| Arc::from(m.as_str())),
stage_timings_ns: Vec::new(),
bootstrap_replicates_requested: Some(self.bootstrap_replicates),
bootstrap_replicates_ok: None,
n_draws: None,
cancelled: false,
early_stopped: false,
}))
}
pub(super) fn execute_change_attribution(
&self,
data: &TabularData,
graph: &Dag,
query: &antecedent_core::ChangeAttributionQuery,
physical: &PhysicalExecutionPlan,
ctx: &ExecutionContext,
) -> Result<CausalAnalysisResult, CausalError> {
let started = Instant::now();
query.validate().map_err(|e| CausalError::Compile { message: e.to_string() })?;
let fitted = fit_gcm(graph.clone(), data)?;
let result = attribute_distribution_change(
&fitted.model,
data,
query,
&antecedent_attribution::DistributionChangeOptions::default(),
ctx,
)?;
let outcome = query.outcome;
let treatment = outcome;
let estimate = EffectEstimate::new(
result.total_change,
f64::NAN,
antecedent_core::AssumptionSet::default(),
OverlapPolicy::ExplicitOverride,
);
let (identification, estimand) = parametric_scm_identification(
CausalQuery::ChangeAttribution(query.clone()),
treatment,
outcome,
);
let mut diagnostics = Vec::new();
let physical_record =
self.apply_callback_plan_marks(physical.record.clone(), &mut diagnostics);
Ok(assemble_result(AssembleArgs {
logical: &physical.logical.record,
physical: &physical_record,
identification,
estimand,
estimate,
distribution: None,
posterior: None,
mediation: None,
counterfactual: None,
anomaly: None,
change_attribution: Some(result),
mechanism_change: None,
unit_change: None,
refutations: Vec::new(),
diagnostics,
provenance: ProvenanceGraph::new(),
treatment,
outcome,
wall_time_ns: u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX),
latency_mode: self.latency_mode.map(|m| Arc::from(m.as_str())),
stage_timings_ns: Vec::new(),
bootstrap_replicates_requested: Some(self.bootstrap_replicates),
bootstrap_replicates_ok: None,
n_draws: None,
cancelled: false,
early_stopped: false,
}))
}
pub(super) fn execute_mechanism_change(
&self,
data: &TabularData,
graph: &Dag,
query: &antecedent_core::MechanismChangeQuery,
physical: &PhysicalExecutionPlan,
ctx: &ExecutionContext,
) -> Result<CausalAnalysisResult, CausalError> {
let started = Instant::now();
query.validate().map_err(|e| CausalError::Compile { message: e.to_string() })?;
let fitted = fit_gcm(graph.clone(), data)?;
let detections = mechanism_change_detection(
&fitted.model,
data,
query,
antecedent_attribution::MechanismChangeMethod::LikelihoodRatio,
ctx,
)?;
let outcome = *query.targets.first().unwrap_or(&VariableId::from_raw(0));
let treatment = outcome;
let estimate = nan_effect();
let (identification, estimand) = parametric_scm_identification(
CausalQuery::MechanismChange(query.clone()),
treatment,
outcome,
);
let mut diagnostics = Vec::new();
let physical_record =
self.apply_callback_plan_marks(physical.record.clone(), &mut diagnostics);
Ok(assemble_result(AssembleArgs {
logical: &physical.logical.record,
physical: &physical_record,
identification,
estimand,
estimate,
distribution: None,
posterior: None,
mediation: None,
counterfactual: None,
anomaly: None,
change_attribution: None,
mechanism_change: Some(detections),
unit_change: None,
refutations: Vec::new(),
diagnostics,
provenance: ProvenanceGraph::new(),
treatment,
outcome,
wall_time_ns: u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX),
latency_mode: self.latency_mode.map(|m| Arc::from(m.as_str())),
stage_timings_ns: Vec::new(),
bootstrap_replicates_requested: Some(self.bootstrap_replicates),
bootstrap_replicates_ok: None,
n_draws: None,
cancelled: false,
early_stopped: false,
}))
}
pub(super) fn execute_unit_change(
&self,
data: &TabularData,
graph: &Dag,
query: &antecedent_core::UnitChangeQuery,
physical: &PhysicalExecutionPlan,
ctx: &ExecutionContext,
) -> Result<CausalAnalysisResult, CausalError> {
let started = Instant::now();
query.validate().map_err(|e| CausalError::Compile { message: e.to_string() })?;
let fitted = fit_gcm(graph.clone(), data)?;
let result = attribute_unit_change(&fitted.model, data, query, ctx)?;
let outcome = query.outcome;
let treatment = outcome;
let estimate = nan_effect();
let (identification, estimand) = parametric_scm_identification(
CausalQuery::UnitChange(query.clone()),
treatment,
outcome,
);
let mut diagnostics = Vec::new();
let physical_record =
self.apply_callback_plan_marks(physical.record.clone(), &mut diagnostics);
Ok(assemble_result(AssembleArgs {
logical: &physical.logical.record,
physical: &physical_record,
identification,
estimand,
estimate,
distribution: None,
posterior: None,
mediation: None,
counterfactual: None,
anomaly: None,
change_attribution: None,
mechanism_change: None,
unit_change: Some(result),
refutations: Vec::new(),
diagnostics,
provenance: ProvenanceGraph::new(),
treatment,
outcome,
wall_time_ns: u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX),
latency_mode: self.latency_mode.map(|m| Arc::from(m.as_str())),
stage_timings_ns: Vec::new(),
bootstrap_replicates_requested: Some(self.bootstrap_replicates),
bootstrap_replicates_ok: None,
n_draws: None,
cancelled: false,
early_stopped: false,
}))
}
}