1use 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#[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 root: usize,
56}
57
58impl CausalExprArena {
59 pub fn compile(&self, root: ExprId) -> Result<CompiledEvaluator, EvalError> {
64 CompiledEvaluator::compile(self, root)
65 }
66}
67
68impl CompiledEvaluator {
69 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 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 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 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 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 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 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 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 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 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 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).unwrap();
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 #[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).unwrap();
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 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 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 #[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).unwrap();
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 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 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 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 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 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 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).unwrap();
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 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 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 assert!((val - 1.5).abs() < 1e-12);
913 }
914
915 #[test]
916 fn ratio_zero_denominator_is_division_by_zero() {
917 let mut arena = CausalExprArena::new();
921 let empty = arena.empty_var_set();
922 let empty_i = arena.empty_intervention_set();
923 let numerator = arena.intern(ExprNode::Distribution {
926 variables: empty,
927 conditioned_on: empty,
928 intervention: empty_i,
929 domain: DomainRef::Observational,
930 });
931 let denominator = arena.intern(ExprNode::Distribution {
932 variables: empty,
933 conditioned_on: empty,
934 intervention: empty_i,
935 domain: DomainRef::Interventional,
936 });
937 let ratio = arena.intern(ExprNode::Ratio { numerator, denominator });
938
939 let mut p = EmpiricalTableProvider::new();
940 let obs_spec = FactorSpec {
941 variables: &[],
942 conditioned_on: &[],
943 intervention: &[],
944 domain: DomainRef::Observational,
945 };
946 let interv_spec = FactorSpec {
947 variables: &[],
948 conditioned_on: &[],
949 intervention: &[],
950 domain: DomainRef::Interventional,
951 };
952 p.insert_probability(&obs_spec, &Assignment::from_pairs([]), 3.0).unwrap();
953 p.insert_probability(&interv_spec, &Assignment::from_pairs([]), 0.0).unwrap();
954
955 let err = arena
956 .compile(ratio)
957 .unwrap()
958 .evaluate(&arena, &p, &EvalContext::default())
959 .unwrap_err();
960 assert_eq!(err, EvalError::DivisionByZero);
961 }
962}