Skip to main content

antecedent_expr/
latex.rs

1//! LaTeX rendering 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 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}