1use antecedent_core::{Value, VariableId};
6
7use crate::{CausalExprArena, ContrastOp, DomainRef, ExprId, ExprNode, InterventionAssignment};
8
9pub(crate) fn latex_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}\\mid {cond})")
21 }
22 }
23 DomainRef::Interventional => {
24 if cond.is_empty() {
25 format!("P({vars}\\mid \\mathrm{{do}}({interv}))")
26 } else {
27 format!("P({vars}\\mid {cond},\\mathrm{{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| latex_expr(arena, *e)).collect();
35 parts.join(" \\cdot ")
36 }
37 ExprNode::SumOut { variables, expr } => {
38 format!(
39 "\\sum_{{{}}}\\left[{}\\right]",
40 fmt_vars(arena.var_set(*variables)),
41 latex_expr(arena, *expr)
42 )
43 }
44 ExprNode::IntegralOut { variables, expr } => {
45 format!(
46 "\\int_{{{}}}\\left[{}\\right]",
47 fmt_vars(arena.var_set(*variables)),
48 latex_expr(arena, *expr)
49 )
50 }
51 ExprNode::Ratio { numerator, denominator } => {
52 format!(
53 "\\frac{{{}}}{{{}}}",
54 latex_expr(arena, *numerator),
55 latex_expr(arena, *denominator)
56 )
57 }
58 ExprNode::Expectation { function, distribution } => {
59 format!(
60 "\\mathbb{{E}}\\left[V{} \\mid {}\\right]",
61 function.variable().raw(),
62 latex_expr(arena, *distribution)
63 )
64 }
65 ExprNode::Contrast { left, right, op } => {
66 let op_s = match op {
67 ContrastOp::Difference => "-",
68 };
69 format!(
70 "\\left({}\\right) {} \\left({}\\right)",
71 latex_expr(arena, *left),
72 op_s,
73 latex_expr(arena, *right)
74 )
75 }
76 }
77}
78
79fn fmt_vars(vars: &[VariableId]) -> String {
80 vars.iter().map(|v| format!("V{}", v.raw())).collect::<Vec<_>>().join(",")
81}
82
83fn fmt_assignments(assignments: &[InterventionAssignment]) -> String {
84 assignments
85 .iter()
86 .map(|a| format!("V{}:={}", a.variable.raw(), fmt_value(&a.value)))
87 .collect::<Vec<_>>()
88 .join(",")
89}
90
91fn fmt_value(v: &Value) -> String {
92 match v {
93 Value::Float64(x) => format!("{x}"),
94 Value::Int64(x) => format!("{x}"),
95 Value::Bool(x) => format!("{x}"),
96 Value::Category(x) => format!("c{x}"),
97 Value::Label(x) => x.to_string(),
98 }
99}