1use 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#[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#[derive(Clone, Debug)]
53pub struct CompiledEvaluator {
54 ops: Vec<EvalOp>,
55 free_vars: Vec<Arc<[VariableId]>>,
59 root: usize,
60}
61
62impl CausalExprArena {
63 pub fn compile(&self, root: ExprId) -> Result<CompiledEvaluator, EvalError> {
68 CompiledEvaluator::compile(self, root)
69 }
70}
71
72impl CompiledEvaluator {
73 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 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 pub fn evaluate_with(
104 &self,
105 arena: &CausalExprArena,
106 provider: &dyn DistributionProvider,
107 ctx: &EvalContext,
108 env: &Assignment,
109 ) -> Result<f64, EvalError> {
110 let mut scratch = env.clone();
115 self.eval_slot(arena, provider, ctx, &mut scratch, self.root)
116 }
117
118 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 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 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 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 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
332fn 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
409fn 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 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 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 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 #[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 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 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 #[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 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 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 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 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 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 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 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 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 assert!((val - 1.5).abs() < 1e-12);
954 }
955
956 #[test]
957 fn evaluation_is_stable_across_repeated_calls() {
958 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 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 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 assert_eq!(env0.entries(), &[(z, f(0.0))]);
1079 }
1080
1081 #[test]
1082 fn ratio_zero_denominator_is_division_by_zero() {
1083 let mut arena = CausalExprArena::new();
1087 let empty = arena.empty_var_set();
1088 let empty_i = arena.empty_intervention_set();
1089 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}