Skip to main content

antecedent_expr/
pretty.rs

1//! Pretty-printing for diagnostics (not equality keys).
2//!
3//! SPDX-License-Identifier: MIT OR Apache-2.0
4
5use 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}