#![forbid(unsafe_code)]
#![deny(missing_docs)]
pub mod estimand;
pub mod eval;
pub mod latex;
pub mod pretty;
pub mod provider;
pub mod simplify;
pub use estimand::{EstimandMethod, IdentifiedEstimand, RdDesignParams};
pub use eval::CompiledEvaluator;
pub use provider::{
Assignment, DistributionProvider, EmpiricalTableProvider, EvalContext, EvalError, FactorSpec,
GaussianDensityProvider, PosteriorDrawProvider, QuadratureNodes,
};
pub use simplify::SimplifyError;
use latex::latex_expr;
use pretty::pretty_expr;
use std::collections::HashMap;
use std::fmt;
use std::sync::Arc;
use antecedent_core::{Value, VariableId};
#[repr(transparent)]
#[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
pub struct ExprId(u32);
impl ExprId {
#[must_use]
pub const fn from_raw(raw: u32) -> Self {
Self(raw)
}
#[must_use]
pub const fn raw(self) -> u32 {
self.0
}
}
#[repr(transparent)]
#[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
pub struct VarSetId(u32);
impl VarSetId {
#[must_use]
pub const fn from_raw(raw: u32) -> Self {
Self(raw)
}
#[must_use]
pub const fn raw(self) -> u32 {
self.0
}
}
#[repr(transparent)]
#[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
pub struct InterventionSetId(u32);
impl InterventionSetId {
#[must_use]
pub const fn from_raw(raw: u32) -> Self {
Self(raw)
}
#[must_use]
pub const fn raw(self) -> u32 {
self.0
}
}
#[derive(Clone, Debug, Eq, PartialEq, Hash)]
pub struct InterventionAssignment {
pub variable: VariableId,
pub value: Value,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
pub enum ContrastOp {
Difference,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
pub enum DomainRef {
Observational,
Interventional,
}
#[repr(transparent)]
#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
pub struct OutcomeExprId(VariableId);
impl OutcomeExprId {
#[must_use]
pub const fn identity(variable: VariableId) -> Self {
Self(variable)
}
#[must_use]
pub const fn variable(self) -> VariableId {
self.0
}
}
#[repr(transparent)]
#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
pub struct ExprListId(u32);
impl ExprListId {
#[must_use]
pub const fn from_raw(raw: u32) -> Self {
Self(raw)
}
#[must_use]
pub const fn raw(self) -> u32 {
self.0
}
}
#[derive(Clone, Debug, Eq, PartialEq, Hash)]
pub enum ExprNode {
Distribution {
variables: VarSetId,
conditioned_on: VarSetId,
intervention: InterventionSetId,
domain: DomainRef,
},
Product(ExprListId),
SumOut {
variables: VarSetId,
expr: ExprId,
},
IntegralOut {
variables: VarSetId,
expr: ExprId,
},
Ratio {
numerator: ExprId,
denominator: ExprId,
},
Expectation {
function: OutcomeExprId,
distribution: ExprId,
},
Contrast {
left: ExprId,
right: ExprId,
op: ContrastOp,
},
}
#[derive(Clone, Debug, Default, Eq, PartialEq)]
pub struct DerivationMeta {
pub rule: Arc<str>,
pub note: Option<Arc<str>>,
}
#[derive(Clone, Debug, Default)]
pub struct CausalExprArena {
nodes: Vec<ExprNode>,
var_sets: Vec<Arc<[VariableId]>>,
var_set_index: HashMap<Arc<[VariableId]>, VarSetId>,
interventions: Vec<Arc<[InterventionAssignment]>>,
intervention_index: HashMap<Arc<[InterventionAssignment]>, InterventionSetId>,
lists: Vec<Arc<[ExprId]>>,
list_index: HashMap<Arc<[ExprId]>, ExprListId>,
node_index: HashMap<ExprNode, ExprId>,
derivation: HashMap<u32, DerivationMeta>,
empty_var_set_id: Option<VarSetId>,
empty_intervention_set_id: Option<InterventionSetId>,
}
impl CausalExprArena {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn intern_var_set(&mut self, vars: impl IntoIterator<Item = VariableId>) -> VarSetId {
let mut v: Vec<VariableId> = vars.into_iter().collect();
v.sort_unstable();
v.dedup();
if let Some(id) = self.var_set_index.get(v.as_slice()) {
return *id;
}
let key: Arc<[VariableId]> = Arc::from(v);
let id = VarSetId(u32::try_from(self.var_sets.len()).expect("var set id"));
self.var_sets.push(Arc::clone(&key));
self.var_set_index.insert(key, id);
id
}
pub fn intern_intervention_assignments(
&mut self,
assignments: impl IntoIterator<Item = InterventionAssignment>,
) -> InterventionSetId {
let mut v: Vec<InterventionAssignment> = assignments.into_iter().collect();
v.sort_by_key(|a| a.variable.raw());
v.dedup_by_key(|a| a.variable.raw());
if let Some(id) = self.intervention_index.get(v.as_slice()) {
return *id;
}
let key: Arc<[InterventionAssignment]> = Arc::from(v);
let id = InterventionSetId(u32::try_from(self.interventions.len()).expect("id"));
self.interventions.push(Arc::clone(&key));
self.intervention_index.insert(key, id);
id
}
pub fn intern_intervention_set(
&mut self,
vars: impl IntoIterator<Item = VariableId>,
) -> InterventionSetId {
self.intern_intervention_assignments(
vars.into_iter()
.map(|variable| InterventionAssignment { variable, value: Value::f64(f64::NAN) }),
)
}
pub fn empty_var_set(&mut self) -> VarSetId {
if let Some(id) = self.empty_var_set_id {
return id;
}
let id = self.intern_var_set([]);
self.empty_var_set_id = Some(id);
id
}
pub fn empty_intervention_set(&mut self) -> InterventionSetId {
if let Some(id) = self.empty_intervention_set_id {
return id;
}
let id = self.intern_intervention_assignments([]);
self.empty_intervention_set_id = Some(id);
id
}
#[must_use]
pub fn var_set(&self, id: VarSetId) -> &[VariableId] {
&self.var_sets[id.0 as usize]
}
#[must_use]
pub fn intervention_assignments(&self, id: InterventionSetId) -> &[InterventionAssignment] {
&self.interventions[id.0 as usize]
}
#[must_use]
pub fn intervention_set(&self, id: InterventionSetId) -> Vec<VariableId> {
self.intervention_assignments(id).iter().map(|a| a.variable).collect()
}
pub fn intern_list(&mut self, exprs: impl IntoIterator<Item = ExprId>) -> ExprListId {
let v: Vec<ExprId> = exprs.into_iter().collect();
if let Some(id) = self.list_index.get(v.as_slice()) {
return *id;
}
let key: Arc<[ExprId]> = Arc::from(v);
let id = ExprListId(u32::try_from(self.lists.len()).expect("list id"));
self.lists.push(Arc::clone(&key));
self.list_index.insert(key, id);
id
}
#[must_use]
pub fn list(&self, id: ExprListId) -> &[ExprId] {
&self.lists[id.0 as usize]
}
pub fn intern(&mut self, node: ExprNode) -> ExprId {
if let Some(id) = self.node_index.get(&node) {
return *id;
}
let id = ExprId(u32::try_from(self.nodes.len()).expect("expr id"));
self.nodes.push(node.clone());
self.node_index.insert(node, id);
id
}
pub fn set_derivation(&mut self, id: ExprId, meta: DerivationMeta) {
self.derivation.insert(id.0, meta);
}
pub fn set_derivation_if_absent(&mut self, id: ExprId, meta: DerivationMeta) {
self.derivation.entry(id.0).or_insert(meta);
}
pub fn simplify(&mut self, root: ExprId) -> Result<ExprId, SimplifyError> {
simplify::simplify(self, root)
}
#[must_use]
pub fn derivation(&self, id: ExprId) -> Option<&DerivationMeta> {
self.derivation.get(&id.0)
}
#[must_use]
pub fn node(&self, id: ExprId) -> &ExprNode {
&self.nodes[id.0 as usize]
}
#[must_use]
pub fn len(&self) -> usize {
self.nodes.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.nodes.is_empty()
}
#[must_use]
pub fn var_set_count(&self) -> usize {
self.var_sets.len()
}
#[must_use]
pub fn intervention_set_count(&self) -> usize {
self.interventions.len()
}
#[must_use]
pub fn list_count(&self) -> usize {
self.lists.len()
}
pub fn backdoor_ate(
&mut self,
treatment: VariableId,
outcome: VariableId,
adjustment: &[VariableId],
active: Value,
control: Value,
) -> ExprId {
let left = self.backdoor_potential_outcome(treatment, outcome, adjustment, active);
let right = self.backdoor_potential_outcome(treatment, outcome, adjustment, control);
let contrast = self.intern(ExprNode::Contrast { left, right, op: ContrastOp::Difference });
self.set_derivation(
contrast,
DerivationMeta {
rule: Arc::from("backdoor.adjustment"),
note: Some(Arc::from(format!("ATE adjustment set size {}", adjustment.len()))),
},
);
contrast
}
fn backdoor_potential_outcome(
&mut self,
treatment: VariableId,
outcome: VariableId,
adjustment: &[VariableId],
level: Value,
) -> ExprId {
let z = self.intern_var_set(adjustment.iter().copied());
let y = self.intern_var_set([outcome]);
let empty = self.empty_var_set();
let empty_i = self.empty_intervention_set();
let do_t = self.intern_intervention_assignments([InterventionAssignment {
variable: treatment,
value: level,
}]);
let dist_body = self.intern(ExprNode::Distribution {
variables: y,
conditioned_on: z,
intervention: do_t,
domain: DomainRef::Interventional,
});
let z_marg = self.intern(ExprNode::Distribution {
variables: z,
conditioned_on: empty,
intervention: empty_i,
domain: DomainRef::Observational,
});
let product = {
let list = self.intern_list([dist_body, z_marg]);
self.intern(ExprNode::Product(list))
};
let summed = self.intern(ExprNode::SumOut { variables: z, expr: product });
self.intern(ExprNode::Expectation {
function: OutcomeExprId::identity(outcome),
distribution: summed,
})
}
pub fn frontdoor_ate(
&mut self,
treatment: VariableId,
outcome: VariableId,
mediators: &[VariableId],
active: Value,
control: Value,
) -> ExprId {
let left = self.frontdoor_potential_outcome(treatment, outcome, mediators, active);
let right = self.frontdoor_potential_outcome(treatment, outcome, mediators, control);
let contrast = self.intern(ExprNode::Contrast { left, right, op: ContrastOp::Difference });
self.set_derivation(
contrast,
DerivationMeta {
rule: Arc::from("frontdoor"),
note: Some(Arc::from(format!("front-door mediator set size {}", mediators.len()))),
},
);
contrast
}
pub fn temporal_mediation_ate(
&mut self,
treatment: VariableId,
outcome: VariableId,
mediators: &[VariableId],
active: Value,
control: Value,
) -> ExprId {
let left = self.frontdoor_potential_outcome(treatment, outcome, mediators, active);
let right = self.frontdoor_potential_outcome(treatment, outcome, mediators, control);
let contrast = self.intern(ExprNode::Contrast { left, right, op: ContrastOp::Difference });
self.set_derivation(
contrast,
DerivationMeta {
rule: Arc::from("temporal_mediation"),
note: Some(Arc::from(format!(
"linear temporal mediation path-product; mediator set size {}",
mediators.len()
))),
},
);
contrast
}
fn frontdoor_potential_outcome(
&mut self,
treatment: VariableId,
outcome: VariableId,
mediators: &[VariableId],
level: Value,
) -> ExprId {
let m = self.intern_var_set(mediators.iter().copied());
let y = self.intern_var_set([outcome]);
let t = self.intern_var_set([treatment]);
let m_and_t = self.intern_var_set(mediators.iter().copied().chain([treatment]));
let empty = self.empty_var_set();
let empty_i = self.empty_intervention_set();
let do_t = self.intern_intervention_assignments([InterventionAssignment {
variable: treatment,
value: level,
}]);
let m_given_t = self.intern(ExprNode::Distribution {
variables: m,
conditioned_on: t,
intervention: do_t,
domain: DomainRef::Observational,
});
let y_given_m_t = self.intern(ExprNode::Distribution {
variables: y,
conditioned_on: m_and_t,
intervention: empty_i,
domain: DomainRef::Observational,
});
let t_marginal = self.intern(ExprNode::Distribution {
variables: t,
conditioned_on: empty,
intervention: empty_i,
domain: DomainRef::Observational,
});
let inner_product = {
let list = self.intern_list([y_given_m_t, t_marginal]);
self.intern(ExprNode::Product(list))
};
let inner_summed = self.intern(ExprNode::SumOut { variables: t, expr: inner_product });
let outer_product = {
let list = self.intern_list([m_given_t, inner_summed]);
self.intern(ExprNode::Product(list))
};
let outer_summed = self.intern(ExprNode::SumOut { variables: m, expr: outer_product });
self.intern(ExprNode::Expectation {
function: OutcomeExprId::identity(outcome),
distribution: outer_summed,
})
}
pub fn iv_wald(
&mut self,
treatment: VariableId,
outcome: VariableId,
instruments: &[VariableId],
active: &Value,
control: &Value,
) -> ExprId {
let z = instruments.first().copied().unwrap_or(treatment);
let z1 = Value::f64(1.0);
let z0 = Value::f64(0.0);
let outcome_given_z1 = self.observational_conditional_mean(outcome, z, z1.clone());
let outcome_given_z0 = self.observational_conditional_mean(outcome, z, z0.clone());
let treatment_given_z1 = self.observational_conditional_mean(treatment, z, z1);
let treatment_given_z0 = self.observational_conditional_mean(treatment, z, z0);
let num = self.intern(ExprNode::Contrast {
left: outcome_given_z1,
right: outcome_given_z0,
op: ContrastOp::Difference,
});
let den = self.intern(ExprNode::Contrast {
left: treatment_given_z1,
right: treatment_given_z0,
op: ContrastOp::Difference,
});
let ratio = self.intern(ExprNode::Ratio { numerator: num, denominator: den });
self.set_derivation(
ratio,
DerivationMeta {
rule: Arc::from("iv.wald"),
note: Some(Arc::from(format!(
"Wald IV ratio using {} instrument(s); treatment contrast [{active:?}, {control:?}]",
instruments.len()
))),
},
);
ratio
}
fn observational_conditional_mean(
&mut self,
outcome: VariableId,
conditioner: VariableId,
level: Value,
) -> ExprId {
let y = self.intern_var_set([outcome]);
let z = self.intern_var_set([conditioner]);
let bind = self.intern_intervention_assignments([InterventionAssignment {
variable: conditioner,
value: level,
}]);
let dist = self.intern(ExprNode::Distribution {
variables: y,
conditioned_on: z,
intervention: bind,
domain: DomainRef::Observational,
});
self.intern(ExprNode::Expectation {
function: OutcomeExprId::identity(outcome),
distribution: dist,
})
}
#[must_use]
pub fn pretty(&self, id: ExprId) -> String {
pretty_expr(self, id)
}
#[must_use]
pub fn latex(&self, id: ExprId) -> String {
latex_expr(self, id)
}
}
impl fmt::Display for ExprId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "E{}", self.0)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn var_sets_are_sorted_and_interned() {
let mut a = CausalExprArena::new();
let s1 = a.intern_var_set([VariableId::from_raw(2), VariableId::from_raw(1)]);
let s2 = a.intern_var_set([VariableId::from_raw(1), VariableId::from_raw(2)]);
assert_eq!(s1, s2);
assert_eq!(a.var_set(s1), &[VariableId::from_raw(1), VariableId::from_raw(2)]);
}
#[test]
fn hash_cons_reuses_nodes() {
let mut a = CausalExprArena::new();
let empty = a.empty_var_set();
let empty_i = a.empty_intervention_set();
let n1 = a.intern(ExprNode::Distribution {
variables: empty,
conditioned_on: empty,
intervention: empty_i,
domain: DomainRef::Observational,
});
let n2 = a.intern(ExprNode::Distribution {
variables: empty,
conditioned_on: empty,
intervention: empty_i,
domain: DomainRef::Observational,
});
assert_eq!(n1, n2);
assert_eq!(a.len(), 1);
}
#[test]
fn backdoor_ate_contrasts_distinct_levels() {
let mut a = CausalExprArena::new();
let id = a.backdoor_ate(
VariableId::from_raw(0),
VariableId::from_raw(1),
&[VariableId::from_raw(2)],
Value::f64(1.0),
Value::f64(0.0),
);
let meta = a.derivation(id).unwrap();
assert_eq!(&*meta.rule, "backdoor.adjustment");
let ExprNode::Contrast { left, right, .. } = a.node(id) else {
panic!("expected contrast");
};
assert_ne!(left, right);
let pretty = a.pretty(id);
assert!(pretty.contains('−') || pretty.contains("E["));
let latex = a.latex(id);
assert!(latex.contains("\\mathbb{E}") || latex.contains("\\mathrm{do}"));
assert!(latex.contains('-'));
}
#[test]
fn frontdoor_ate_contrasts_distinct_levels() {
let mut a = CausalExprArena::new();
let id = a.frontdoor_ate(
VariableId::from_raw(0),
VariableId::from_raw(1),
&[VariableId::from_raw(2)],
Value::f64(1.0),
Value::f64(0.0),
);
let meta = a.derivation(id).unwrap();
assert_eq!(&*meta.rule, "frontdoor");
let ExprNode::Contrast { left, right, .. } = a.node(id) else {
panic!("expected contrast");
};
assert_ne!(left, right);
}
#[test]
fn iv_wald_is_ratio_of_instrument_contrasts() {
let mut a = CausalExprArena::new();
let id = a.iv_wald(
VariableId::from_raw(0),
VariableId::from_raw(1),
&[VariableId::from_raw(2)],
&Value::f64(1.0),
&Value::f64(0.0),
);
let meta = a.derivation(id).unwrap();
assert_eq!(&*meta.rule, "iv.wald");
let ExprNode::Ratio { numerator, denominator } = a.node(id) else {
panic!("expected Wald ratio");
};
assert!(matches!(a.node(*numerator), ExprNode::Contrast { .. }));
assert!(matches!(a.node(*denominator), ExprNode::Contrast { .. }));
assert_ne!(*numerator, *denominator);
}
}