antecedent_expr/
pretty.rs1use antecedent_core::{Value, VariableId};
6
7use crate::{CausalExprArena, ContrastOp, DomainRef, ExprId, ExprNode, InterventionAssignment};
8
9pub(crate) fn pretty_expr(arena: &CausalExprArena, id: ExprId) -> String {
10 match arena.node(id) {
11 ExprNode::Distribution { variables, conditioned_on, intervention, domain } => {
12 let vars = fmt_vars(arena.var_set(*variables));
13 let cond = fmt_vars(arena.var_set(*conditioned_on));
14 let interv = fmt_assignments(arena.intervention_assignments(*intervention));
15 match domain {
16 DomainRef::Observational => {
17 if cond.is_empty() {
18 format!("P({vars})")
19 } else {
20 format!("P({vars}|{cond})")
21 }
22 }
23 DomainRef::Interventional => {
24 if cond.is_empty() {
25 format!("P({vars}|do({interv}))")
26 } else {
27 format!("P({vars}|{cond},do({interv}))")
28 }
29 }
30 }
31 }
32 ExprNode::Product(list) => {
33 let parts: Vec<String> =
34 arena.lists[list.0 as usize].iter().map(|e| pretty_expr(arena, *e)).collect();
35 parts.join(" * ")
36 }
37 ExprNode::SumOut { variables, expr } => {
38 format!("Σ_{{{}}}[{}]", fmt_vars(arena.var_set(*variables)), pretty_expr(arena, *expr))
39 }
40 ExprNode::IntegralOut { variables, expr } => {
41 format!("∫_{{{}}}[{}]", fmt_vars(arena.var_set(*variables)), pretty_expr(arena, *expr))
42 }
43 ExprNode::Ratio { numerator, denominator } => {
44 format!("({})/({})", pretty_expr(arena, *numerator), pretty_expr(arena, *denominator))
45 }
46 ExprNode::Expectation { function, distribution } => {
47 format!("E[V{} | {}]", function.variable().raw(), pretty_expr(arena, *distribution))
48 }
49 ExprNode::Contrast { left, right, op } => {
50 let op_s = match op {
51 ContrastOp::Difference => "−",
52 };
53 format!("({}) {} ({})", pretty_expr(arena, *left), op_s, pretty_expr(arena, *right))
54 }
55 }
56}
57
58fn fmt_vars(vars: &[VariableId]) -> String {
59 vars.iter().map(|v| format!("V{}", v.raw())).collect::<Vec<_>>().join(",")
60}
61
62fn fmt_assignments(assignments: &[InterventionAssignment]) -> String {
63 assignments
64 .iter()
65 .map(|a| format!("V{}:={}", a.variable.raw(), fmt_value(&a.value)))
66 .collect::<Vec<_>>()
67 .join(",")
68}
69
70fn fmt_value(v: &Value) -> String {
71 match v {
72 Value::Float64(x) => format!("{x}"),
73 Value::Int64(x) => format!("{x}"),
74 Value::Bool(x) => format!("{x}"),
75 Value::Category(x) => format!("c{x}"),
76 Value::Label(x) => x.to_string(),
77 }
78}