use std::collections::HashMap;
use std::sync::Arc;
use antecedent_core::{Value, VariableId};
use crate::provider::{Assignment, DistributionProvider, EvalContext, EvalError, FactorSpec};
use crate::{
CausalExprArena, ContrastOp, DomainRef, ExprId, ExprNode, InterventionSetId, OutcomeExprId,
VarSetId,
};
#[derive(Clone, Debug)]
enum EvalOp {
Distribution {
variables: VarSetId,
conditioned_on: VarSetId,
intervention: InterventionSetId,
domain: DomainRef,
},
Product {
children: Arc<[usize]>,
},
SumOut {
variables: VarSetId,
body: usize,
},
IntegralOut {
variables: VarSetId,
body: usize,
},
Ratio {
numerator: usize,
denominator: usize,
},
Expectation {
function: OutcomeExprId,
distribution: usize,
},
Contrast {
left: usize,
right: usize,
op: ContrastOp,
},
}
#[derive(Clone, Debug)]
pub struct CompiledEvaluator {
ops: Vec<EvalOp>,
free_vars: Vec<Arc<[VariableId]>>,
root: usize,
}
impl CausalExprArena {
pub fn compile(&self, root: ExprId) -> Result<CompiledEvaluator, EvalError> {
CompiledEvaluator::compile(self, root)
}
}
impl CompiledEvaluator {
pub fn compile(arena: &CausalExprArena, root: ExprId) -> Result<Self, EvalError> {
let mut ops = Vec::new();
let mut expr_to_slot = HashMap::new();
let root_slot = compile_rec(arena, root, &mut ops, &mut expr_to_slot)?;
let free_vars = compute_free_vars(&ops, arena);
Ok(Self { ops, free_vars, root: root_slot })
}
pub fn evaluate(
&self,
arena: &CausalExprArena,
provider: &dyn DistributionProvider,
ctx: &EvalContext,
) -> Result<f64, EvalError> {
self.evaluate_with(arena, provider, ctx, &Assignment::new())
}
pub fn evaluate_with(
&self,
arena: &CausalExprArena,
provider: &dyn DistributionProvider,
ctx: &EvalContext,
env: &Assignment,
) -> Result<f64, EvalError> {
let mut scratch = env.clone();
self.eval_slot(arena, provider, ctx, &mut scratch, self.root)
}
pub fn evaluate_batch(
&self,
arena: &CausalExprArena,
provider: &dyn DistributionProvider,
) -> Result<Vec<f64>, EvalError> {
match provider.n_draws() {
None => Ok(vec![self.evaluate(arena, provider, &EvalContext::default())?]),
Some(n) => {
let mut out = Vec::with_capacity(n);
for draw in 0..n {
let ctx = EvalContext { draw: Some(draw) };
out.push(self.evaluate(arena, provider, &ctx)?);
}
Ok(out)
}
}
}
fn eval_slot(
&self,
arena: &CausalExprArena,
provider: &dyn DistributionProvider,
ctx: &EvalContext,
env: &mut Assignment,
slot: usize,
) -> Result<f64, EvalError> {
match &self.ops[slot] {
EvalOp::Distribution { variables, conditioned_on, intervention, domain } => {
let spec = FactorSpec {
variables: arena.var_set(*variables),
conditioned_on: arena.var_set(*conditioned_on),
intervention: arena.intervention_assignments(*intervention),
domain: *domain,
};
with_scoped_bindings(env, spec.intervention.iter().map(|a| a.variable), |env| {
for a in spec.intervention {
env.set(a.variable, a.value.clone());
}
provider.probability(&spec, env, ctx)
})
}
EvalOp::Product { children } => {
let mut prod = 1.0;
for &c in children.iter() {
prod *= self.eval_slot(arena, provider, ctx, env, c)?;
}
Ok(prod)
}
EvalOp::SumOut { variables, body } => {
self.eval_sum_out(arena, provider, ctx, env, *variables, *body)
}
EvalOp::IntegralOut { variables, body } => {
self.eval_integral_out(arena, provider, ctx, env, *variables, *body)
}
EvalOp::Ratio { numerator, denominator } => {
let num = self.eval_slot(arena, provider, ctx, env, *numerator)?;
let den = self.eval_slot(arena, provider, ctx, env, *denominator)?;
if den == 0.0 {
return Err(EvalError::DivisionByZero);
}
Ok(num / den)
}
EvalOp::Expectation { function, distribution } => {
self.eval_expectation(arena, provider, ctx, env, function.variable(), *distribution)
}
EvalOp::Contrast { left, right, op } => {
let l = self.eval_slot(arena, provider, ctx, env, *left)?;
let r = self.eval_slot(arena, provider, ctx, env, *right)?;
match op {
ContrastOp::Difference => Ok(l - r),
}
}
}
}
fn eval_sum_out(
&self,
arena: &CausalExprArena,
provider: &dyn DistributionProvider,
ctx: &EvalContext,
env: &mut Assignment,
variables: VarSetId,
body: usize,
) -> Result<f64, EvalError> {
let vars = arena.var_set(variables);
let rows = provider.support(vars, ctx)?;
with_scoped_bindings(env, vars.iter().copied(), |env| {
let mut sum = 0.0;
for row in rows.iter() {
if row.len() != vars.len() {
return Err(EvalError::SupportShape {
expected: vars.len(),
actual: row.len(),
});
}
for (i, &v) in vars.iter().enumerate() {
env.set(v, row[i].clone());
}
sum += self.eval_slot(arena, provider, ctx, env, body)?;
}
Ok(sum)
})
}
fn eval_integral_out(
&self,
arena: &CausalExprArena,
provider: &dyn DistributionProvider,
ctx: &EvalContext,
env: &mut Assignment,
variables: VarSetId,
body: usize,
) -> Result<f64, EvalError> {
let vars = arena.var_set(variables);
if let Some(nodes) = provider.quadrature(vars, ctx)? {
return with_scoped_bindings(env, vars.iter().copied(), |env| {
let mut acc = 0.0;
for (row, weight) in nodes.iter() {
if row.len() != vars.len() {
return Err(EvalError::SupportShape {
expected: vars.len(),
actual: row.len(),
});
}
for (i, &v) in vars.iter().enumerate() {
env.set(v, row[i].clone());
}
acc += *weight * self.eval_slot(arena, provider, ctx, env, body)?;
}
Ok(acc)
});
}
let rows = provider.support(vars, ctx).map_err(|e| match e {
EvalError::EmptySupport(_) => EvalError::UnsupportedIntegralOut,
other => other,
})?;
with_scoped_bindings(env, vars.iter().copied(), |env| {
let mut sum = 0.0;
for row in rows.iter() {
if row.len() != vars.len() {
return Err(EvalError::SupportShape {
expected: vars.len(),
actual: row.len(),
});
}
for (i, &v) in vars.iter().enumerate() {
env.set(v, row[i].clone());
}
sum += self.eval_slot(arena, provider, ctx, env, body)?;
}
Ok(sum)
})
}
fn eval_expectation(
&self,
arena: &CausalExprArena,
provider: &dyn DistributionProvider,
ctx: &EvalContext,
env: &mut Assignment,
outcome_var: VariableId,
distribution: usize,
) -> Result<f64, EvalError> {
let free = &self.free_vars[distribution];
let mut enum_vars: Vec<VariableId> =
free.iter().copied().filter(|v| env.get(*v).is_none()).collect();
if !enum_vars.contains(&outcome_var) && env.get(outcome_var).is_none() {
enum_vars.push(outcome_var);
}
enum_vars.sort_by_key(|v| v.raw());
enum_vars.dedup();
if enum_vars.is_empty() {
let dens = self.eval_slot(arena, provider, ctx, env, distribution)?;
let y = provider.outcome(outcome_var, env, ctx)?;
return Ok(y * dens);
}
let rows = provider.support(&enum_vars, ctx)?;
with_scoped_bindings(env, enum_vars.iter().copied(), |env| {
let mut acc = 0.0;
for row in rows.iter() {
if row.len() != enum_vars.len() {
return Err(EvalError::SupportShape {
expected: enum_vars.len(),
actual: row.len(),
});
}
for (i, &v) in enum_vars.iter().enumerate() {
env.set(v, row[i].clone());
}
let dens = self.eval_slot(arena, provider, ctx, env, distribution)?;
let y = provider.outcome(outcome_var, env, ctx)?;
acc += y * dens;
}
Ok(acc)
})
}
}
fn with_scoped_bindings<T>(
env: &mut Assignment,
vars: impl IntoIterator<Item = VariableId>,
f: impl FnOnce(&mut Assignment) -> Result<T, EvalError>,
) -> Result<T, EvalError> {
let saved: Vec<(VariableId, Option<Value>)> =
vars.into_iter().map(|v| (v, env.get(v).cloned())).collect();
let result = f(env);
for (v, prev) in saved {
match prev {
Some(value) => env.set(v, value),
None => {
env.remove(v);
}
}
}
result
}
fn compile_rec(
arena: &CausalExprArena,
id: ExprId,
ops: &mut Vec<EvalOp>,
expr_to_slot: &mut HashMap<u32, usize>,
) -> Result<usize, EvalError> {
if let Some(&slot) = expr_to_slot.get(&id.raw()) {
return Ok(slot);
}
let op = match arena.node(id).clone() {
ExprNode::Distribution { variables, conditioned_on, intervention, domain } => {
EvalOp::Distribution { variables, conditioned_on, intervention, domain }
}
ExprNode::Product(list) => {
let mut children = Vec::new();
for &c in arena.list(list) {
children.push(compile_rec(arena, c, ops, expr_to_slot)?);
}
EvalOp::Product { children: Arc::from(children) }
}
ExprNode::SumOut { variables, expr } => {
let body = compile_rec(arena, expr, ops, expr_to_slot)?;
EvalOp::SumOut { variables, body }
}
ExprNode::IntegralOut { variables, expr } => {
let body = compile_rec(arena, expr, ops, expr_to_slot)?;
EvalOp::IntegralOut { variables, body }
}
ExprNode::Ratio { numerator, denominator } => {
let n = compile_rec(arena, numerator, ops, expr_to_slot)?;
let d = compile_rec(arena, denominator, ops, expr_to_slot)?;
EvalOp::Ratio { numerator: n, denominator: d }
}
ExprNode::Expectation { function, distribution } => {
let dist = compile_rec(arena, distribution, ops, expr_to_slot)?;
EvalOp::Expectation { function, distribution: dist }
}
ExprNode::Contrast { left, right, op } => {
let l = compile_rec(arena, left, ops, expr_to_slot)?;
let r = compile_rec(arena, right, ops, expr_to_slot)?;
EvalOp::Contrast { left: l, right: r, op }
}
};
let slot = ops.len();
ops.push(op);
expr_to_slot.insert(id.raw(), slot);
Ok(slot)
}
fn compute_free_vars(ops: &[EvalOp], arena: &CausalExprArena) -> Vec<Arc<[VariableId]>> {
let mut out: Vec<Arc<[VariableId]>> = Vec::with_capacity(ops.len());
for op in ops {
let mut vars: Vec<VariableId> = match op {
EvalOp::Distribution { variables, conditioned_on, intervention, .. } => {
let mut vars = arena.var_set(*variables).to_vec();
let bound = arena.intervention_assignments(*intervention);
for &v in arena.var_set(*conditioned_on) {
if !bound.iter().any(|a| a.variable == v) {
vars.push(v);
}
}
vars
}
EvalOp::Product { children } => {
children.iter().flat_map(|&c| out[c].iter().copied()).collect()
}
EvalOp::SumOut { variables, body } | EvalOp::IntegralOut { variables, body } => {
let bound = arena.var_set(*variables);
out[*body].iter().copied().filter(|v| !bound.contains(v)).collect()
}
EvalOp::Ratio { numerator, denominator } => {
out[*numerator].iter().chain(out[*denominator].iter()).copied().collect()
}
EvalOp::Expectation { function, distribution } => {
let mut vars = out[*distribution].to_vec();
vars.push(function.variable());
vars
}
EvalOp::Contrast { left, right, .. } => {
out[*left].iter().chain(out[*right].iter()).copied().collect()
}
};
vars.sort_by_key(|v| v.raw());
vars.dedup();
out.push(Arc::from(vars));
}
out
}
#[cfg(test)]
mod tests {
use super::*;
use crate::provider::{EmpiricalTableProvider, PosteriorDrawProvider};
use crate::{InterventionAssignment, OutcomeExprId};
use antecedent_core::Value;
fn v(id: u32) -> VariableId {
VariableId::from_raw(id)
}
fn f(x: f64) -> Value {
Value::f64(x)
}
fn backdoor_provider(t: VariableId, y: VariableId, z: VariableId) -> EmpiricalTableProvider {
let mut p = EmpiricalTableProvider::new();
p.set_domain(z, [f(0.0), f(1.0)]);
p.set_domain(y, [f(0.0), f(1.0)]);
p.set_domain(t, [f(0.0), f(1.0)]);
for (zval, prob) in [(0.0, 0.5), (1.0, 0.5)] {
let spec = FactorSpec {
variables: &[z],
conditioned_on: &[],
intervention: &[],
domain: DomainRef::Observational,
};
let assign = Assignment::from_pairs([(z, f(zval))]);
p.insert_probability(&spec, &assign, prob).unwrap();
}
let ey = |tlev: f64, zlev: f64| -> f64 {
match (tlev.to_bits(), zlev.to_bits()) {
(t, z) if t == 1.0f64.to_bits() && z == 0.0f64.to_bits() => 0.8,
(t, z) if t == 1.0f64.to_bits() && z == 1.0f64.to_bits() => 0.6,
(t, z) if t == 0.0f64.to_bits() && z == 0.0f64.to_bits() => 0.3,
(t, z) if t == 0.0f64.to_bits() && z == 1.0f64.to_bits() => 0.2,
_ => panic!("bad levels"),
}
};
for tlev in [0.0, 1.0] {
let interv = [InterventionAssignment { variable: t, value: f(tlev) }];
for zlev in [0.0, 1.0] {
let p_y1 = ey(tlev, zlev);
for (yval, prob) in [(1.0, p_y1), (0.0, 1.0 - p_y1)] {
let spec = FactorSpec {
variables: &[y],
conditioned_on: &[z],
intervention: &interv,
domain: DomainRef::Interventional,
};
let assign = Assignment::from_pairs([(y, f(yval)), (z, f(zlev))]);
p.insert_probability(&spec, &assign, prob).unwrap();
}
}
}
p
}
#[test]
fn backdoor_ate_matches_closed_form() {
let mut arena = CausalExprArena::new();
let t = v(0);
let y = v(1);
let z = v(2);
let expr = arena.backdoor_ate(t, y, &[z], f(1.0), f(0.0));
let provider = backdoor_provider(t, y, z);
let compiled = arena.compile(expr).unwrap();
let ate = compiled.evaluate(&arena, &provider, &EvalContext::default()).unwrap();
assert!((ate - 0.45).abs() < 1e-12, "ate={ate}");
}
#[test]
fn simplify_preserves_backdoor_evaluation() {
let mut arena = CausalExprArena::new();
let t = v(0);
let y = v(1);
let z = v(2);
let expr = arena.backdoor_ate(t, y, &[z], f(1.0), f(0.0));
let provider = backdoor_provider(t, y, z);
let before = arena
.compile(expr)
.unwrap()
.evaluate(&arena, &provider, &EvalContext::default())
.unwrap();
let simplified = arena.simplify(expr).unwrap();
let after = arena
.compile(simplified)
.unwrap()
.evaluate(&arena, &provider, &EvalContext::default())
.unwrap();
assert!((before - after).abs() < 1e-12, "before={before} after={after}");
assert!((after - 0.45).abs() < 1e-12);
}
#[test]
fn simplify_preserves_backdoor_empty_evaluation() {
fn assert_simplify_preserves(
arena: &mut CausalExprArena,
expr: ExprId,
provider: &EmpiricalTableProvider,
expected: f64,
label: &str,
) {
let before = arena
.compile(expr)
.unwrap()
.evaluate(arena, provider, &EvalContext::default())
.unwrap();
let simplified = arena.simplify(expr).unwrap();
let after = arena
.compile(simplified)
.unwrap()
.evaluate(arena, provider, &EvalContext::default())
.unwrap();
assert!((before - after).abs() < 1e-12, "{label}: before={before} after={after}");
assert!((after - expected).abs() < 1e-12, "{label}: after={after}");
}
let mut arena = CausalExprArena::new();
let t = v(0);
let y = v(1);
let expr = arena.backdoor_ate(t, y, &[], f(1.0), f(0.0));
let mut p = EmpiricalTableProvider::new();
p.set_domain(y, [f(0.0), f(1.0)]);
p.set_domain(t, [f(0.0), f(1.0)]);
let empty_spec = FactorSpec {
variables: &[],
conditioned_on: &[],
intervention: &[],
domain: DomainRef::Observational,
};
p.insert_probability(&empty_spec, &Assignment::from_pairs([]), 1.0).unwrap();
for tlev in [0.0, 1.0] {
let ey = if (tlev - 1.0_f64).abs() < f64::EPSILON { 0.7 } else { 0.2 };
let interv = [InterventionAssignment { variable: t, value: f(tlev) }];
for (yval, prob) in [(1.0, ey), (0.0, 1.0 - ey)] {
let spec = FactorSpec {
variables: &[y],
conditioned_on: &[],
intervention: &interv,
domain: DomainRef::Interventional,
};
p.insert_probability(&spec, &Assignment::from_pairs([(y, f(yval))]), prob).unwrap();
}
}
assert_simplify_preserves(&mut arena, expr, &p, 0.5, "backdoor_empty_z");
}
#[test]
fn simplify_preserves_frontdoor_evaluation() {
fn assert_simplify_preserves(
arena: &mut CausalExprArena,
expr: ExprId,
provider: &EmpiricalTableProvider,
expected: f64,
label: &str,
) {
let before = arena
.compile(expr)
.unwrap()
.evaluate(arena, provider, &EvalContext::default())
.unwrap();
let simplified = arena.simplify(expr).unwrap();
let after = arena
.compile(simplified)
.unwrap()
.evaluate(arena, provider, &EvalContext::default())
.unwrap();
assert!((before - after).abs() < 1e-12, "{label}: before={before} after={after}");
assert!((after - expected).abs() < 1e-12, "{label}: after={after}");
}
let mut arena = CausalExprArena::new();
let t = v(0);
let y = v(1);
let m = v(2);
let expr = arena.frontdoor_ate(t, y, &[m], f(1.0), f(0.0));
let mut p = EmpiricalTableProvider::new();
p.set_domain(t, [f(0.0), f(1.0)]);
p.set_domain(y, [f(0.0), f(1.0)]);
p.set_domain(m, [f(0.0), f(1.0)]);
for (tval, prob) in [(0.0, 0.5), (1.0, 0.5)] {
let spec = FactorSpec {
variables: &[t],
conditioned_on: &[],
intervention: &[],
domain: DomainRef::Observational,
};
p.insert_probability(&spec, &Assignment::from_pairs([(t, f(tval))]), prob).unwrap();
}
for tlev in [0.0, 1.0] {
let pm1 = if (tlev - 1.0_f64).abs() < f64::EPSILON { 0.7 } else { 0.3 };
let interv = [InterventionAssignment { variable: t, value: f(tlev) }];
for (mval, prob) in [(1.0, pm1), (0.0, 1.0 - pm1)] {
let spec = FactorSpec {
variables: &[m],
conditioned_on: &[t],
intervention: &interv,
domain: DomainRef::Observational,
};
p.insert_probability(
&spec,
&Assignment::from_pairs([(m, f(mval)), (t, f(tlev))]),
prob,
)
.unwrap();
}
}
for tlev in [0.0, 1.0] {
for mlev in [0.0, 1.0] {
let py1 = if (mlev - 1.0_f64).abs() < f64::EPSILON { 0.9 } else { 0.1 };
for (yval, prob) in [(1.0, py1), (0.0, 1.0 - py1)] {
let spec = FactorSpec {
variables: &[y],
conditioned_on: &[t, m],
intervention: &[],
domain: DomainRef::Observational,
};
let assign = Assignment::from_pairs([(y, f(yval)), (m, f(mlev)), (t, f(tlev))]);
p.insert_probability(&spec, &assign, prob).unwrap();
}
}
}
assert_simplify_preserves(&mut arena, expr, &p, 0.32, "frontdoor");
}
#[test]
fn shallow_frontdoor_evaluates() {
let mut arena = CausalExprArena::new();
let t = v(0);
let y = v(1);
let m = v(2);
let expr = arena.frontdoor_ate(t, y, &[m], f(1.0), f(0.0));
let mut p = EmpiricalTableProvider::new();
p.set_domain(t, [f(0.0), f(1.0)]);
p.set_domain(y, [f(0.0), f(1.0)]);
p.set_domain(m, [f(0.0), f(1.0)]);
for (tval, prob) in [(0.0, 0.5), (1.0, 0.5)] {
let spec = FactorSpec {
variables: &[t],
conditioned_on: &[],
intervention: &[],
domain: DomainRef::Observational,
};
p.insert_probability(&spec, &Assignment::from_pairs([(t, f(tval))]), prob).unwrap();
}
for tlev in [0.0, 1.0] {
let pm1 = if (tlev - 1.0_f64).abs() < f64::EPSILON { 0.7 } else { 0.3 };
let interv = [InterventionAssignment { variable: t, value: f(tlev) }];
for (mval, prob) in [(1.0, pm1), (0.0, 1.0 - pm1)] {
let spec = FactorSpec {
variables: &[m],
conditioned_on: &[t],
intervention: &interv,
domain: DomainRef::Observational,
};
p.insert_probability(
&spec,
&Assignment::from_pairs([(m, f(mval)), (t, f(tlev))]),
prob,
)
.unwrap();
}
}
for tlev in [0.0, 1.0] {
for mlev in [0.0, 1.0] {
let py1 = if (mlev - 1.0_f64).abs() < f64::EPSILON { 0.9 } else { 0.1 };
for (yval, prob) in [(1.0, py1), (0.0, 1.0 - py1)] {
let spec = FactorSpec {
variables: &[y],
conditioned_on: &[t, m],
intervention: &[],
domain: DomainRef::Observational,
};
let assign = Assignment::from_pairs([(y, f(yval)), (m, f(mlev)), (t, f(tlev))]);
p.insert_probability(&spec, &assign, prob).unwrap();
}
}
}
let compiled = arena.compile(expr).unwrap();
let ate = compiled.evaluate(&arena, &p, &EvalContext::default()).unwrap();
assert!((ate - 0.32).abs() < 1e-12, "ate={ate}");
let simplified = arena.simplify(expr).unwrap();
let ate2 = arena
.compile(simplified)
.unwrap()
.evaluate(&arena, &p, &EvalContext::default())
.unwrap();
assert!((ate - ate2).abs() < 1e-12);
}
#[test]
fn discrete_integral_out_matches_sum_out() {
let mut arena = CausalExprArena::new();
let empty = arena.empty_var_set();
let empty_i = arena.empty_intervention_set();
let z = v(0);
let zset = arena.intern_var_set([z]);
let dist = arena.intern(ExprNode::Distribution {
variables: zset,
conditioned_on: empty,
intervention: empty_i,
domain: DomainRef::Observational,
});
let sum = arena.intern(ExprNode::SumOut { variables: zset, expr: dist });
let integ = arena.intern(ExprNode::IntegralOut { variables: zset, expr: dist });
let mut p = EmpiricalTableProvider::new();
p.set_domain(z, [f(0.0), f(1.0)]);
for (zval, prob) in [(0.0, 0.3), (1.0, 0.7)] {
let spec = FactorSpec {
variables: &[z],
conditioned_on: &[],
intervention: &[],
domain: DomainRef::Observational,
};
p.insert_probability(&spec, &Assignment::from_pairs([(z, f(zval))]), prob).unwrap();
}
let s = arena.compile(sum).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
let i =
arena.compile(integ).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
assert!((s - 1.0).abs() < 1e-12);
assert!((i - s).abs() < 1e-12);
}
#[test]
fn continuous_gaussian_integral_out_normalizes() {
use crate::provider::GaussianDensityProvider;
let mut arena = CausalExprArena::new();
let empty = arena.empty_var_set();
let empty_i = arena.empty_intervention_set();
let x = v(0);
let xset = arena.intern_var_set([x]);
let dist = arena.intern(ExprNode::Distribution {
variables: xset,
conditioned_on: empty,
intervention: empty_i,
domain: DomainRef::Observational,
});
let integ = arena.intern(ExprNode::IntegralOut { variables: xset, expr: dist });
let mut p = GaussianDensityProvider::new();
p.set_gaussian(x, 0.0, 1.0);
let mass =
arena.compile(integ).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
assert!((mass - 1.0).abs() < 1e-6, "∫ φ = {mass}");
}
#[test]
fn nested_integral_out_product_gaussian() {
use crate::provider::GaussianDensityProvider;
let mut arena = CausalExprArena::new();
let empty = arena.empty_var_set();
let empty_i = arena.empty_intervention_set();
let x = v(0);
let y = v(1);
let xset = arena.intern_var_set([x]);
let yset = arena.intern_var_set([y]);
let both = arena.intern_var_set([x, y]);
let dist = arena.intern(ExprNode::Distribution {
variables: both,
conditioned_on: empty,
intervention: empty_i,
domain: DomainRef::Observational,
});
let inner = arena.intern(ExprNode::IntegralOut { variables: yset, expr: dist });
let outer = arena.intern(ExprNode::IntegralOut { variables: xset, expr: inner });
let mut p = GaussianDensityProvider::new();
p.set_gaussian(x, 1.0, 0.25);
p.set_gaussian(y, -0.5, 4.0);
let mass =
arena.compile(outer).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
assert!((mass - 1.0).abs() < 1e-5, "∬ φ = {mass}");
}
#[test]
fn posterior_evaluate_batch() {
let mut arena = CausalExprArena::new();
let t = v(0);
let y = v(1);
let z = v(2);
let expr = arena.backdoor_ate(t, y, &[z], f(1.0), f(0.0));
let draw0 = backdoor_provider(t, y, z);
let mut draw1 = EmpiricalTableProvider::new();
draw1.set_domain(z, [f(0.0), f(1.0)]);
draw1.set_domain(y, [f(0.0), f(1.0)]);
draw1.set_domain(t, [f(0.0), f(1.0)]);
for (zval, prob) in [(0.0, 0.5), (1.0, 0.5)] {
let spec = FactorSpec {
variables: &[z],
conditioned_on: &[],
intervention: &[],
domain: DomainRef::Observational,
};
draw1.insert_probability(&spec, &Assignment::from_pairs([(z, f(zval))]), prob).unwrap();
}
for tlev in [0.0, 1.0] {
let interv = [InterventionAssignment { variable: t, value: f(tlev) }];
let py1 = tlev;
for zlev in [0.0, 1.0] {
for (yval, prob) in [(1.0, py1), (0.0, 1.0 - py1)] {
let spec = FactorSpec {
variables: &[y],
conditioned_on: &[z],
intervention: &interv,
domain: DomainRef::Interventional,
};
draw1
.insert_probability(
&spec,
&Assignment::from_pairs([(y, f(yval)), (z, f(zlev))]),
prob,
)
.unwrap();
}
}
}
let posterior = PosteriorDrawProvider::from_draws(vec![draw0, draw1]);
let compiled = arena.compile(expr).unwrap();
let batch = compiled.evaluate_batch(&arena, &posterior).unwrap();
assert_eq!(batch.len(), 2);
assert!((batch[0] - 0.45).abs() < 1e-12, "draw0={}", batch[0]);
assert!((batch[1] - 1.0).abs() < 1e-12, "draw1={}", batch[1]);
let single0 =
compiled.evaluate(&arena, &posterior, &EvalContext { draw: Some(0) }).unwrap();
let single1 =
compiled.evaluate(&arena, &posterior, &EvalContext { draw: Some(1) }).unwrap();
assert!((single0 - batch[0]).abs() < 1e-15);
assert!((single1 - batch[1]).abs() < 1e-15);
}
#[test]
fn expectation_of_simple_marginal() {
let mut arena = CausalExprArena::new();
let y = v(0);
let yset = arena.intern_var_set([y]);
let empty = arena.empty_var_set();
let empty_i = arena.empty_intervention_set();
let dist = arena.intern(ExprNode::Distribution {
variables: yset,
conditioned_on: empty,
intervention: empty_i,
domain: DomainRef::Observational,
});
let exp = arena.intern(ExprNode::Expectation {
function: OutcomeExprId::identity(y),
distribution: dist,
});
let mut p = EmpiricalTableProvider::new();
p.set_domain(y, [f(0.0), f(2.0)]);
let spec = FactorSpec {
variables: &[y],
conditioned_on: &[],
intervention: &[],
domain: DomainRef::Observational,
};
p.insert_probability(&spec, &Assignment::from_pairs([(y, f(0.0))]), 0.25).unwrap();
p.insert_probability(&spec, &Assignment::from_pairs([(y, f(2.0))]), 0.75).unwrap();
let val =
arena.compile(exp).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
assert!((val - 1.5).abs() < 1e-12);
}
#[test]
fn evaluation_is_stable_across_repeated_calls() {
let mut arena = CausalExprArena::new();
let t = v(0);
let y = v(1);
let z = v(2);
let expr = arena.backdoor_ate(t, y, &[z], f(1.0), f(0.0));
let provider = backdoor_provider(t, y, z);
let compiled = arena.compile(expr).unwrap();
let first = compiled.evaluate(&arena, &provider, &EvalContext::default()).unwrap();
let second = compiled.evaluate(&arena, &provider, &EvalContext::default()).unwrap();
let third = compiled.evaluate(&arena, &provider, &EvalContext::default()).unwrap();
assert_eq!(first.to_bits(), second.to_bits());
assert_eq!(first.to_bits(), third.to_bits());
assert!((first - 0.45).abs() < 1e-12, "ate={first}");
}
#[test]
fn scoped_intervention_binding_restores_between_siblings() {
let mut arena = CausalExprArena::new();
let z = v(0);
let zset = arena.intern_var_set([z]);
let empty = arena.empty_var_set();
let empty_i = arena.empty_intervention_set();
let do_z1 = arena.intern_intervention_assignments([InterventionAssignment {
variable: z,
value: f(1.0),
}]);
let shadowed = arena.intern(ExprNode::Distribution {
variables: empty,
conditioned_on: zset,
intervention: do_z1,
domain: DomainRef::Observational,
});
let z_marginal = arena.intern(ExprNode::Distribution {
variables: zset,
conditioned_on: empty,
intervention: empty_i,
domain: DomainRef::Observational,
});
let product = {
let list = arena.intern_list([shadowed, z_marginal]);
arena.intern(ExprNode::Product(list))
};
let sum = arena.intern(ExprNode::SumOut { variables: zset, expr: product });
let mut p = EmpiricalTableProvider::new();
p.set_domain(z, [f(0.0), f(1.0)]);
let interv = [InterventionAssignment { variable: z, value: f(1.0) }];
let shadow_spec = FactorSpec {
variables: &[],
conditioned_on: &[z],
intervention: &interv,
domain: DomainRef::Observational,
};
p.insert_probability(&shadow_spec, &Assignment::from_pairs([(z, f(1.0))]), 2.0).unwrap();
let marg_spec = FactorSpec {
variables: &[z],
conditioned_on: &[],
intervention: &[],
domain: DomainRef::Observational,
};
p.insert_probability(&marg_spec, &Assignment::from_pairs([(z, f(0.0))]), 0.3).unwrap();
p.insert_probability(&marg_spec, &Assignment::from_pairs([(z, f(1.0))]), 0.7).unwrap();
let val =
arena.compile(sum).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
assert!((val - 2.0).abs() < 1e-12, "val={val}");
}
#[test]
fn expectation_respects_env_bound_conditioning() {
let mut arena = CausalExprArena::new();
let y = v(0);
let z = v(1);
let yset = arena.intern_var_set([y]);
let zset = arena.intern_var_set([z]);
let empty_i = arena.empty_intervention_set();
let dist = arena.intern(ExprNode::Distribution {
variables: yset,
conditioned_on: zset,
intervention: empty_i,
domain: DomainRef::Observational,
});
let exp = arena.intern(ExprNode::Expectation {
function: OutcomeExprId::identity(y),
distribution: dist,
});
let mut p = EmpiricalTableProvider::new();
p.set_domain(y, [f(0.0), f(2.0)]);
p.set_domain(z, [f(0.0), f(1.0)]);
let spec = FactorSpec {
variables: &[y],
conditioned_on: &[z],
intervention: &[],
domain: DomainRef::Observational,
};
for (yv, zv, prob) in [(0.0, 0.0, 0.25), (2.0, 0.0, 0.75), (0.0, 1.0, 1.0), (2.0, 1.0, 0.0)]
{
p.insert_probability(&spec, &Assignment::from_pairs([(y, f(yv)), (z, f(zv))]), prob)
.unwrap();
}
let compiled = arena.compile(exp).unwrap();
let env0 = Assignment::from_pairs([(z, f(0.0))]);
let e0 = compiled.evaluate_with(&arena, &p, &EvalContext::default(), &env0).unwrap();
assert!((e0 - 1.5).abs() < 1e-12, "E[Y|z=0]={e0}");
let env1 = Assignment::from_pairs([(z, f(1.0))]);
let e1 = compiled.evaluate_with(&arena, &p, &EvalContext::default(), &env1).unwrap();
assert!(e1.abs() < 1e-12, "E[Y|z=1]={e1}");
assert_eq!(env0.entries(), &[(z, f(0.0))]);
}
#[test]
fn ratio_zero_denominator_is_division_by_zero() {
let mut arena = CausalExprArena::new();
let empty = arena.empty_var_set();
let empty_i = arena.empty_intervention_set();
let numerator = arena.intern(ExprNode::Distribution {
variables: empty,
conditioned_on: empty,
intervention: empty_i,
domain: DomainRef::Observational,
});
let denominator = arena.intern(ExprNode::Distribution {
variables: empty,
conditioned_on: empty,
intervention: empty_i,
domain: DomainRef::Interventional,
});
let ratio = arena.intern(ExprNode::Ratio { numerator, denominator });
let mut p = EmpiricalTableProvider::new();
let obs_spec = FactorSpec {
variables: &[],
conditioned_on: &[],
intervention: &[],
domain: DomainRef::Observational,
};
let interv_spec = FactorSpec {
variables: &[],
conditioned_on: &[],
intervention: &[],
domain: DomainRef::Interventional,
};
p.insert_probability(&obs_spec, &Assignment::from_pairs([]), 3.0).unwrap();
p.insert_probability(&interv_spec, &Assignment::from_pairs([]), 0.0).unwrap();
let err = arena
.compile(ratio)
.unwrap()
.evaluate(&arena, &p, &EvalContext::default())
.unwrap_err();
assert_eq!(err, EvalError::DivisionByZero);
}
}