Skip to main content

xlog_prob/kc/
ddnnf.rs

1//! Decision-DNNF parser and CPU reference evaluator.
2
3use std::collections::{BTreeMap, BTreeSet, HashMap, HashSet};
4
5use xlog_core::{Result, XlogError};
6
7use crate::logsumexp::{
8    circuit_logsumexp, validate_circuit_log_weight_pair, validate_circuit_value,
9};
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq)]
12pub enum DdnnfNodeKind {
13    Or,
14    And,
15    True,
16    False,
17}
18
19#[derive(Debug, Clone)]
20pub struct DdnnfNode {
21    pub kind: DdnnfNodeKind,
22}
23
24#[derive(Debug, Clone)]
25pub struct DdnnfEdge {
26    pub from: u32,
27    pub to: u32,
28    pub lits: Vec<i32>,
29}
30
31#[derive(Debug, Clone)]
32pub struct DecisionDnnf {
33    root: u32,
34    nodes: BTreeMap<u32, DdnnfNode>,
35    edges: Vec<DdnnfEdge>,
36    outgoing: BTreeMap<u32, Vec<usize>>,
37    max_var: u32,
38}
39
40impl DecisionDnnf {
41    pub fn root(&self) -> u32 {
42        self.root
43    }
44
45    pub fn max_var(&self) -> u32 {
46        self.max_var
47    }
48
49    pub fn node_kind(&self, node_id: u32) -> Option<DdnnfNodeKind> {
50        self.nodes.get(&node_id).map(|n| n.kind)
51    }
52
53    pub fn outgoing_edge_indices(&self, node_id: u32) -> Option<&[usize]> {
54        self.outgoing.get(&node_id).map(|v| v.as_slice())
55    }
56
57    pub fn edge(&self, edge_idx: usize) -> Option<&DdnnfEdge> {
58        self.edges.get(edge_idx)
59    }
60
61    pub fn parse_str(input: &str) -> Result<Self> {
62        let mut nodes: BTreeMap<u32, DdnnfNode> = BTreeMap::new();
63        let mut edges: Vec<DdnnfEdge> = Vec::new();
64        let mut targets: HashSet<u32> = HashSet::new();
65        let mut max_var: u32 = 0;
66
67        for (line_no, raw_line) in input.lines().enumerate() {
68            let line = raw_line.trim();
69            if line.is_empty() {
70                continue;
71            }
72
73            let mut tokens: Vec<&str> = line.split_whitespace().collect();
74            if tokens.is_empty() {
75                continue;
76            }
77
78            if tokens.last() != Some(&"0") {
79                return Err(XlogError::Compilation(format!(
80                    "Decision-DNNF parse error at line {}: missing 0 terminator",
81                    line_no + 1
82                )));
83            }
84            tokens.pop();
85            if tokens.is_empty() {
86                return Err(XlogError::Compilation(format!(
87                    "Decision-DNNF parse error at line {}: empty record before terminator",
88                    line_no + 1
89                )));
90            }
91
92            match tokens[0] {
93                "o" | "a" | "t" | "f" => {
94                    if tokens.len() < 2 {
95                        return Err(XlogError::Compilation(format!(
96                            "Decision-DNNF parse error at line {}: node record missing id",
97                            line_no + 1
98                        )));
99                    }
100                    let id: u32 = tokens[1].parse().map_err(|_| {
101                        XlogError::Compilation(format!(
102                            "Decision-DNNF parse error at line {}: invalid node id '{}'",
103                            line_no + 1,
104                            tokens[1]
105                        ))
106                    })?;
107
108                    let kind = match tokens[0] {
109                        "o" => DdnnfNodeKind::Or,
110                        "a" => DdnnfNodeKind::And,
111                        "t" => DdnnfNodeKind::True,
112                        "f" => DdnnfNodeKind::False,
113                        _ => unreachable!(),
114                    };
115
116                    if nodes.insert(id, DdnnfNode { kind }).is_some() {
117                        return Err(XlogError::Compilation(format!(
118                            "Decision-DNNF parse error at line {}: duplicate node id {}",
119                            line_no + 1,
120                            id
121                        )));
122                    }
123                }
124                _ => {
125                    if tokens.len() < 2 {
126                        return Err(XlogError::Compilation(format!(
127                            "Decision-DNNF parse error at line {}: edge record missing dst",
128                            line_no + 1
129                        )));
130                    }
131                    let from: u32 = tokens[0].parse().map_err(|_| {
132                        XlogError::Compilation(format!(
133                            "Decision-DNNF parse error at line {}: invalid edge src '{}'",
134                            line_no + 1,
135                            tokens[0]
136                        ))
137                    })?;
138                    let to: u32 = tokens[1].parse().map_err(|_| {
139                        XlogError::Compilation(format!(
140                            "Decision-DNNF parse error at line {}: invalid edge dst '{}'",
141                            line_no + 1,
142                            tokens[1]
143                        ))
144                    })?;
145
146                    let mut lits: Vec<i32> = Vec::new();
147                    for &tok in &tokens[2..] {
148                        let lit: i32 = tok.parse().map_err(|_| {
149                            XlogError::Compilation(format!(
150                                "Decision-DNNF parse error at line {}: invalid literal '{}'",
151                                line_no + 1,
152                                tok
153                            ))
154                        })?;
155                        if lit == 0 {
156                            return Err(XlogError::Compilation(format!(
157                                "Decision-DNNF parse error at line {}: literal cannot be 0",
158                                line_no + 1
159                            )));
160                        }
161                        max_var = max_var.max(lit.unsigned_abs());
162                        lits.push(lit);
163                    }
164
165                    let edge_id = edges.len();
166                    edges.push(DdnnfEdge { from, to, lits });
167                    targets.insert(to);
168
169                    // outgoing filled later after validation.
170                    let _ = edge_id;
171                }
172            }
173        }
174
175        if nodes.is_empty() {
176            return Err(XlogError::Compilation(
177                "Decision-DNNF parse error: no nodes found".to_string(),
178            ));
179        }
180
181        for edge in &edges {
182            let from_kind = nodes.get(&edge.from).ok_or_else(|| {
183                XlogError::Compilation(format!(
184                    "Decision-DNNF parse error: edge references unknown src node {}",
185                    edge.from
186                ))
187            })?;
188            let _to_kind = nodes.get(&edge.to).ok_or_else(|| {
189                XlogError::Compilation(format!(
190                    "Decision-DNNF parse error: edge references unknown dst node {}",
191                    edge.to
192                ))
193            })?;
194
195            match from_kind.kind {
196                DdnnfNodeKind::Or | DdnnfNodeKind::And => {}
197                DdnnfNodeKind::True | DdnnfNodeKind::False => {
198                    return Err(XlogError::Compilation(format!(
199                        "Decision-DNNF parse error: leaf node {} cannot have outgoing edges",
200                        edge.from
201                    )));
202                }
203            }
204        }
205
206        let declared: BTreeSet<u32> = nodes.keys().copied().collect();
207        let target_set: BTreeSet<u32> = targets.into_iter().collect();
208        let roots: Vec<u32> = declared.difference(&target_set).copied().collect();
209        let root = match roots.as_slice() {
210            [only] => *only,
211            [] => {
212                return Err(XlogError::Compilation(
213                    "Decision-DNNF parse error: could not infer root (no root candidates)"
214                        .to_string(),
215                ))
216            }
217            many => {
218                return Err(XlogError::Compilation(format!(
219                    "Decision-DNNF parse error: could not infer unique root (candidates: {:?})",
220                    many
221                )))
222            }
223        };
224
225        let mut outgoing: BTreeMap<u32, Vec<usize>> = BTreeMap::new();
226        for (idx, edge) in edges.iter().enumerate() {
227            outgoing.entry(edge.from).or_default().push(idx);
228        }
229
230        // Optional cycle check (defensive).
231        Self::check_acyclic(root, &nodes, &edges, &outgoing)?;
232
233        Ok(Self {
234            root,
235            nodes,
236            edges,
237            outgoing,
238            max_var,
239        })
240    }
241
242    fn check_acyclic(
243        root: u32,
244        nodes: &BTreeMap<u32, DdnnfNode>,
245        edges: &[DdnnfEdge],
246        outgoing: &BTreeMap<u32, Vec<usize>>,
247    ) -> Result<()> {
248        let mut visiting: HashSet<u32> = HashSet::new();
249        let mut visited: HashSet<u32> = HashSet::new();
250
251        fn dfs(
252            node_id: u32,
253            nodes: &BTreeMap<u32, DdnnfNode>,
254            edges: &[DdnnfEdge],
255            outgoing: &BTreeMap<u32, Vec<usize>>,
256            visiting: &mut HashSet<u32>,
257            visited: &mut HashSet<u32>,
258        ) -> Result<()> {
259            if visited.contains(&node_id) {
260                return Ok(());
261            }
262            if !visiting.insert(node_id) {
263                return Err(XlogError::Compilation(format!(
264                    "Decision-DNNF parse error: cycle detected at node {}",
265                    node_id
266                )));
267            }
268
269            let node = nodes.get(&node_id).ok_or_else(|| {
270                XlogError::Compilation(format!(
271                    "Decision-DNNF parse error: unknown node {} during cycle check",
272                    node_id
273                ))
274            })?;
275
276            match node.kind {
277                DdnnfNodeKind::True | DdnnfNodeKind::False => {}
278                DdnnfNodeKind::Or | DdnnfNodeKind::And => {
279                    if let Some(out) = outgoing.get(&node_id) {
280                        for &edge_idx in out {
281                            let edge = &edges[edge_idx];
282                            dfs(edge.to, nodes, edges, outgoing, visiting, visited)?;
283                        }
284                    }
285                }
286            }
287
288            visiting.remove(&node_id);
289            visited.insert(node_id);
290            Ok(())
291        }
292
293        dfs(root, nodes, edges, outgoing, &mut visiting, &mut visited)
294    }
295
296    pub fn eval_log_wmc<F>(&self, var_log_weights: F) -> Result<f64>
297    where
298        F: Fn(u32) -> (f64, f64),
299    {
300        let mut memo: HashMap<u32, f64> = HashMap::new();
301
302        fn eval_node<F>(
303            node_id: u32,
304            ddnnf: &DecisionDnnf,
305            memo: &mut HashMap<u32, f64>,
306            var_log_weights: &F,
307        ) -> Result<f64>
308        where
309            F: Fn(u32) -> (f64, f64),
310        {
311            if let Some(&v) = memo.get(&node_id) {
312                return Ok(v);
313            }
314
315            let node = ddnnf.nodes.get(&node_id).ok_or_else(|| {
316                XlogError::Compilation(format!(
317                    "Decision-DNNF eval error: unknown node {}",
318                    node_id
319                ))
320            })?;
321
322            let value = match node.kind {
323                DdnnfNodeKind::True => 0.0,
324                DdnnfNodeKind::False => f64::NEG_INFINITY,
325                DdnnfNodeKind::And => {
326                    let out = ddnnf.outgoing.get(&node_id).ok_or_else(|| {
327                        XlogError::Compilation(format!(
328                            "Decision-DNNF eval error: AND node {} has no children",
329                            node_id
330                        ))
331                    })?;
332
333                    let mut acc = 0.0;
334                    for &edge_idx in out {
335                        let edge = &ddnnf.edges[edge_idx];
336                        let child = eval_node(edge.to, ddnnf, memo, var_log_weights)?;
337                        let mut lit_sum = 0.0;
338                        for &lit in &edge.lits {
339                            let var = lit.unsigned_abs();
340                            let (t, f) = validate_circuit_log_weight_pair(var_log_weights(var))?;
341                            lit_sum += if lit > 0 { t } else { f };
342                        }
343                        acc += lit_sum + child;
344                    }
345                    acc
346                }
347                DdnnfNodeKind::Or => {
348                    let out = ddnnf.outgoing.get(&node_id).ok_or_else(|| {
349                        XlogError::Compilation(format!(
350                            "Decision-DNNF eval error: OR node {} has no children",
351                            node_id
352                        ))
353                    })?;
354
355                    let mut branch_vals: Vec<f64> = Vec::with_capacity(out.len());
356                    for &edge_idx in out {
357                        let edge = &ddnnf.edges[edge_idx];
358                        let child = eval_node(edge.to, ddnnf, memo, var_log_weights)?;
359                        let mut lit_sum = 0.0;
360                        for &lit in &edge.lits {
361                            let var = lit.unsigned_abs();
362                            let (t, f) = validate_circuit_log_weight_pair(var_log_weights(var))?;
363                            lit_sum += if lit > 0 { t } else { f };
364                        }
365                        branch_vals.push(lit_sum + child);
366                    }
367                    circuit_logsumexp(&branch_vals)?
368                }
369            };
370
371            let value = validate_circuit_value(value)?;
372            memo.insert(node_id, value);
373            Ok(value)
374        }
375
376        eval_node(self.root, self, &mut memo, &var_log_weights)
377    }
378}
379
380#[cfg(test)]
381mod tests {
382    use super::*;
383
384    #[test]
385    fn test_parse_and_eval_identity_variable() {
386        // Represents the formula: x1
387        let nnf = r#"
388o 1 0
389t 2 0
390f 3 0
3911 2 1 0
3921 3 -1 0
393"#;
394
395        let ddnnf = DecisionDnnf::parse_str(nnf).unwrap();
396        assert_eq!(ddnnf.root(), 1);
397        assert_eq!(ddnnf.max_var(), 1);
398
399        let p = 0.3_f64;
400        let log_wmc = ddnnf
401            .eval_log_wmc(|var| match var {
402                1 => (p.ln(), (1.0 - p).ln()),
403                _ => panic!("unexpected var {}", var),
404            })
405            .unwrap();
406
407        assert!((log_wmc - p.ln()).abs() < 1e-9, "log_wmc={}", log_wmc);
408    }
409
410    #[test]
411    fn test_parse_detects_missing_terminator() {
412        let nnf = "t 1";
413        let err = DecisionDnnf::parse_str(nnf).unwrap_err();
414        let msg = err.to_string();
415        assert!(msg.contains("terminator"), "msg={}", msg);
416    }
417}