Skip to main content

antecedent_expr/
eval.rs

1//! Compiled topological evaluators for causal expressions.
2//!
3//! SPDX-License-Identifier: MIT OR Apache-2.0
4
5use std::collections::HashMap;
6use std::sync::Arc;
7
8use antecedent_core::{Value, VariableId};
9
10use crate::provider::{Assignment, DistributionProvider, EvalContext, EvalError, FactorSpec};
11use crate::{
12    CausalExprArena, ContrastOp, DomainRef, ExprId, ExprNode, InterventionSetId, OutcomeExprId,
13    VarSetId,
14};
15
16/// One step in a compiled evaluation plan (child references are slot indices).
17#[derive(Clone, Debug)]
18enum EvalOp {
19    Distribution {
20        variables: VarSetId,
21        conditioned_on: VarSetId,
22        intervention: InterventionSetId,
23        domain: DomainRef,
24    },
25    Product {
26        children: Arc<[usize]>,
27    },
28    SumOut {
29        variables: VarSetId,
30        body: usize,
31    },
32    IntegralOut {
33        variables: VarSetId,
34        body: usize,
35    },
36    Ratio {
37        numerator: usize,
38        denominator: usize,
39    },
40    Expectation {
41        function: OutcomeExprId,
42        distribution: usize,
43    },
44    Contrast {
45        left: usize,
46        right: usize,
47        op: ContrastOp,
48    },
49}
50
51/// Topologically ordered compiled evaluator for repeated provider evaluation.
52#[derive(Clone, Debug)]
53pub struct CompiledEvaluator {
54    ops: Vec<EvalOp>,
55    /// Sorted, deduplicated free variables per slot. A static property of the
56    /// plan, computed once at compile time; `Expectation` evaluation reads it
57    /// on every call instead of re-deriving it per evaluation.
58    free_vars: Vec<Arc<[VariableId]>>,
59    root: usize,
60}
61
62impl CausalExprArena {
63    /// Compile `root` into a topological evaluation plan.
64    ///
65    /// Continuous [`ExprNode::IntegralOut`] compiles successfully; evaluation uses
66    /// [`DistributionProvider::quadrature`] or discrete [`DistributionProvider::support`].
67    pub fn compile(&self, root: ExprId) -> Result<CompiledEvaluator, EvalError> {
68        CompiledEvaluator::compile(self, root)
69    }
70}
71
72impl CompiledEvaluator {
73    /// Compile an expression DAG into slot-addressed ops (post-order).
74    ///
75    /// Continuous [`ExprNode::IntegralOut`] is supported (see [`CausalExprArena::compile`]).
76    pub fn compile(arena: &CausalExprArena, root: ExprId) -> Result<Self, EvalError> {
77        let mut ops = Vec::new();
78        let mut expr_to_slot = HashMap::new();
79        let root_slot = compile_rec(arena, root, &mut ops, &mut expr_to_slot)?;
80        let free_vars = compute_free_vars(&ops, arena);
81        Ok(Self { ops, free_vars, root: root_slot })
82    }
83
84    /// Evaluate once against a provider.
85    ///
86    /// # Errors
87    ///
88    /// Provider / numeric failures.
89    pub fn evaluate(
90        &self,
91        arena: &CausalExprArena,
92        provider: &dyn DistributionProvider,
93        ctx: &EvalContext,
94    ) -> Result<f64, EvalError> {
95        self.evaluate_with(arena, provider, ctx, &Assignment::new())
96    }
97
98    /// Evaluate with an initial variable binding (e.g. `do(X=x)` and outcome levels).
99    ///
100    /// # Errors
101    ///
102    /// Provider / numeric failures.
103    pub fn evaluate_with(
104        &self,
105        arena: &CausalExprArena,
106        provider: &dyn DistributionProvider,
107        ctx: &EvalContext,
108        env: &Assignment,
109    ) -> Result<f64, EvalError> {
110        // One clone per evaluation: `eval_slot` threads a single mutable
111        // scratch assignment through the whole plan, with each scope binding
112        // and restoring its own variables (see `with_scoped_bindings`) rather
113        // than cloning the assignment per support row.
114        let mut scratch = env.clone();
115        self.eval_slot(arena, provider, ctx, &mut scratch, self.root)
116    }
117
118    /// Evaluate over all posterior draws (`provider.n_draws()`), or a single
119    /// empirical evaluation when `n_draws` is `None`.
120    ///
121    /// # Errors
122    ///
123    /// Provider / numeric failures.
124    pub fn evaluate_batch(
125        &self,
126        arena: &CausalExprArena,
127        provider: &dyn DistributionProvider,
128    ) -> Result<Vec<f64>, EvalError> {
129        match provider.n_draws() {
130            None => Ok(vec![self.evaluate(arena, provider, &EvalContext::default())?]),
131            Some(n) => {
132                let mut out = Vec::with_capacity(n);
133                for draw in 0..n {
134                    let ctx = EvalContext { draw: Some(draw) };
135                    out.push(self.evaluate(arena, provider, &ctx)?);
136                }
137                Ok(out)
138            }
139        }
140    }
141
142    fn eval_slot(
143        &self,
144        arena: &CausalExprArena,
145        provider: &dyn DistributionProvider,
146        ctx: &EvalContext,
147        env: &mut Assignment,
148        slot: usize,
149    ) -> Result<f64, EvalError> {
150        // Density / scalar under `env`. Expectations and contrasts are scalars;
151        // other ops are densities in the free variables bound by `env`.
152        match &self.ops[slot] {
153            EvalOp::Distribution { variables, conditioned_on, intervention, domain } => {
154                let spec = FactorSpec {
155                    variables: arena.var_set(*variables),
156                    conditioned_on: arena.var_set(*conditioned_on),
157                    intervention: arena.intervention_assignments(*intervention),
158                    domain: *domain,
159                };
160                // Interventions bind targets; bind them into the shared
161                // scratch assignment for the lookup, restored on exit.
162                with_scoped_bindings(env, spec.intervention.iter().map(|a| a.variable), |env| {
163                    for a in spec.intervention {
164                        env.set(a.variable, a.value.clone());
165                    }
166                    provider.probability(&spec, env, ctx)
167                })
168            }
169            EvalOp::Product { children } => {
170                let mut prod = 1.0;
171                for &c in children.iter() {
172                    prod *= self.eval_slot(arena, provider, ctx, env, c)?;
173                }
174                Ok(prod)
175            }
176            EvalOp::SumOut { variables, body } => {
177                self.eval_sum_out(arena, provider, ctx, env, *variables, *body)
178            }
179            EvalOp::IntegralOut { variables, body } => {
180                self.eval_integral_out(arena, provider, ctx, env, *variables, *body)
181            }
182            EvalOp::Ratio { numerator, denominator } => {
183                let num = self.eval_slot(arena, provider, ctx, env, *numerator)?;
184                let den = self.eval_slot(arena, provider, ctx, env, *denominator)?;
185                if den == 0.0 {
186                    return Err(EvalError::DivisionByZero);
187                }
188                Ok(num / den)
189            }
190            EvalOp::Expectation { function, distribution } => {
191                self.eval_expectation(arena, provider, ctx, env, function.variable(), *distribution)
192            }
193            EvalOp::Contrast { left, right, op } => {
194                let l = self.eval_slot(arena, provider, ctx, env, *left)?;
195                let r = self.eval_slot(arena, provider, ctx, env, *right)?;
196                match op {
197                    ContrastOp::Difference => Ok(l - r),
198                }
199            }
200        }
201    }
202
203    fn eval_sum_out(
204        &self,
205        arena: &CausalExprArena,
206        provider: &dyn DistributionProvider,
207        ctx: &EvalContext,
208        env: &mut Assignment,
209        variables: VarSetId,
210        body: usize,
211    ) -> Result<f64, EvalError> {
212        let vars = arena.var_set(variables);
213        let rows = provider.support(vars, ctx)?;
214        with_scoped_bindings(env, vars.iter().copied(), |env| {
215            let mut sum = 0.0;
216            for row in rows.iter() {
217                if row.len() != vars.len() {
218                    return Err(EvalError::SupportShape {
219                        expected: vars.len(),
220                        actual: row.len(),
221                    });
222                }
223                for (i, &v) in vars.iter().enumerate() {
224                    env.set(v, row[i].clone());
225                }
226                sum += self.eval_slot(arena, provider, ctx, env, body)?;
227            }
228            Ok(sum)
229        })
230    }
231
232    fn eval_integral_out(
233        &self,
234        arena: &CausalExprArena,
235        provider: &dyn DistributionProvider,
236        ctx: &EvalContext,
237        env: &mut Assignment,
238        variables: VarSetId,
239        body: usize,
240    ) -> Result<f64, EvalError> {
241        let vars = arena.var_set(variables);
242        if let Some(nodes) = provider.quadrature(vars, ctx)? {
243            return with_scoped_bindings(env, vars.iter().copied(), |env| {
244                let mut acc = 0.0;
245                for (row, weight) in nodes.iter() {
246                    if row.len() != vars.len() {
247                        return Err(EvalError::SupportShape {
248                            expected: vars.len(),
249                            actual: row.len(),
250                        });
251                    }
252                    for (i, &v) in vars.iter().enumerate() {
253                        env.set(v, row[i].clone());
254                    }
255                    acc += *weight * self.eval_slot(arena, provider, ctx, env, body)?;
256                }
257                Ok(acc)
258            });
259        }
260        // Discrete / counting-measure fallback (IntegralOut ≡ SumOut).
261        let rows = provider.support(vars, ctx).map_err(|e| match e {
262            EvalError::EmptySupport(_) => EvalError::UnsupportedIntegralOut,
263            other => other,
264        })?;
265        with_scoped_bindings(env, vars.iter().copied(), |env| {
266            let mut sum = 0.0;
267            for row in rows.iter() {
268                if row.len() != vars.len() {
269                    return Err(EvalError::SupportShape {
270                        expected: vars.len(),
271                        actual: row.len(),
272                    });
273                }
274                for (i, &v) in vars.iter().enumerate() {
275                    env.set(v, row[i].clone());
276                }
277                sum += self.eval_slot(arena, provider, ctx, env, body)?;
278            }
279            Ok(sum)
280        })
281    }
282
283    fn eval_expectation(
284        &self,
285        arena: &CausalExprArena,
286        provider: &dyn DistributionProvider,
287        ctx: &EvalContext,
288        env: &mut Assignment,
289        outcome_var: VariableId,
290        distribution: usize,
291    ) -> Result<f64, EvalError> {
292        // E[f | D] = Σ_{x ∈ support(free(D))} f(x) · dens(D, x)
293        // Free variables per slot are precomputed at compile time; only the
294        // env-dependent filtering happens per evaluation.
295        let free = &self.free_vars[distribution];
296        let mut enum_vars: Vec<VariableId> =
297            free.iter().copied().filter(|v| env.get(*v).is_none()).collect();
298        if !enum_vars.contains(&outcome_var) && env.get(outcome_var).is_none() {
299            enum_vars.push(outcome_var);
300        }
301        enum_vars.sort_by_key(|v| v.raw());
302        enum_vars.dedup();
303
304        if enum_vars.is_empty() {
305            let dens = self.eval_slot(arena, provider, ctx, env, distribution)?;
306            let y = provider.outcome(outcome_var, env, ctx)?;
307            return Ok(y * dens);
308        }
309
310        let rows = provider.support(&enum_vars, ctx)?;
311        with_scoped_bindings(env, enum_vars.iter().copied(), |env| {
312            let mut acc = 0.0;
313            for row in rows.iter() {
314                if row.len() != enum_vars.len() {
315                    return Err(EvalError::SupportShape {
316                        expected: enum_vars.len(),
317                        actual: row.len(),
318                    });
319                }
320                for (i, &v) in enum_vars.iter().enumerate() {
321                    env.set(v, row[i].clone());
322                }
323                let dens = self.eval_slot(arena, provider, ctx, env, distribution)?;
324                let y = provider.outcome(outcome_var, env, ctx)?;
325                acc += y * dens;
326            }
327            Ok(acc)
328        })
329    }
330}
331
332/// Run `f` against the shared scratch assignment, then restore any prior
333/// bindings of `vars` (removing bindings that did not exist before).
334///
335/// Evaluation bindings are strictly stack-scoped — sum/integral/expectation
336/// rows and intervention targets shadow outer bindings only for the duration
337/// of the nested evaluation — so saving and restoring just those variables is
338/// observationally identical to the previous clone-per-row scheme, without the
339/// per-row `Assignment` clone. Restoration also runs on the error path so a
340/// failed inner evaluation leaves the scratch assignment as it found it.
341fn with_scoped_bindings<T>(
342    env: &mut Assignment,
343    vars: impl IntoIterator<Item = VariableId>,
344    f: impl FnOnce(&mut Assignment) -> Result<T, EvalError>,
345) -> Result<T, EvalError> {
346    let saved: Vec<(VariableId, Option<Value>)> =
347        vars.into_iter().map(|v| (v, env.get(v).cloned())).collect();
348    let result = f(env);
349    for (v, prev) in saved {
350        match prev {
351            Some(value) => env.set(v, value),
352            None => {
353                env.remove(v);
354            }
355        }
356    }
357    result
358}
359
360fn compile_rec(
361    arena: &CausalExprArena,
362    id: ExprId,
363    ops: &mut Vec<EvalOp>,
364    expr_to_slot: &mut HashMap<u32, usize>,
365) -> Result<usize, EvalError> {
366    if let Some(&slot) = expr_to_slot.get(&id.raw()) {
367        return Ok(slot);
368    }
369    let op = match arena.node(id).clone() {
370        ExprNode::Distribution { variables, conditioned_on, intervention, domain } => {
371            EvalOp::Distribution { variables, conditioned_on, intervention, domain }
372        }
373        ExprNode::Product(list) => {
374            let mut children = Vec::new();
375            for &c in arena.list(list) {
376                children.push(compile_rec(arena, c, ops, expr_to_slot)?);
377            }
378            EvalOp::Product { children: Arc::from(children) }
379        }
380        ExprNode::SumOut { variables, expr } => {
381            let body = compile_rec(arena, expr, ops, expr_to_slot)?;
382            EvalOp::SumOut { variables, body }
383        }
384        ExprNode::IntegralOut { variables, expr } => {
385            let body = compile_rec(arena, expr, ops, expr_to_slot)?;
386            EvalOp::IntegralOut { variables, body }
387        }
388        ExprNode::Ratio { numerator, denominator } => {
389            let n = compile_rec(arena, numerator, ops, expr_to_slot)?;
390            let d = compile_rec(arena, denominator, ops, expr_to_slot)?;
391            EvalOp::Ratio { numerator: n, denominator: d }
392        }
393        ExprNode::Expectation { function, distribution } => {
394            let dist = compile_rec(arena, distribution, ops, expr_to_slot)?;
395            EvalOp::Expectation { function, distribution: dist }
396        }
397        ExprNode::Contrast { left, right, op } => {
398            let l = compile_rec(arena, left, ops, expr_to_slot)?;
399            let r = compile_rec(arena, right, ops, expr_to_slot)?;
400            EvalOp::Contrast { left: l, right: r, op }
401        }
402    };
403    let slot = ops.len();
404    ops.push(op);
405    expr_to_slot.insert(id.raw(), slot);
406    Ok(slot)
407}
408
409/// Per-slot free variables (sorted, deduplicated), computed once per compile.
410///
411/// Slots are emitted post-order by `compile_rec`, so every child index is
412/// smaller than its parent's and a single forward pass suffices.
413///
414/// The `Distribution` arm must agree with `simplify::free_vars` (see the
415/// comment there): `conditioned_on` variables bound by the accompanying
416/// `intervention` set are do(·)-fixed, not free.
417fn compute_free_vars(ops: &[EvalOp], arena: &CausalExprArena) -> Vec<Arc<[VariableId]>> {
418    let mut out: Vec<Arc<[VariableId]>> = Vec::with_capacity(ops.len());
419    for op in ops {
420        let mut vars: Vec<VariableId> = match op {
421            EvalOp::Distribution { variables, conditioned_on, intervention, .. } => {
422                let mut vars = arena.var_set(*variables).to_vec();
423                let bound = arena.intervention_assignments(*intervention);
424                for &v in arena.var_set(*conditioned_on) {
425                    if !bound.iter().any(|a| a.variable == v) {
426                        vars.push(v);
427                    }
428                }
429                vars
430            }
431            EvalOp::Product { children } => {
432                children.iter().flat_map(|&c| out[c].iter().copied()).collect()
433            }
434            EvalOp::SumOut { variables, body } | EvalOp::IntegralOut { variables, body } => {
435                let bound = arena.var_set(*variables);
436                out[*body].iter().copied().filter(|v| !bound.contains(v)).collect()
437            }
438            EvalOp::Ratio { numerator, denominator } => {
439                out[*numerator].iter().chain(out[*denominator].iter()).copied().collect()
440            }
441            EvalOp::Expectation { function, distribution } => {
442                let mut vars = out[*distribution].to_vec();
443                vars.push(function.variable());
444                vars
445            }
446            EvalOp::Contrast { left, right, .. } => {
447                out[*left].iter().chain(out[*right].iter()).copied().collect()
448            }
449        };
450        vars.sort_by_key(|v| v.raw());
451        vars.dedup();
452        out.push(Arc::from(vars));
453    }
454    out
455}
456
457#[cfg(test)]
458mod tests {
459    use super::*;
460    use crate::provider::{EmpiricalTableProvider, PosteriorDrawProvider};
461    use crate::{InterventionAssignment, OutcomeExprId};
462    use antecedent_core::Value;
463
464    fn v(id: u32) -> VariableId {
465        VariableId::from_raw(id)
466    }
467
468    fn f(x: f64) -> Value {
469        Value::f64(x)
470    }
471
472    /// Binary confounder Z, binary Y; backdoor ATE = 0.45.
473    fn backdoor_provider(t: VariableId, y: VariableId, z: VariableId) -> EmpiricalTableProvider {
474        let mut p = EmpiricalTableProvider::new();
475        p.set_domain(z, [f(0.0), f(1.0)]);
476        p.set_domain(y, [f(0.0), f(1.0)]);
477        p.set_domain(t, [f(0.0), f(1.0)]);
478
479        // P(Z)
480        for (zval, prob) in [(0.0, 0.5), (1.0, 0.5)] {
481            let spec = FactorSpec {
482                variables: &[z],
483                conditioned_on: &[],
484                intervention: &[],
485                domain: DomainRef::Observational,
486            };
487            let assign = Assignment::from_pairs([(z, f(zval))]);
488            p.insert_probability(&spec, &assign, prob).unwrap();
489        }
490
491        // P(Y | Z, do(T=t)) = P(Y | T=t, Z) under backdoor.
492        // E[Y|T=1,Z=0]=0.8, E[Y|T=1,Z=1]=0.6, E[Y|T=0,Z=0]=0.3, E[Y|T=0,Z=1]=0.2
493        let ey = |tlev: f64, zlev: f64| -> f64 {
494            match (tlev.to_bits(), zlev.to_bits()) {
495                (t, z) if t == 1.0f64.to_bits() && z == 0.0f64.to_bits() => 0.8,
496                (t, z) if t == 1.0f64.to_bits() && z == 1.0f64.to_bits() => 0.6,
497                (t, z) if t == 0.0f64.to_bits() && z == 0.0f64.to_bits() => 0.3,
498                (t, z) if t == 0.0f64.to_bits() && z == 1.0f64.to_bits() => 0.2,
499                _ => panic!("bad levels"),
500            }
501        };
502        for tlev in [0.0, 1.0] {
503            let interv = [InterventionAssignment { variable: t, value: f(tlev) }];
504            for zlev in [0.0, 1.0] {
505                let p_y1 = ey(tlev, zlev);
506                for (yval, prob) in [(1.0, p_y1), (0.0, 1.0 - p_y1)] {
507                    let spec = FactorSpec {
508                        variables: &[y],
509                        conditioned_on: &[z],
510                        intervention: &interv,
511                        domain: DomainRef::Interventional,
512                    };
513                    let assign = Assignment::from_pairs([(y, f(yval)), (z, f(zlev))]);
514                    p.insert_probability(&spec, &assign, prob).unwrap();
515                }
516            }
517        }
518        p
519    }
520
521    #[test]
522    fn backdoor_ate_matches_closed_form() {
523        let mut arena = CausalExprArena::new();
524        let t = v(0);
525        let y = v(1);
526        let z = v(2);
527        let expr = arena.backdoor_ate(t, y, &[z], f(1.0), f(0.0));
528        let provider = backdoor_provider(t, y, z);
529        let compiled = arena.compile(expr).unwrap();
530        let ate = compiled.evaluate(&arena, &provider, &EvalContext::default()).unwrap();
531        assert!((ate - 0.45).abs() < 1e-12, "ate={ate}");
532    }
533
534    #[test]
535    fn simplify_preserves_backdoor_evaluation() {
536        let mut arena = CausalExprArena::new();
537        let t = v(0);
538        let y = v(1);
539        let z = v(2);
540        let expr = arena.backdoor_ate(t, y, &[z], f(1.0), f(0.0));
541        let provider = backdoor_provider(t, y, z);
542        let before = arena
543            .compile(expr)
544            .unwrap()
545            .evaluate(&arena, &provider, &EvalContext::default())
546            .unwrap();
547        let simplified = arena.simplify(expr).unwrap();
548        let after = arena
549            .compile(simplified)
550            .unwrap()
551            .evaluate(&arena, &provider, &EvalContext::default())
552            .unwrap();
553        assert!((before - after).abs() < 1e-12, "before={before} after={after}");
554        assert!((after - 0.45).abs() < 1e-12);
555    }
556
557    /// Empty adjustment (second Z set): simplify must preserve numeric eval.
558    #[test]
559    fn simplify_preserves_backdoor_empty_evaluation() {
560        fn assert_simplify_preserves(
561            arena: &mut CausalExprArena,
562            expr: ExprId,
563            provider: &EmpiricalTableProvider,
564            expected: f64,
565            label: &str,
566        ) {
567            let before = arena
568                .compile(expr)
569                .unwrap()
570                .evaluate(arena, provider, &EvalContext::default())
571                .unwrap();
572            let simplified = arena.simplify(expr).unwrap();
573            let after = arena
574                .compile(simplified)
575                .unwrap()
576                .evaluate(arena, provider, &EvalContext::default())
577                .unwrap();
578            assert!((before - after).abs() < 1e-12, "{label}: before={before} after={after}");
579            assert!((after - expected).abs() < 1e-12, "{label}: after={after}");
580        }
581
582        // Backdoor with empty Z: E[Y|do(1)]=0.7, E[Y|do(0)]=0.2 → ATE = 0.5.
583        // Exercises simplify.empty_sum_out / singleton product on the adjustment set.
584        let mut arena = CausalExprArena::new();
585        let t = v(0);
586        let y = v(1);
587        let expr = arena.backdoor_ate(t, y, &[], f(1.0), f(0.0));
588        let mut p = EmpiricalTableProvider::new();
589        p.set_domain(y, [f(0.0), f(1.0)]);
590        p.set_domain(t, [f(0.0), f(1.0)]);
591        // Vacuous P(∅) factor from empty adjustment marginal.
592        let empty_spec = FactorSpec {
593            variables: &[],
594            conditioned_on: &[],
595            intervention: &[],
596            domain: DomainRef::Observational,
597        };
598        p.insert_probability(&empty_spec, &Assignment::from_pairs([]), 1.0).unwrap();
599        for tlev in [0.0, 1.0] {
600            let ey = if (tlev - 1.0_f64).abs() < f64::EPSILON { 0.7 } else { 0.2 };
601            let interv = [InterventionAssignment { variable: t, value: f(tlev) }];
602            for (yval, prob) in [(1.0, ey), (0.0, 1.0 - ey)] {
603                let spec = FactorSpec {
604                    variables: &[y],
605                    conditioned_on: &[],
606                    intervention: &interv,
607                    domain: DomainRef::Interventional,
608                };
609                p.insert_probability(&spec, &Assignment::from_pairs([(y, f(yval))]), prob).unwrap();
610            }
611        }
612        assert_simplify_preserves(&mut arena, expr, &p, 0.5, "backdoor_empty_z");
613    }
614
615    /// Frontdoor: simplify must preserve numeric eval.
616    #[test]
617    fn simplify_preserves_frontdoor_evaluation() {
618        fn assert_simplify_preserves(
619            arena: &mut CausalExprArena,
620            expr: ExprId,
621            provider: &EmpiricalTableProvider,
622            expected: f64,
623            label: &str,
624        ) {
625            let before = arena
626                .compile(expr)
627                .unwrap()
628                .evaluate(arena, provider, &EvalContext::default())
629                .unwrap();
630            let simplified = arena.simplify(expr).unwrap();
631            let after = arena
632                .compile(simplified)
633                .unwrap()
634                .evaluate(arena, provider, &EvalContext::default())
635                .unwrap();
636            assert!((before - after).abs() < 1e-12, "{label}: before={before} after={after}");
637            assert!((after - expected).abs() < 1e-12, "{label}: after={after}");
638        }
639
640        // Frontdoor (same tables as shallow_frontdoor_evaluates): ATE = 0.32.
641        let mut arena = CausalExprArena::new();
642        let t = v(0);
643        let y = v(1);
644        let m = v(2);
645        let expr = arena.frontdoor_ate(t, y, &[m], f(1.0), f(0.0));
646        let mut p = EmpiricalTableProvider::new();
647        p.set_domain(t, [f(0.0), f(1.0)]);
648        p.set_domain(y, [f(0.0), f(1.0)]);
649        p.set_domain(m, [f(0.0), f(1.0)]);
650        for (tval, prob) in [(0.0, 0.5), (1.0, 0.5)] {
651            let spec = FactorSpec {
652                variables: &[t],
653                conditioned_on: &[],
654                intervention: &[],
655                domain: DomainRef::Observational,
656            };
657            p.insert_probability(&spec, &Assignment::from_pairs([(t, f(tval))]), prob).unwrap();
658        }
659        for tlev in [0.0, 1.0] {
660            let pm1 = if (tlev - 1.0_f64).abs() < f64::EPSILON { 0.7 } else { 0.3 };
661            let interv = [InterventionAssignment { variable: t, value: f(tlev) }];
662            for (mval, prob) in [(1.0, pm1), (0.0, 1.0 - pm1)] {
663                let spec = FactorSpec {
664                    variables: &[m],
665                    conditioned_on: &[t],
666                    intervention: &interv,
667                    domain: DomainRef::Observational,
668                };
669                p.insert_probability(
670                    &spec,
671                    &Assignment::from_pairs([(m, f(mval)), (t, f(tlev))]),
672                    prob,
673                )
674                .unwrap();
675            }
676        }
677        for tlev in [0.0, 1.0] {
678            for mlev in [0.0, 1.0] {
679                let py1 = if (mlev - 1.0_f64).abs() < f64::EPSILON { 0.9 } else { 0.1 };
680                for (yval, prob) in [(1.0, py1), (0.0, 1.0 - py1)] {
681                    let spec = FactorSpec {
682                        variables: &[y],
683                        conditioned_on: &[t, m],
684                        intervention: &[],
685                        domain: DomainRef::Observational,
686                    };
687                    let assign = Assignment::from_pairs([(y, f(yval)), (m, f(mlev)), (t, f(tlev))]);
688                    p.insert_probability(&spec, &assign, prob).unwrap();
689                }
690            }
691        }
692        assert_simplify_preserves(&mut arena, expr, &p, 0.32, "frontdoor");
693    }
694
695    #[test]
696    fn shallow_frontdoor_evaluates() {
697        // Minimal front-door: T→M→Y with no hidden confounding encoded in tables.
698        // P(M|T=t); P(Y|M,T'); P(T').
699        let mut arena = CausalExprArena::new();
700        let t = v(0);
701        let y = v(1);
702        let m = v(2);
703        let expr = arena.frontdoor_ate(t, y, &[m], f(1.0), f(0.0));
704
705        let mut p = EmpiricalTableProvider::new();
706        p.set_domain(t, [f(0.0), f(1.0)]);
707        p.set_domain(y, [f(0.0), f(1.0)]);
708        p.set_domain(m, [f(0.0), f(1.0)]);
709
710        // P(T')
711        for (tval, prob) in [(0.0, 0.5), (1.0, 0.5)] {
712            let spec = FactorSpec {
713                variables: &[t],
714                conditioned_on: &[],
715                intervention: &[],
716                domain: DomainRef::Observational,
717            };
718            p.insert_probability(&spec, &Assignment::from_pairs([(t, f(tval))]), prob).unwrap();
719        }
720
721        // P(M | T=t): P(M=1|T=1)=0.7, P(M=1|T=0)=0.3 (FD condition 2).
722        for tlev in [0.0, 1.0] {
723            let pm1 = if (tlev - 1.0_f64).abs() < f64::EPSILON { 0.7 } else { 0.3 };
724            let interv = [InterventionAssignment { variable: t, value: f(tlev) }];
725            for (mval, prob) in [(1.0, pm1), (0.0, 1.0 - pm1)] {
726                let spec = FactorSpec {
727                    variables: &[m],
728                    conditioned_on: &[t],
729                    intervention: &interv,
730                    domain: DomainRef::Observational,
731                };
732                p.insert_probability(
733                    &spec,
734                    &Assignment::from_pairs([(m, f(mval)), (t, f(tlev))]),
735                    prob,
736                )
737                .unwrap();
738            }
739        }
740
741        // P(Y | M, T'): E[Y|M=1,*]=0.9, E[Y|M=0,*]=0.1 (T' irrelevant)
742        // Arena sorts m_and_t as [t, m] when t.raw() < m.raw().
743        for tlev in [0.0, 1.0] {
744            for mlev in [0.0, 1.0] {
745                let py1 = if (mlev - 1.0_f64).abs() < f64::EPSILON { 0.9 } else { 0.1 };
746                for (yval, prob) in [(1.0, py1), (0.0, 1.0 - py1)] {
747                    let spec = FactorSpec {
748                        variables: &[y],
749                        conditioned_on: &[t, m],
750                        intervention: &[],
751                        domain: DomainRef::Observational,
752                    };
753                    let assign = Assignment::from_pairs([(y, f(yval)), (m, f(mlev)), (t, f(tlev))]);
754                    p.insert_probability(&spec, &assign, prob).unwrap();
755                }
756            }
757        }
758
759        // Front-door: E[Y|do(T=t)] = Σ_m P(m|t) Σ_t' P(y|m,t') P(t')
760        // With P(Y|M) independent of T': E[Y|do(T=1)] = 0.7*0.9 + 0.3*0.1 = 0.66
761        // E[Y|do(T=0)] = 0.3*0.9 + 0.7*0.1 = 0.34
762        // ATE = 0.32
763        let compiled = arena.compile(expr).unwrap();
764        let ate = compiled.evaluate(&arena, &p, &EvalContext::default()).unwrap();
765        assert!((ate - 0.32).abs() < 1e-12, "ate={ate}");
766
767        let simplified = arena.simplify(expr).unwrap();
768        let ate2 = arena
769            .compile(simplified)
770            .unwrap()
771            .evaluate(&arena, &p, &EvalContext::default())
772            .unwrap();
773        assert!((ate - ate2).abs() < 1e-12);
774    }
775
776    #[test]
777    fn discrete_integral_out_matches_sum_out() {
778        let mut arena = CausalExprArena::new();
779        let empty = arena.empty_var_set();
780        let empty_i = arena.empty_intervention_set();
781        let z = v(0);
782        let zset = arena.intern_var_set([z]);
783        let dist = arena.intern(ExprNode::Distribution {
784            variables: zset,
785            conditioned_on: empty,
786            intervention: empty_i,
787            domain: DomainRef::Observational,
788        });
789        let sum = arena.intern(ExprNode::SumOut { variables: zset, expr: dist });
790        let integ = arena.intern(ExprNode::IntegralOut { variables: zset, expr: dist });
791
792        let mut p = EmpiricalTableProvider::new();
793        p.set_domain(z, [f(0.0), f(1.0)]);
794        for (zval, prob) in [(0.0, 0.3), (1.0, 0.7)] {
795            let spec = FactorSpec {
796                variables: &[z],
797                conditioned_on: &[],
798                intervention: &[],
799                domain: DomainRef::Observational,
800            };
801            p.insert_probability(&spec, &Assignment::from_pairs([(z, f(zval))]), prob).unwrap();
802        }
803        let s = arena.compile(sum).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
804        let i =
805            arena.compile(integ).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
806        assert!((s - 1.0).abs() < 1e-12);
807        assert!((i - s).abs() < 1e-12);
808    }
809
810    #[test]
811    fn continuous_gaussian_integral_out_normalizes() {
812        use crate::provider::GaussianDensityProvider;
813        let mut arena = CausalExprArena::new();
814        let empty = arena.empty_var_set();
815        let empty_i = arena.empty_intervention_set();
816        let x = v(0);
817        let xset = arena.intern_var_set([x]);
818        let dist = arena.intern(ExprNode::Distribution {
819            variables: xset,
820            conditioned_on: empty,
821            intervention: empty_i,
822            domain: DomainRef::Observational,
823        });
824        let integ = arena.intern(ExprNode::IntegralOut { variables: xset, expr: dist });
825        let mut p = GaussianDensityProvider::new();
826        p.set_gaussian(x, 0.0, 1.0);
827        let mass =
828            arena.compile(integ).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
829        assert!((mass - 1.0).abs() < 1e-6, "∫ φ = {mass}");
830    }
831
832    #[test]
833    fn nested_integral_out_product_gaussian() {
834        use crate::provider::GaussianDensityProvider;
835        let mut arena = CausalExprArena::new();
836        let empty = arena.empty_var_set();
837        let empty_i = arena.empty_intervention_set();
838        let x = v(0);
839        let y = v(1);
840        let xset = arena.intern_var_set([x]);
841        let yset = arena.intern_var_set([y]);
842        let both = arena.intern_var_set([x, y]);
843        let dist = arena.intern(ExprNode::Distribution {
844            variables: both,
845            conditioned_on: empty,
846            intervention: empty_i,
847            domain: DomainRef::Observational,
848        });
849        let inner = arena.intern(ExprNode::IntegralOut { variables: yset, expr: dist });
850        let outer = arena.intern(ExprNode::IntegralOut { variables: xset, expr: inner });
851        let mut p = GaussianDensityProvider::new();
852        p.set_gaussian(x, 1.0, 0.25);
853        p.set_gaussian(y, -0.5, 4.0);
854        let mass =
855            arena.compile(outer).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
856        assert!((mass - 1.0).abs() < 1e-5, "∬ φ = {mass}");
857    }
858
859    #[test]
860    fn posterior_evaluate_batch() {
861        let mut arena = CausalExprArena::new();
862        let t = v(0);
863        let y = v(1);
864        let z = v(2);
865        let expr = arena.backdoor_ate(t, y, &[z], f(1.0), f(0.0));
866
867        let draw0 = backdoor_provider(t, y, z);
868        // Perturb P(Z) in draw1 so ATE still 0.45 if conditionals unchanged...
869        // Actually change E[Y|T=1,Z=*] so ATE differs.
870        let mut draw1 = EmpiricalTableProvider::new();
871        draw1.set_domain(z, [f(0.0), f(1.0)]);
872        draw1.set_domain(y, [f(0.0), f(1.0)]);
873        draw1.set_domain(t, [f(0.0), f(1.0)]);
874        for (zval, prob) in [(0.0, 0.5), (1.0, 0.5)] {
875            let spec = FactorSpec {
876                variables: &[z],
877                conditioned_on: &[],
878                intervention: &[],
879                domain: DomainRef::Observational,
880            };
881            draw1.insert_probability(&spec, &Assignment::from_pairs([(z, f(zval))]), prob).unwrap();
882        }
883        // E[Y|T=1,*]=1.0, E[Y|T=0,*]=0.0 → ATE = 1.0
884        for tlev in [0.0, 1.0] {
885            let interv = [InterventionAssignment { variable: t, value: f(tlev) }];
886            let py1 = tlev;
887            for zlev in [0.0, 1.0] {
888                for (yval, prob) in [(1.0, py1), (0.0, 1.0 - py1)] {
889                    let spec = FactorSpec {
890                        variables: &[y],
891                        conditioned_on: &[z],
892                        intervention: &interv,
893                        domain: DomainRef::Interventional,
894                    };
895                    draw1
896                        .insert_probability(
897                            &spec,
898                            &Assignment::from_pairs([(y, f(yval)), (z, f(zlev))]),
899                            prob,
900                        )
901                        .unwrap();
902                }
903            }
904        }
905
906        let posterior = PosteriorDrawProvider::from_draws(vec![draw0, draw1]);
907        let compiled = arena.compile(expr).unwrap();
908        let batch = compiled.evaluate_batch(&arena, &posterior).unwrap();
909        assert_eq!(batch.len(), 2);
910        assert!((batch[0] - 0.45).abs() < 1e-12, "draw0={}", batch[0]);
911        assert!((batch[1] - 1.0).abs() < 1e-12, "draw1={}", batch[1]);
912
913        let single0 =
914            compiled.evaluate(&arena, &posterior, &EvalContext { draw: Some(0) }).unwrap();
915        let single1 =
916            compiled.evaluate(&arena, &posterior, &EvalContext { draw: Some(1) }).unwrap();
917        assert!((single0 - batch[0]).abs() < 1e-15);
918        assert!((single1 - batch[1]).abs() < 1e-15);
919    }
920
921    #[test]
922    fn expectation_of_simple_marginal() {
923        let mut arena = CausalExprArena::new();
924        let y = v(0);
925        let yset = arena.intern_var_set([y]);
926        let empty = arena.empty_var_set();
927        let empty_i = arena.empty_intervention_set();
928        let dist = arena.intern(ExprNode::Distribution {
929            variables: yset,
930            conditioned_on: empty,
931            intervention: empty_i,
932            domain: DomainRef::Observational,
933        });
934        let exp = arena.intern(ExprNode::Expectation {
935            function: OutcomeExprId::identity(y),
936            distribution: dist,
937        });
938
939        let mut p = EmpiricalTableProvider::new();
940        p.set_domain(y, [f(0.0), f(2.0)]);
941        let spec = FactorSpec {
942            variables: &[y],
943            conditioned_on: &[],
944            intervention: &[],
945            domain: DomainRef::Observational,
946        };
947        p.insert_probability(&spec, &Assignment::from_pairs([(y, f(0.0))]), 0.25).unwrap();
948        p.insert_probability(&spec, &Assignment::from_pairs([(y, f(2.0))]), 0.75).unwrap();
949
950        let val =
951            arena.compile(exp).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
952        // 0*0.25 + 2*0.75 = 1.5
953        assert!((val - 1.5).abs() < 1e-12);
954    }
955
956    #[test]
957    fn evaluation_is_stable_across_repeated_calls() {
958        // Support memoization and the shared scratch assignment must leave
959        // repeated evaluations bitwise identical (nested SumOut + Expectation
960        // exercise both caches on the second call).
961        let mut arena = CausalExprArena::new();
962        let t = v(0);
963        let y = v(1);
964        let z = v(2);
965        let expr = arena.backdoor_ate(t, y, &[z], f(1.0), f(0.0));
966        let provider = backdoor_provider(t, y, z);
967        let compiled = arena.compile(expr).unwrap();
968        let first = compiled.evaluate(&arena, &provider, &EvalContext::default()).unwrap();
969        let second = compiled.evaluate(&arena, &provider, &EvalContext::default()).unwrap();
970        let third = compiled.evaluate(&arena, &provider, &EvalContext::default()).unwrap();
971        assert_eq!(first.to_bits(), second.to_bits());
972        assert_eq!(first.to_bits(), third.to_bits());
973        assert!((first - 0.45).abs() < 1e-12, "ate={first}");
974    }
975
976    #[test]
977    fn scoped_intervention_binding_restores_between_siblings() {
978        // SumOut_z Product[ P(· | z, do(z:=1)), P(z) ]: the first factor binds
979        // z:=1 for its own lookup only; the sibling P(z) must still see the
980        // row's z. Correct scoping gives Σ_z 2.0 · P(z) = 2.0; a leaked
981        // binding would give 2.0 · P(z=1) per row = 2.8.
982        let mut arena = CausalExprArena::new();
983        let z = v(0);
984        let zset = arena.intern_var_set([z]);
985        let empty = arena.empty_var_set();
986        let empty_i = arena.empty_intervention_set();
987        let do_z1 = arena.intern_intervention_assignments([InterventionAssignment {
988            variable: z,
989            value: f(1.0),
990        }]);
991        let shadowed = arena.intern(ExprNode::Distribution {
992            variables: empty,
993            conditioned_on: zset,
994            intervention: do_z1,
995            domain: DomainRef::Observational,
996        });
997        let z_marginal = arena.intern(ExprNode::Distribution {
998            variables: zset,
999            conditioned_on: empty,
1000            intervention: empty_i,
1001            domain: DomainRef::Observational,
1002        });
1003        let product = {
1004            let list = arena.intern_list([shadowed, z_marginal]);
1005            arena.intern(ExprNode::Product(list))
1006        };
1007        let sum = arena.intern(ExprNode::SumOut { variables: zset, expr: product });
1008
1009        let mut p = EmpiricalTableProvider::new();
1010        p.set_domain(z, [f(0.0), f(1.0)]);
1011        let interv = [InterventionAssignment { variable: z, value: f(1.0) }];
1012        let shadow_spec = FactorSpec {
1013            variables: &[],
1014            conditioned_on: &[z],
1015            intervention: &interv,
1016            domain: DomainRef::Observational,
1017        };
1018        p.insert_probability(&shadow_spec, &Assignment::from_pairs([(z, f(1.0))]), 2.0).unwrap();
1019        let marg_spec = FactorSpec {
1020            variables: &[z],
1021            conditioned_on: &[],
1022            intervention: &[],
1023            domain: DomainRef::Observational,
1024        };
1025        p.insert_probability(&marg_spec, &Assignment::from_pairs([(z, f(0.0))]), 0.3).unwrap();
1026        p.insert_probability(&marg_spec, &Assignment::from_pairs([(z, f(1.0))]), 0.7).unwrap();
1027
1028        let val =
1029            arena.compile(sum).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
1030        assert!((val - 2.0).abs() < 1e-12, "val={val}");
1031    }
1032
1033    #[test]
1034    fn expectation_respects_env_bound_conditioning() {
1035        // E[Y | z] with z pre-bound in the environment: the compile-time
1036        // free-variable set of the distribution slot is filtered against the
1037        // environment, so only Y is enumerated and the bound z selects the
1038        // right conditional column.
1039        let mut arena = CausalExprArena::new();
1040        let y = v(0);
1041        let z = v(1);
1042        let yset = arena.intern_var_set([y]);
1043        let zset = arena.intern_var_set([z]);
1044        let empty_i = arena.empty_intervention_set();
1045        let dist = arena.intern(ExprNode::Distribution {
1046            variables: yset,
1047            conditioned_on: zset,
1048            intervention: empty_i,
1049            domain: DomainRef::Observational,
1050        });
1051        let exp = arena.intern(ExprNode::Expectation {
1052            function: OutcomeExprId::identity(y),
1053            distribution: dist,
1054        });
1055
1056        let mut p = EmpiricalTableProvider::new();
1057        p.set_domain(y, [f(0.0), f(2.0)]);
1058        p.set_domain(z, [f(0.0), f(1.0)]);
1059        let spec = FactorSpec {
1060            variables: &[y],
1061            conditioned_on: &[z],
1062            intervention: &[],
1063            domain: DomainRef::Observational,
1064        };
1065        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)]
1066        {
1067            p.insert_probability(&spec, &Assignment::from_pairs([(y, f(yv)), (z, f(zv))]), prob)
1068                .unwrap();
1069        }
1070        let compiled = arena.compile(exp).unwrap();
1071        let env0 = Assignment::from_pairs([(z, f(0.0))]);
1072        let e0 = compiled.evaluate_with(&arena, &p, &EvalContext::default(), &env0).unwrap();
1073        assert!((e0 - 1.5).abs() < 1e-12, "E[Y|z=0]={e0}");
1074        let env1 = Assignment::from_pairs([(z, f(1.0))]);
1075        let e1 = compiled.evaluate_with(&arena, &p, &EvalContext::default(), &env1).unwrap();
1076        assert!(e1.abs() < 1e-12, "E[Y|z=1]={e1}");
1077        // The caller's environment is never mutated by evaluation.
1078        assert_eq!(env0.entries(), &[(z, f(0.0))]);
1079    }
1080
1081    #[test]
1082    fn ratio_zero_denominator_is_division_by_zero() {
1083        // `EvalOp::Ratio` must reject an exactly-zero denominator rather than
1084        // returning `f64::INFINITY`/NaN. `iv_wald_is_ratio_of_instrument_contrasts`
1085        // (lib.rs) only checks the compiled shape, not this evaluation-time guard.
1086        let mut arena = CausalExprArena::new();
1087        let empty = arena.empty_var_set();
1088        let empty_i = arena.empty_intervention_set();
1089        // Two vacuous (no free variables) factors, distinguished by domain so they
1090        // hash-cons to distinct nodes with independently settable probabilities.
1091        let numerator = arena.intern(ExprNode::Distribution {
1092            variables: empty,
1093            conditioned_on: empty,
1094            intervention: empty_i,
1095            domain: DomainRef::Observational,
1096        });
1097        let denominator = arena.intern(ExprNode::Distribution {
1098            variables: empty,
1099            conditioned_on: empty,
1100            intervention: empty_i,
1101            domain: DomainRef::Interventional,
1102        });
1103        let ratio = arena.intern(ExprNode::Ratio { numerator, denominator });
1104
1105        let mut p = EmpiricalTableProvider::new();
1106        let obs_spec = FactorSpec {
1107            variables: &[],
1108            conditioned_on: &[],
1109            intervention: &[],
1110            domain: DomainRef::Observational,
1111        };
1112        let interv_spec = FactorSpec {
1113            variables: &[],
1114            conditioned_on: &[],
1115            intervention: &[],
1116            domain: DomainRef::Interventional,
1117        };
1118        p.insert_probability(&obs_spec, &Assignment::from_pairs([]), 3.0).unwrap();
1119        p.insert_probability(&interv_spec, &Assignment::from_pairs([]), 0.0).unwrap();
1120
1121        let err = arena
1122            .compile(ratio)
1123            .unwrap()
1124            .evaluate(&arena, &p, &EvalContext::default())
1125            .unwrap_err();
1126        assert_eq!(err, EvalError::DivisionByZero);
1127    }
1128}