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::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    root: usize,
56}
57
58impl CausalExprArena {
59    /// Compile `root` into a topological evaluation plan.
60    ///
61    /// Continuous [`ExprNode::IntegralOut`] compiles successfully; evaluation uses
62    /// [`DistributionProvider::quadrature`] or discrete [`DistributionProvider::support`].
63    pub fn compile(&self, root: ExprId) -> Result<CompiledEvaluator, EvalError> {
64        CompiledEvaluator::compile(self, root)
65    }
66}
67
68impl CompiledEvaluator {
69    /// Compile an expression DAG into slot-addressed ops (post-order).
70    ///
71    /// Continuous [`ExprNode::IntegralOut`] is supported (see [`CausalExprArena::compile`]).
72    pub fn compile(arena: &CausalExprArena, root: ExprId) -> Result<Self, EvalError> {
73        let mut ops = Vec::new();
74        let mut expr_to_slot = HashMap::new();
75        let root_slot = compile_rec(arena, root, &mut ops, &mut expr_to_slot)?;
76        Ok(Self { ops, root: root_slot })
77    }
78
79    /// Evaluate once against a provider.
80    ///
81    /// # Errors
82    ///
83    /// Provider / numeric failures.
84    pub fn evaluate(
85        &self,
86        arena: &CausalExprArena,
87        provider: &dyn DistributionProvider,
88        ctx: &EvalContext,
89    ) -> Result<f64, EvalError> {
90        self.evaluate_with(arena, provider, ctx, &Assignment::new())
91    }
92
93    /// Evaluate with an initial variable binding (e.g. `do(X=x)` and outcome levels).
94    ///
95    /// # Errors
96    ///
97    /// Provider / numeric failures.
98    pub fn evaluate_with(
99        &self,
100        arena: &CausalExprArena,
101        provider: &dyn DistributionProvider,
102        ctx: &EvalContext,
103        env: &Assignment,
104    ) -> Result<f64, EvalError> {
105        self.eval_slot(arena, provider, ctx, env, self.root)
106    }
107
108    /// Evaluate over all posterior draws (`provider.n_draws()`), or a single
109    /// empirical evaluation when `n_draws` is `None`.
110    ///
111    /// # Errors
112    ///
113    /// Provider / numeric failures.
114    pub fn evaluate_batch(
115        &self,
116        arena: &CausalExprArena,
117        provider: &dyn DistributionProvider,
118    ) -> Result<Vec<f64>, EvalError> {
119        match provider.n_draws() {
120            None => Ok(vec![self.evaluate(arena, provider, &EvalContext::default())?]),
121            Some(n) => {
122                let mut out = Vec::with_capacity(n);
123                for draw in 0..n {
124                    let ctx = EvalContext { draw: Some(draw) };
125                    out.push(self.evaluate(arena, provider, &ctx)?);
126                }
127                Ok(out)
128            }
129        }
130    }
131
132    fn eval_slot(
133        &self,
134        arena: &CausalExprArena,
135        provider: &dyn DistributionProvider,
136        ctx: &EvalContext,
137        env: &Assignment,
138        slot: usize,
139    ) -> Result<f64, EvalError> {
140        // Density / scalar under `env`. Expectations and contrasts are scalars;
141        // other ops are densities in the free variables bound by `env`.
142        match &self.ops[slot] {
143            EvalOp::Distribution { variables, conditioned_on, intervention, domain } => {
144                let spec = FactorSpec {
145                    variables: arena.var_set(*variables),
146                    conditioned_on: arena.var_set(*conditioned_on),
147                    intervention: arena.intervention_assignments(*intervention),
148                    domain: *domain,
149                };
150                // Interventions bind targets; merge into lookup assignment.
151                let mut lookup = env.clone();
152                for a in spec.intervention {
153                    lookup.set(a.variable, a.value.clone());
154                }
155                provider.probability(&spec, &lookup, ctx)
156            }
157            EvalOp::Product { children } => {
158                let mut prod = 1.0;
159                for &c in children.iter() {
160                    prod *= self.eval_slot(arena, provider, ctx, env, c)?;
161                }
162                Ok(prod)
163            }
164            EvalOp::SumOut { variables, body } => {
165                self.eval_sum_out(arena, provider, ctx, env, *variables, *body)
166            }
167            EvalOp::IntegralOut { variables, body } => {
168                self.eval_integral_out(arena, provider, ctx, env, *variables, *body)
169            }
170            EvalOp::Ratio { numerator, denominator } => {
171                let num = self.eval_slot(arena, provider, ctx, env, *numerator)?;
172                let den = self.eval_slot(arena, provider, ctx, env, *denominator)?;
173                if den == 0.0 {
174                    return Err(EvalError::DivisionByZero);
175                }
176                Ok(num / den)
177            }
178            EvalOp::Expectation { function, distribution } => {
179                self.eval_expectation(arena, provider, ctx, env, function.variable(), *distribution)
180            }
181            EvalOp::Contrast { left, right, op } => {
182                let l = self.eval_slot(arena, provider, ctx, env, *left)?;
183                let r = self.eval_slot(arena, provider, ctx, env, *right)?;
184                match op {
185                    ContrastOp::Difference => Ok(l - r),
186                }
187            }
188        }
189    }
190
191    fn eval_sum_out(
192        &self,
193        arena: &CausalExprArena,
194        provider: &dyn DistributionProvider,
195        ctx: &EvalContext,
196        env: &Assignment,
197        variables: VarSetId,
198        body: usize,
199    ) -> Result<f64, EvalError> {
200        let vars = arena.var_set(variables);
201        let rows = provider.support(vars, ctx)?;
202        let mut sum = 0.0;
203        for row in rows.iter() {
204            if row.len() != vars.len() {
205                return Err(EvalError::SupportShape { expected: vars.len(), actual: row.len() });
206            }
207            let mut extended = env.clone();
208            for (i, &v) in vars.iter().enumerate() {
209                extended.set(v, row[i].clone());
210            }
211            sum += self.eval_slot(arena, provider, ctx, &extended, body)?;
212        }
213        Ok(sum)
214    }
215
216    fn eval_integral_out(
217        &self,
218        arena: &CausalExprArena,
219        provider: &dyn DistributionProvider,
220        ctx: &EvalContext,
221        env: &Assignment,
222        variables: VarSetId,
223        body: usize,
224    ) -> Result<f64, EvalError> {
225        let vars = arena.var_set(variables);
226        if let Some(nodes) = provider.quadrature(vars, ctx)? {
227            let mut acc = 0.0;
228            for (row, weight) in nodes.iter() {
229                if row.len() != vars.len() {
230                    return Err(EvalError::SupportShape {
231                        expected: vars.len(),
232                        actual: row.len(),
233                    });
234                }
235                let mut extended = env.clone();
236                for (i, &v) in vars.iter().enumerate() {
237                    extended.set(v, row[i].clone());
238                }
239                acc += *weight * self.eval_slot(arena, provider, ctx, &extended, body)?;
240            }
241            return Ok(acc);
242        }
243        // Discrete / counting-measure fallback (IntegralOut ≡ SumOut).
244        let rows = provider.support(vars, ctx).map_err(|e| match e {
245            EvalError::EmptySupport(_) => EvalError::UnsupportedIntegralOut,
246            other => other,
247        })?;
248        let mut sum = 0.0;
249        for row in rows.iter() {
250            if row.len() != vars.len() {
251                return Err(EvalError::SupportShape { expected: vars.len(), actual: row.len() });
252            }
253            let mut extended = env.clone();
254            for (i, &v) in vars.iter().enumerate() {
255                extended.set(v, row[i].clone());
256            }
257            sum += self.eval_slot(arena, provider, ctx, &extended, body)?;
258        }
259        Ok(sum)
260    }
261
262    fn eval_expectation(
263        &self,
264        arena: &CausalExprArena,
265        provider: &dyn DistributionProvider,
266        ctx: &EvalContext,
267        env: &Assignment,
268        outcome_var: VariableId,
269        distribution: usize,
270    ) -> Result<f64, EvalError> {
271        // E[f | D] = Σ_{x ∈ support(free(D))} f(x) · dens(D, x)
272        let free = free_vars_of_slot(self, arena, distribution);
273        let unbound: Vec<VariableId> = free.into_iter().filter(|v| env.get(*v).is_none()).collect();
274        let mut enum_vars = unbound;
275        if !enum_vars.contains(&outcome_var) && env.get(outcome_var).is_none() {
276            enum_vars.push(outcome_var);
277        }
278        enum_vars.sort_by_key(|v| v.raw());
279        enum_vars.dedup();
280
281        if enum_vars.is_empty() {
282            let dens = self.eval_slot(arena, provider, ctx, env, distribution)?;
283            let y = provider.outcome(outcome_var, env, ctx)?;
284            return Ok(y * dens);
285        }
286
287        let rows = provider.support(&enum_vars, ctx)?;
288        let mut acc = 0.0;
289        for row in rows.iter() {
290            if row.len() != enum_vars.len() {
291                return Err(EvalError::SupportShape {
292                    expected: enum_vars.len(),
293                    actual: row.len(),
294                });
295            }
296            let mut extended = env.clone();
297            for (i, &v) in enum_vars.iter().enumerate() {
298                extended.set(v, row[i].clone());
299            }
300            let dens = self.eval_slot(arena, provider, ctx, &extended, distribution)?;
301            let y = provider.outcome(outcome_var, &extended, ctx)?;
302            acc += y * dens;
303        }
304        Ok(acc)
305    }
306}
307
308fn compile_rec(
309    arena: &CausalExprArena,
310    id: ExprId,
311    ops: &mut Vec<EvalOp>,
312    expr_to_slot: &mut HashMap<u32, usize>,
313) -> Result<usize, EvalError> {
314    if let Some(&slot) = expr_to_slot.get(&id.raw()) {
315        return Ok(slot);
316    }
317    let op = match arena.node(id).clone() {
318        ExprNode::Distribution { variables, conditioned_on, intervention, domain } => {
319            EvalOp::Distribution { variables, conditioned_on, intervention, domain }
320        }
321        ExprNode::Product(list) => {
322            let mut children = Vec::new();
323            for &c in arena.list(list) {
324                children.push(compile_rec(arena, c, ops, expr_to_slot)?);
325            }
326            EvalOp::Product { children: Arc::from(children) }
327        }
328        ExprNode::SumOut { variables, expr } => {
329            let body = compile_rec(arena, expr, ops, expr_to_slot)?;
330            EvalOp::SumOut { variables, body }
331        }
332        ExprNode::IntegralOut { variables, expr } => {
333            let body = compile_rec(arena, expr, ops, expr_to_slot)?;
334            EvalOp::IntegralOut { variables, body }
335        }
336        ExprNode::Ratio { numerator, denominator } => {
337            let n = compile_rec(arena, numerator, ops, expr_to_slot)?;
338            let d = compile_rec(arena, denominator, ops, expr_to_slot)?;
339            EvalOp::Ratio { numerator: n, denominator: d }
340        }
341        ExprNode::Expectation { function, distribution } => {
342            let dist = compile_rec(arena, distribution, ops, expr_to_slot)?;
343            EvalOp::Expectation { function, distribution: dist }
344        }
345        ExprNode::Contrast { left, right, op } => {
346            let l = compile_rec(arena, left, ops, expr_to_slot)?;
347            let r = compile_rec(arena, right, ops, expr_to_slot)?;
348            EvalOp::Contrast { left: l, right: r, op }
349        }
350    };
351    let slot = ops.len();
352    ops.push(op);
353    expr_to_slot.insert(id.raw(), slot);
354    Ok(slot)
355}
356
357fn free_vars_of_slot(
358    compiled: &CompiledEvaluator,
359    arena: &CausalExprArena,
360    slot: usize,
361) -> Vec<VariableId> {
362    let mut out = Vec::new();
363    free_vars_rec(compiled, arena, slot, &mut out);
364    out.sort_by_key(|v| v.raw());
365    out.dedup();
366    out
367}
368
369fn free_vars_rec(
370    compiled: &CompiledEvaluator,
371    arena: &CausalExprArena,
372    slot: usize,
373    out: &mut Vec<VariableId>,
374) {
375    match &compiled.ops[slot] {
376        EvalOp::Distribution { variables, conditioned_on, intervention, .. } => {
377            out.extend_from_slice(arena.var_set(*variables));
378            let bound: Vec<VariableId> =
379                arena.intervention_assignments(*intervention).iter().map(|a| a.variable).collect();
380            for &v in arena.var_set(*conditioned_on) {
381                if !bound.iter().any(|b| *b == v) {
382                    out.push(v);
383                }
384            }
385        }
386        EvalOp::Product { children } => {
387            for &c in children.iter() {
388                free_vars_rec(compiled, arena, c, out);
389            }
390        }
391        EvalOp::SumOut { variables, body } | EvalOp::IntegralOut { variables, body } => {
392            let mut inner = Vec::new();
393            free_vars_rec(compiled, arena, *body, &mut inner);
394            let bound = arena.var_set(*variables);
395            for v in inner {
396                if !bound.iter().any(|b| *b == v) {
397                    out.push(v);
398                }
399            }
400        }
401        EvalOp::Ratio { numerator, denominator } => {
402            free_vars_rec(compiled, arena, *numerator, out);
403            free_vars_rec(compiled, arena, *denominator, out);
404        }
405        EvalOp::Expectation { function, distribution } => {
406            free_vars_rec(compiled, arena, *distribution, out);
407            out.push(function.variable());
408        }
409        EvalOp::Contrast { left, right, .. } => {
410            free_vars_rec(compiled, arena, *left, out);
411            free_vars_rec(compiled, arena, *right, out);
412        }
413    }
414}
415
416#[cfg(test)]
417mod tests {
418    use super::*;
419    use crate::provider::{EmpiricalTableProvider, PosteriorDrawProvider};
420    use crate::{InterventionAssignment, OutcomeExprId};
421    use antecedent_core::Value;
422
423    fn v(id: u32) -> VariableId {
424        VariableId::from_raw(id)
425    }
426
427    fn f(x: f64) -> Value {
428        Value::f64(x)
429    }
430
431    /// Binary confounder Z, binary Y; backdoor ATE = 0.45.
432    fn backdoor_provider(t: VariableId, y: VariableId, z: VariableId) -> EmpiricalTableProvider {
433        let mut p = EmpiricalTableProvider::new();
434        p.set_domain(z, [f(0.0), f(1.0)]);
435        p.set_domain(y, [f(0.0), f(1.0)]);
436        p.set_domain(t, [f(0.0), f(1.0)]);
437
438        // P(Z)
439        for (zval, prob) in [(0.0, 0.5), (1.0, 0.5)] {
440            let spec = FactorSpec {
441                variables: &[z],
442                conditioned_on: &[],
443                intervention: &[],
444                domain: DomainRef::Observational,
445            };
446            let assign = Assignment::from_pairs([(z, f(zval))]);
447            p.insert_probability(&spec, &assign, prob).unwrap();
448        }
449
450        // P(Y | Z, do(T=t)) = P(Y | T=t, Z) under backdoor.
451        // 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
452        let ey = |tlev: f64, zlev: f64| -> f64 {
453            match (tlev.to_bits(), zlev.to_bits()) {
454                (t, z) if t == 1.0f64.to_bits() && z == 0.0f64.to_bits() => 0.8,
455                (t, z) if t == 1.0f64.to_bits() && z == 1.0f64.to_bits() => 0.6,
456                (t, z) if t == 0.0f64.to_bits() && z == 0.0f64.to_bits() => 0.3,
457                (t, z) if t == 0.0f64.to_bits() && z == 1.0f64.to_bits() => 0.2,
458                _ => panic!("bad levels"),
459            }
460        };
461        for tlev in [0.0, 1.0] {
462            let interv = [InterventionAssignment { variable: t, value: f(tlev) }];
463            for zlev in [0.0, 1.0] {
464                let p_y1 = ey(tlev, zlev);
465                for (yval, prob) in [(1.0, p_y1), (0.0, 1.0 - p_y1)] {
466                    let spec = FactorSpec {
467                        variables: &[y],
468                        conditioned_on: &[z],
469                        intervention: &interv,
470                        domain: DomainRef::Interventional,
471                    };
472                    let assign = Assignment::from_pairs([(y, f(yval)), (z, f(zlev))]);
473                    p.insert_probability(&spec, &assign, prob).unwrap();
474                }
475            }
476        }
477        p
478    }
479
480    #[test]
481    fn backdoor_ate_matches_closed_form() {
482        let mut arena = CausalExprArena::new();
483        let t = v(0);
484        let y = v(1);
485        let z = v(2);
486        let expr = arena.backdoor_ate(t, y, &[z], f(1.0), f(0.0));
487        let provider = backdoor_provider(t, y, z);
488        let compiled = arena.compile(expr).unwrap();
489        let ate = compiled.evaluate(&arena, &provider, &EvalContext::default()).unwrap();
490        assert!((ate - 0.45).abs() < 1e-12, "ate={ate}");
491    }
492
493    #[test]
494    fn simplify_preserves_backdoor_evaluation() {
495        let mut arena = CausalExprArena::new();
496        let t = v(0);
497        let y = v(1);
498        let z = v(2);
499        let expr = arena.backdoor_ate(t, y, &[z], f(1.0), f(0.0));
500        let provider = backdoor_provider(t, y, z);
501        let before = arena
502            .compile(expr)
503            .unwrap()
504            .evaluate(&arena, &provider, &EvalContext::default())
505            .unwrap();
506        let simplified = arena.simplify(expr);
507        let after = arena
508            .compile(simplified)
509            .unwrap()
510            .evaluate(&arena, &provider, &EvalContext::default())
511            .unwrap();
512        assert!((before - after).abs() < 1e-12, "before={before} after={after}");
513        assert!((after - 0.45).abs() < 1e-12);
514    }
515
516    /// Empty adjustment (second Z set): simplify must preserve numeric eval.
517    #[test]
518    fn simplify_preserves_backdoor_empty_evaluation() {
519        fn assert_simplify_preserves(
520            arena: &mut CausalExprArena,
521            expr: ExprId,
522            provider: &EmpiricalTableProvider,
523            expected: f64,
524            label: &str,
525        ) {
526            let before = arena
527                .compile(expr)
528                .unwrap()
529                .evaluate(arena, provider, &EvalContext::default())
530                .unwrap();
531            let simplified = arena.simplify(expr);
532            let after = arena
533                .compile(simplified)
534                .unwrap()
535                .evaluate(arena, provider, &EvalContext::default())
536                .unwrap();
537            assert!((before - after).abs() < 1e-12, "{label}: before={before} after={after}");
538            assert!((after - expected).abs() < 1e-12, "{label}: after={after}");
539        }
540
541        // Backdoor with empty Z: E[Y|do(1)]=0.7, E[Y|do(0)]=0.2 → ATE = 0.5.
542        // Exercises simplify.empty_sum_out / singleton product on the adjustment set.
543        let mut arena = CausalExprArena::new();
544        let t = v(0);
545        let y = v(1);
546        let expr = arena.backdoor_ate(t, y, &[], f(1.0), f(0.0));
547        let mut p = EmpiricalTableProvider::new();
548        p.set_domain(y, [f(0.0), f(1.0)]);
549        p.set_domain(t, [f(0.0), f(1.0)]);
550        // Vacuous P(∅) factor from empty adjustment marginal.
551        let empty_spec = FactorSpec {
552            variables: &[],
553            conditioned_on: &[],
554            intervention: &[],
555            domain: DomainRef::Observational,
556        };
557        p.insert_probability(&empty_spec, &Assignment::from_pairs([]), 1.0).unwrap();
558        for tlev in [0.0, 1.0] {
559            let ey = if (tlev - 1.0_f64).abs() < f64::EPSILON { 0.7 } else { 0.2 };
560            let interv = [InterventionAssignment { variable: t, value: f(tlev) }];
561            for (yval, prob) in [(1.0, ey), (0.0, 1.0 - ey)] {
562                let spec = FactorSpec {
563                    variables: &[y],
564                    conditioned_on: &[],
565                    intervention: &interv,
566                    domain: DomainRef::Interventional,
567                };
568                p.insert_probability(&spec, &Assignment::from_pairs([(y, f(yval))]), prob).unwrap();
569            }
570        }
571        assert_simplify_preserves(&mut arena, expr, &p, 0.5, "backdoor_empty_z");
572    }
573
574    /// Frontdoor: simplify must preserve numeric eval.
575    #[test]
576    fn simplify_preserves_frontdoor_evaluation() {
577        fn assert_simplify_preserves(
578            arena: &mut CausalExprArena,
579            expr: ExprId,
580            provider: &EmpiricalTableProvider,
581            expected: f64,
582            label: &str,
583        ) {
584            let before = arena
585                .compile(expr)
586                .unwrap()
587                .evaluate(arena, provider, &EvalContext::default())
588                .unwrap();
589            let simplified = arena.simplify(expr);
590            let after = arena
591                .compile(simplified)
592                .unwrap()
593                .evaluate(arena, provider, &EvalContext::default())
594                .unwrap();
595            assert!((before - after).abs() < 1e-12, "{label}: before={before} after={after}");
596            assert!((after - expected).abs() < 1e-12, "{label}: after={after}");
597        }
598
599        // Frontdoor (same tables as shallow_frontdoor_evaluates): ATE = 0.32.
600        let mut arena = CausalExprArena::new();
601        let t = v(0);
602        let y = v(1);
603        let m = v(2);
604        let expr = arena.frontdoor_ate(t, y, &[m], f(1.0), f(0.0));
605        let mut p = EmpiricalTableProvider::new();
606        p.set_domain(t, [f(0.0), f(1.0)]);
607        p.set_domain(y, [f(0.0), f(1.0)]);
608        p.set_domain(m, [f(0.0), f(1.0)]);
609        for (tval, prob) in [(0.0, 0.5), (1.0, 0.5)] {
610            let spec = FactorSpec {
611                variables: &[t],
612                conditioned_on: &[],
613                intervention: &[],
614                domain: DomainRef::Observational,
615            };
616            p.insert_probability(&spec, &Assignment::from_pairs([(t, f(tval))]), prob).unwrap();
617        }
618        for tlev in [0.0, 1.0] {
619            let pm1 = if (tlev - 1.0_f64).abs() < f64::EPSILON { 0.7 } else { 0.3 };
620            let interv = [InterventionAssignment { variable: t, value: f(tlev) }];
621            for (mval, prob) in [(1.0, pm1), (0.0, 1.0 - pm1)] {
622                let spec = FactorSpec {
623                    variables: &[m],
624                    conditioned_on: &[t],
625                    intervention: &interv,
626                    domain: DomainRef::Observational,
627                };
628                p.insert_probability(
629                    &spec,
630                    &Assignment::from_pairs([(m, f(mval)), (t, f(tlev))]),
631                    prob,
632                )
633                .unwrap();
634            }
635        }
636        for tlev in [0.0, 1.0] {
637            for mlev in [0.0, 1.0] {
638                let py1 = if (mlev - 1.0_f64).abs() < f64::EPSILON { 0.9 } else { 0.1 };
639                for (yval, prob) in [(1.0, py1), (0.0, 1.0 - py1)] {
640                    let spec = FactorSpec {
641                        variables: &[y],
642                        conditioned_on: &[t, m],
643                        intervention: &[],
644                        domain: DomainRef::Observational,
645                    };
646                    let assign = Assignment::from_pairs([(y, f(yval)), (m, f(mlev)), (t, f(tlev))]);
647                    p.insert_probability(&spec, &assign, prob).unwrap();
648                }
649            }
650        }
651        assert_simplify_preserves(&mut arena, expr, &p, 0.32, "frontdoor");
652    }
653
654    #[test]
655    fn shallow_frontdoor_evaluates() {
656        // Minimal front-door: T→M→Y with no hidden confounding encoded in tables.
657        // P(M|T=t); P(Y|M,T'); P(T').
658        let mut arena = CausalExprArena::new();
659        let t = v(0);
660        let y = v(1);
661        let m = v(2);
662        let expr = arena.frontdoor_ate(t, y, &[m], f(1.0), f(0.0));
663
664        let mut p = EmpiricalTableProvider::new();
665        p.set_domain(t, [f(0.0), f(1.0)]);
666        p.set_domain(y, [f(0.0), f(1.0)]);
667        p.set_domain(m, [f(0.0), f(1.0)]);
668
669        // P(T')
670        for (tval, prob) in [(0.0, 0.5), (1.0, 0.5)] {
671            let spec = FactorSpec {
672                variables: &[t],
673                conditioned_on: &[],
674                intervention: &[],
675                domain: DomainRef::Observational,
676            };
677            p.insert_probability(&spec, &Assignment::from_pairs([(t, f(tval))]), prob).unwrap();
678        }
679
680        // P(M | T=t): P(M=1|T=1)=0.7, P(M=1|T=0)=0.3 (FD condition 2).
681        for tlev in [0.0, 1.0] {
682            let pm1 = if (tlev - 1.0_f64).abs() < f64::EPSILON { 0.7 } else { 0.3 };
683            let interv = [InterventionAssignment { variable: t, value: f(tlev) }];
684            for (mval, prob) in [(1.0, pm1), (0.0, 1.0 - pm1)] {
685                let spec = FactorSpec {
686                    variables: &[m],
687                    conditioned_on: &[t],
688                    intervention: &interv,
689                    domain: DomainRef::Observational,
690                };
691                p.insert_probability(
692                    &spec,
693                    &Assignment::from_pairs([(m, f(mval)), (t, f(tlev))]),
694                    prob,
695                )
696                .unwrap();
697            }
698        }
699
700        // P(Y | M, T'): E[Y|M=1,*]=0.9, E[Y|M=0,*]=0.1 (T' irrelevant)
701        // Arena sorts m_and_t as [t, m] when t.raw() < m.raw().
702        for tlev in [0.0, 1.0] {
703            for mlev in [0.0, 1.0] {
704                let py1 = if (mlev - 1.0_f64).abs() < f64::EPSILON { 0.9 } else { 0.1 };
705                for (yval, prob) in [(1.0, py1), (0.0, 1.0 - py1)] {
706                    let spec = FactorSpec {
707                        variables: &[y],
708                        conditioned_on: &[t, m],
709                        intervention: &[],
710                        domain: DomainRef::Observational,
711                    };
712                    let assign = Assignment::from_pairs([(y, f(yval)), (m, f(mlev)), (t, f(tlev))]);
713                    p.insert_probability(&spec, &assign, prob).unwrap();
714                }
715            }
716        }
717
718        // Front-door: E[Y|do(T=t)] = Σ_m P(m|t) Σ_t' P(y|m,t') P(t')
719        // With P(Y|M) independent of T': E[Y|do(T=1)] = 0.7*0.9 + 0.3*0.1 = 0.66
720        // E[Y|do(T=0)] = 0.3*0.9 + 0.7*0.1 = 0.34
721        // ATE = 0.32
722        let compiled = arena.compile(expr).unwrap();
723        let ate = compiled.evaluate(&arena, &p, &EvalContext::default()).unwrap();
724        assert!((ate - 0.32).abs() < 1e-12, "ate={ate}");
725
726        let simplified = arena.simplify(expr);
727        let ate2 = arena
728            .compile(simplified)
729            .unwrap()
730            .evaluate(&arena, &p, &EvalContext::default())
731            .unwrap();
732        assert!((ate - ate2).abs() < 1e-12);
733    }
734
735    #[test]
736    fn discrete_integral_out_matches_sum_out() {
737        let mut arena = CausalExprArena::new();
738        let empty = arena.empty_var_set();
739        let empty_i = arena.empty_intervention_set();
740        let z = v(0);
741        let zset = arena.intern_var_set([z]);
742        let dist = arena.intern(ExprNode::Distribution {
743            variables: zset,
744            conditioned_on: empty,
745            intervention: empty_i,
746            domain: DomainRef::Observational,
747        });
748        let sum = arena.intern(ExprNode::SumOut { variables: zset, expr: dist });
749        let integ = arena.intern(ExprNode::IntegralOut { variables: zset, expr: dist });
750
751        let mut p = EmpiricalTableProvider::new();
752        p.set_domain(z, [f(0.0), f(1.0)]);
753        for (zval, prob) in [(0.0, 0.3), (1.0, 0.7)] {
754            let spec = FactorSpec {
755                variables: &[z],
756                conditioned_on: &[],
757                intervention: &[],
758                domain: DomainRef::Observational,
759            };
760            p.insert_probability(&spec, &Assignment::from_pairs([(z, f(zval))]), prob).unwrap();
761        }
762        let s = arena.compile(sum).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
763        let i =
764            arena.compile(integ).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
765        assert!((s - 1.0).abs() < 1e-12);
766        assert!((i - s).abs() < 1e-12);
767    }
768
769    #[test]
770    fn continuous_gaussian_integral_out_normalizes() {
771        use crate::provider::GaussianDensityProvider;
772        let mut arena = CausalExprArena::new();
773        let empty = arena.empty_var_set();
774        let empty_i = arena.empty_intervention_set();
775        let x = v(0);
776        let xset = arena.intern_var_set([x]);
777        let dist = arena.intern(ExprNode::Distribution {
778            variables: xset,
779            conditioned_on: empty,
780            intervention: empty_i,
781            domain: DomainRef::Observational,
782        });
783        let integ = arena.intern(ExprNode::IntegralOut { variables: xset, expr: dist });
784        let mut p = GaussianDensityProvider::new();
785        p.set_gaussian(x, 0.0, 1.0);
786        let mass =
787            arena.compile(integ).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
788        assert!((mass - 1.0).abs() < 1e-6, "∫ φ = {mass}");
789    }
790
791    #[test]
792    fn nested_integral_out_product_gaussian() {
793        use crate::provider::GaussianDensityProvider;
794        let mut arena = CausalExprArena::new();
795        let empty = arena.empty_var_set();
796        let empty_i = arena.empty_intervention_set();
797        let x = v(0);
798        let y = v(1);
799        let xset = arena.intern_var_set([x]);
800        let yset = arena.intern_var_set([y]);
801        let both = arena.intern_var_set([x, y]);
802        let dist = arena.intern(ExprNode::Distribution {
803            variables: both,
804            conditioned_on: empty,
805            intervention: empty_i,
806            domain: DomainRef::Observational,
807        });
808        let inner = arena.intern(ExprNode::IntegralOut { variables: yset, expr: dist });
809        let outer = arena.intern(ExprNode::IntegralOut { variables: xset, expr: inner });
810        let mut p = GaussianDensityProvider::new();
811        p.set_gaussian(x, 1.0, 0.25);
812        p.set_gaussian(y, -0.5, 4.0);
813        let mass =
814            arena.compile(outer).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
815        assert!((mass - 1.0).abs() < 1e-5, "∬ φ = {mass}");
816    }
817
818    #[test]
819    fn posterior_evaluate_batch() {
820        let mut arena = CausalExprArena::new();
821        let t = v(0);
822        let y = v(1);
823        let z = v(2);
824        let expr = arena.backdoor_ate(t, y, &[z], f(1.0), f(0.0));
825
826        let draw0 = backdoor_provider(t, y, z);
827        // Perturb P(Z) in draw1 so ATE still 0.45 if conditionals unchanged...
828        // Actually change E[Y|T=1,Z=*] so ATE differs.
829        let mut draw1 = EmpiricalTableProvider::new();
830        draw1.set_domain(z, [f(0.0), f(1.0)]);
831        draw1.set_domain(y, [f(0.0), f(1.0)]);
832        draw1.set_domain(t, [f(0.0), f(1.0)]);
833        for (zval, prob) in [(0.0, 0.5), (1.0, 0.5)] {
834            let spec = FactorSpec {
835                variables: &[z],
836                conditioned_on: &[],
837                intervention: &[],
838                domain: DomainRef::Observational,
839            };
840            draw1.insert_probability(&spec, &Assignment::from_pairs([(z, f(zval))]), prob).unwrap();
841        }
842        // E[Y|T=1,*]=1.0, E[Y|T=0,*]=0.0 → ATE = 1.0
843        for tlev in [0.0, 1.0] {
844            let interv = [InterventionAssignment { variable: t, value: f(tlev) }];
845            let py1 = tlev;
846            for zlev in [0.0, 1.0] {
847                for (yval, prob) in [(1.0, py1), (0.0, 1.0 - py1)] {
848                    let spec = FactorSpec {
849                        variables: &[y],
850                        conditioned_on: &[z],
851                        intervention: &interv,
852                        domain: DomainRef::Interventional,
853                    };
854                    draw1
855                        .insert_probability(
856                            &spec,
857                            &Assignment::from_pairs([(y, f(yval)), (z, f(zlev))]),
858                            prob,
859                        )
860                        .unwrap();
861                }
862            }
863        }
864
865        let posterior = PosteriorDrawProvider::from_draws(vec![draw0, draw1]);
866        let compiled = arena.compile(expr).unwrap();
867        let batch = compiled.evaluate_batch(&arena, &posterior).unwrap();
868        assert_eq!(batch.len(), 2);
869        assert!((batch[0] - 0.45).abs() < 1e-12, "draw0={}", batch[0]);
870        assert!((batch[1] - 1.0).abs() < 1e-12, "draw1={}", batch[1]);
871
872        let single0 =
873            compiled.evaluate(&arena, &posterior, &EvalContext { draw: Some(0) }).unwrap();
874        let single1 =
875            compiled.evaluate(&arena, &posterior, &EvalContext { draw: Some(1) }).unwrap();
876        assert!((single0 - batch[0]).abs() < 1e-15);
877        assert!((single1 - batch[1]).abs() < 1e-15);
878    }
879
880    #[test]
881    fn expectation_of_simple_marginal() {
882        let mut arena = CausalExprArena::new();
883        let y = v(0);
884        let yset = arena.intern_var_set([y]);
885        let empty = arena.empty_var_set();
886        let empty_i = arena.empty_intervention_set();
887        let dist = arena.intern(ExprNode::Distribution {
888            variables: yset,
889            conditioned_on: empty,
890            intervention: empty_i,
891            domain: DomainRef::Observational,
892        });
893        let exp = arena.intern(ExprNode::Expectation {
894            function: OutcomeExprId::identity(y),
895            distribution: dist,
896        });
897
898        let mut p = EmpiricalTableProvider::new();
899        p.set_domain(y, [f(0.0), f(2.0)]);
900        let spec = FactorSpec {
901            variables: &[y],
902            conditioned_on: &[],
903            intervention: &[],
904            domain: DomainRef::Observational,
905        };
906        p.insert_probability(&spec, &Assignment::from_pairs([(y, f(0.0))]), 0.25).unwrap();
907        p.insert_probability(&spec, &Assignment::from_pairs([(y, f(2.0))]), 0.75).unwrap();
908
909        let val =
910            arena.compile(exp).unwrap().evaluate(&arena, &p, &EvalContext::default()).unwrap();
911        // 0*0.25 + 2*0.75 = 1.5
912        assert!((val - 1.5).abs() < 1e-12);
913    }
914}