Skip to main content

xlog_prob/
xgcf.rs

1//! XLOG GPU Circuit Format (XGCF) - host-side representation and CPU evaluator.
2
3use std::collections::HashMap;
4
5use xlog_core::{Result, XlogError};
6
7use crate::kc::ddnnf::{DdnnfEdge, DdnnfNodeKind, DecisionDnnf};
8use crate::logsumexp::{
9    circuit_logsumexp, validate_circuit_gradient_values, validate_circuit_log_weight_pair,
10    validate_circuit_value,
11};
12
13#[derive(Debug, Clone, Copy, PartialEq, Eq)]
14#[repr(u8)]
15pub enum XgcfNodeType {
16    Const0 = 0,
17    Const1 = 1,
18    Lit = 2,
19    And = 3,
20    Or = 4,
21    Decision = 5,
22}
23
24#[derive(Debug, Clone)]
25pub struct Xgcf {
26    pub node_type: Vec<XgcfNodeType>,
27    pub child_offsets: Vec<u32>,
28    pub child_indices: Vec<u32>,
29    pub lit: Vec<i32>,
30    pub decision_var: Vec<u32>,
31    pub decision_child_false: Vec<u32>,
32    pub decision_child_true: Vec<u32>,
33    pub roots: Vec<u32>,
34    pub level_offsets: Vec<u32>,
35    pub level_nodes: Vec<u32>,
36}
37
38impl Xgcf {
39    pub fn from_ddnnf(ddnnf: &DecisionDnnf) -> Result<Self> {
40        XgcfBuilder::new(ddnnf).build()
41    }
42
43    /// Return a semantically equivalent circuit that is **smooth** with respect to
44    /// the subset of variables marked as random in `is_random_var`.
45    ///
46    /// Smoothness guarantees that, for any OR/DECISION node, all branches mention
47    /// the same set of random variables. This makes WMC evaluation and gradients
48    /// correct even when evidence/queries force random variables in non-smooth
49    /// Decision-DNNF outputs.
50    pub fn smooth_random_vars(&self, is_random_var: &[bool]) -> Result<Self> {
51        XgcfSmoother::new(self, is_random_var)?.smooth()
52    }
53
54    pub fn eval_log_wmc<F>(&self, var_log_weights: F) -> Result<f64>
55    where
56        F: Fn(u32) -> (f64, f64),
57    {
58        if self.roots.len() != 1 {
59            return Err(XlogError::Compilation(format!(
60                "XGCF eval expects exactly 1 root, got {}",
61                self.roots.len()
62            )));
63        }
64
65        let n = self.node_type.len();
66        if self.child_offsets.len() != n + 1 {
67            return Err(XlogError::Compilation(format!(
68                "XGCF invariant violation: child_offsets len {} != num_nodes+1 ({})",
69                self.child_offsets.len(),
70                n + 1
71            )));
72        }
73        if self.lit.len() != n
74            || self.decision_var.len() != n
75            || self.decision_child_false.len() != n
76            || self.decision_child_true.len() != n
77        {
78            return Err(XlogError::Compilation(
79                "XGCF invariant violation: per-node arrays length mismatch".to_string(),
80            ));
81        }
82        if self.level_offsets.is_empty() || *self.level_offsets.first().unwrap() != 0 {
83            return Err(XlogError::Compilation(
84                "XGCF invariant violation: level_offsets must start at 0".to_string(),
85            ));
86        }
87        if *self.level_offsets.last().unwrap() != self.level_nodes.len() as u32 {
88            return Err(XlogError::Compilation(
89                "XGCF invariant violation: level_offsets last != level_nodes.len".to_string(),
90            ));
91        }
92
93        let mut values: Vec<f64> = vec![0.0; n];
94
95        let num_levels = self.level_offsets.len() - 1;
96        for level in 0..num_levels {
97            let start = self.level_offsets[level] as usize;
98            let end = self.level_offsets[level + 1] as usize;
99            for &node_u32 in &self.level_nodes[start..end] {
100                let idx = node_u32 as usize;
101                let v = match self.node_type[idx] {
102                    XgcfNodeType::Const0 => f64::NEG_INFINITY,
103                    XgcfNodeType::Const1 => 0.0,
104                    XgcfNodeType::Lit => {
105                        let lit = self.lit[idx];
106                        if lit == 0 {
107                            return Err(XlogError::Compilation(format!(
108                                "XGCF invariant violation: LIT node {} has lit=0",
109                                idx
110                            )));
111                        }
112                        let var = lit.unsigned_abs();
113                        let (t, f) = validate_circuit_log_weight_pair(var_log_weights(var))?;
114                        if lit > 0 {
115                            t
116                        } else {
117                            f
118                        }
119                    }
120                    XgcfNodeType::And => {
121                        let c0 = self.child_offsets[idx] as usize;
122                        let c1 = self.child_offsets[idx + 1] as usize;
123                        if c0 == c1 {
124                            return Err(XlogError::Compilation(format!(
125                                "XGCF eval error: AND node {} has no children",
126                                idx
127                            )));
128                        }
129                        let mut acc = 0.0;
130                        for &child in &self.child_indices[c0..c1] {
131                            acc += values[child as usize];
132                        }
133                        acc
134                    }
135                    XgcfNodeType::Or => {
136                        let c0 = self.child_offsets[idx] as usize;
137                        let c1 = self.child_offsets[idx + 1] as usize;
138                        if c0 == c1 {
139                            return Err(XlogError::Compilation(format!(
140                                "XGCF eval error: OR node {} has no children",
141                                idx
142                            )));
143                        }
144                        let mut branch_vals: Vec<f64> = Vec::with_capacity(c1 - c0);
145                        for &child in &self.child_indices[c0..c1] {
146                            branch_vals.push(values[child as usize]);
147                        }
148                        circuit_logsumexp(&branch_vals)?
149                    }
150                    XgcfNodeType::Decision => {
151                        let var = self.decision_var[idx];
152                        if var == 0 {
153                            return Err(XlogError::Compilation(format!(
154                                "XGCF invariant violation: DECISION node {} has var=0",
155                                idx
156                            )));
157                        }
158                        let child_false = self.decision_child_false[idx] as usize;
159                        let child_true = self.decision_child_true[idx] as usize;
160                        let (t, f) = validate_circuit_log_weight_pair(var_log_weights(var))?;
161                        circuit_logsumexp(&[f + values[child_false], t + values[child_true]])?
162                    }
163                };
164                values[idx] = validate_circuit_value(v)?;
165            }
166        }
167
168        Ok(values[self.roots[0] as usize])
169    }
170
171    pub fn eval_log_wmc_and_grads(
172        &self,
173        var_log_weights: &[(f64, f64)],
174    ) -> Result<(f64, Vec<f64>, Vec<f64>)> {
175        if self.roots.len() != 1 {
176            return Err(XlogError::Compilation(format!(
177                "XGCF eval expects exactly 1 root, got {}",
178                self.roots.len()
179            )));
180        }
181
182        let n = self.node_type.len();
183        if self.child_offsets.len() != n + 1 {
184            return Err(XlogError::Compilation(format!(
185                "XGCF invariant violation: child_offsets len {} != num_nodes+1 ({})",
186                self.child_offsets.len(),
187                n + 1
188            )));
189        }
190        if self.lit.len() != n
191            || self.decision_var.len() != n
192            || self.decision_child_false.len() != n
193            || self.decision_child_true.len() != n
194        {
195            return Err(XlogError::Compilation(
196                "XGCF invariant violation: per-node arrays length mismatch".to_string(),
197            ));
198        }
199        if self.level_offsets.is_empty() || *self.level_offsets.first().unwrap() != 0 {
200            return Err(XlogError::Compilation(
201                "XGCF invariant violation: level_offsets must start at 0".to_string(),
202            ));
203        }
204        if *self.level_offsets.last().unwrap() != self.level_nodes.len() as u32 {
205            return Err(XlogError::Compilation(
206                "XGCF invariant violation: level_offsets last != level_nodes.len".to_string(),
207            ));
208        }
209
210        let mut max_var: u32 = 0;
211        for (&ty, &lit) in self.node_type.iter().zip(self.lit.iter()) {
212            if ty == XgcfNodeType::Lit && lit != 0 {
213                max_var = max_var.max(lit.unsigned_abs());
214            }
215        }
216        for &var in &self.decision_var {
217            max_var = max_var.max(var);
218        }
219
220        let weights_len = (max_var as usize) + 1;
221        if var_log_weights.len() < weights_len {
222            return Err(XlogError::Compilation(format!(
223                "XGCF eval expects weight table len >= {}, got {}",
224                weights_len,
225                var_log_weights.len()
226            )));
227        }
228        let mut values: Vec<f64> = vec![0.0; n];
229
230        let num_levels = self.level_offsets.len() - 1;
231        for level in 0..num_levels {
232            let start = self.level_offsets[level] as usize;
233            let end = self.level_offsets[level + 1] as usize;
234            for &node_u32 in &self.level_nodes[start..end] {
235                let idx = node_u32 as usize;
236                let v = match self.node_type[idx] {
237                    XgcfNodeType::Const0 => f64::NEG_INFINITY,
238                    XgcfNodeType::Const1 => 0.0,
239                    XgcfNodeType::Lit => {
240                        let lit = self.lit[idx];
241                        if lit == 0 {
242                            return Err(XlogError::Compilation(format!(
243                                "XGCF invariant violation: LIT node {} has lit=0",
244                                idx
245                            )));
246                        }
247                        let var = lit.unsigned_abs();
248                        let (t, f) =
249                            validate_circuit_log_weight_pair(var_log_weights[var as usize])?;
250                        if lit > 0 {
251                            t
252                        } else {
253                            f
254                        }
255                    }
256                    XgcfNodeType::And => {
257                        let c0 = self.child_offsets[idx] as usize;
258                        let c1 = self.child_offsets[idx + 1] as usize;
259                        if c0 == c1 {
260                            return Err(XlogError::Compilation(format!(
261                                "XGCF eval error: AND node {} has no children",
262                                idx
263                            )));
264                        }
265                        let mut acc = 0.0;
266                        for &child in &self.child_indices[c0..c1] {
267                            acc += values[child as usize];
268                        }
269                        acc
270                    }
271                    XgcfNodeType::Or => {
272                        let c0 = self.child_offsets[idx] as usize;
273                        let c1 = self.child_offsets[idx + 1] as usize;
274                        if c0 == c1 {
275                            return Err(XlogError::Compilation(format!(
276                                "XGCF eval error: OR node {} has no children",
277                                idx
278                            )));
279                        }
280                        let mut branch_vals: Vec<f64> = Vec::with_capacity(c1 - c0);
281                        for &child in &self.child_indices[c0..c1] {
282                            branch_vals.push(values[child as usize]);
283                        }
284                        circuit_logsumexp(&branch_vals)?
285                    }
286                    XgcfNodeType::Decision => {
287                        let var = self.decision_var[idx];
288                        if var == 0 {
289                            return Err(XlogError::Compilation(format!(
290                                "XGCF invariant violation: DECISION node {} has var=0",
291                                idx
292                            )));
293                        }
294                        let child_false = self.decision_child_false[idx] as usize;
295                        let child_true = self.decision_child_true[idx] as usize;
296                        let (t, f) =
297                            validate_circuit_log_weight_pair(var_log_weights[var as usize])?;
298                        circuit_logsumexp(&[f + values[child_false], t + values[child_true]])?
299                    }
300                };
301                values[idx] = validate_circuit_value(v)?;
302            }
303        }
304
305        let root_idx = self.roots[0] as usize;
306        let log_z = values[root_idx];
307
308        let mut adj: Vec<f64> = vec![0.0; n];
309        adj[root_idx] = 1.0;
310
311        let mut grad_true: Vec<f64> = vec![0.0; weights_len];
312        let mut grad_false: Vec<f64> = vec![0.0; weights_len];
313
314        for level in (0..num_levels).rev() {
315            let start = self.level_offsets[level] as usize;
316            let end = self.level_offsets[level + 1] as usize;
317            for &node_u32 in &self.level_nodes[start..end] {
318                let idx = node_u32 as usize;
319                let a = adj[idx];
320                if a == 0.0 {
321                    continue;
322                }
323
324                match self.node_type[idx] {
325                    XgcfNodeType::Const0 | XgcfNodeType::Const1 => {}
326                    XgcfNodeType::Lit => {
327                        let lit = self.lit[idx];
328                        if lit == 0 {
329                            return Err(XlogError::Compilation(format!(
330                                "XGCF invariant violation: LIT node {} has lit=0",
331                                idx
332                            )));
333                        }
334                        let var = lit.unsigned_abs() as usize;
335                        if lit > 0 {
336                            grad_true[var] += a;
337                        } else {
338                            grad_false[var] += a;
339                        }
340                    }
341                    XgcfNodeType::And => {
342                        let c0 = self.child_offsets[idx] as usize;
343                        let c1 = self.child_offsets[idx + 1] as usize;
344                        for &child in &self.child_indices[c0..c1] {
345                            adj[child as usize] += a;
346                        }
347                    }
348                    XgcfNodeType::Or => {
349                        let parent_v = values[idx];
350                        if parent_v.is_infinite() && parent_v.is_sign_negative() {
351                            continue;
352                        }
353                        let c0 = self.child_offsets[idx] as usize;
354                        let c1 = self.child_offsets[idx + 1] as usize;
355                        for &child in &self.child_indices[c0..c1] {
356                            let child_idx = child as usize;
357                            let child_v = values[child_idx];
358                            if child_v.is_infinite() && child_v.is_sign_negative() {
359                                continue;
360                            }
361                            let ratio = (child_v - parent_v).exp();
362                            if ratio != 0.0 {
363                                adj[child_idx] += a * ratio;
364                            }
365                        }
366                    }
367                    XgcfNodeType::Decision => {
368                        let var = self.decision_var[idx] as usize;
369                        if var == 0 {
370                            return Err(XlogError::Compilation(format!(
371                                "XGCF invariant violation: DECISION node {} has var=0",
372                                idx
373                            )));
374                        }
375
376                        let parent_v = values[idx];
377                        if parent_v.is_infinite() && parent_v.is_sign_negative() {
378                            continue;
379                        }
380
381                        let child_false = self.decision_child_false[idx] as usize;
382                        let child_true = self.decision_child_true[idx] as usize;
383                        let (t, f) = var_log_weights[var];
384
385                        let vf = values[child_false];
386                        let vt = values[child_true];
387
388                        let mut p_false = 0.0;
389                        let mut p_true = 0.0;
390                        if !(vf.is_infinite() && vf.is_sign_negative()) {
391                            p_false = (f + vf - parent_v).exp();
392                        }
393                        if !(vt.is_infinite() && vt.is_sign_negative()) {
394                            p_true = (t + vt - parent_v).exp();
395                        }
396
397                        if p_false != 0.0 {
398                            adj[child_false] += a * p_false;
399                            grad_false[var] += a * p_false;
400                        }
401                        if p_true != 0.0 {
402                            adj[child_true] += a * p_true;
403                            grad_true[var] += a * p_true;
404                        }
405                    }
406                }
407            }
408        }
409
410        validate_circuit_gradient_values(&grad_true, &grad_false)?;
411        Ok((log_z, grad_true, grad_false))
412    }
413}
414
415fn max_var_in_circuit(circuit: &Xgcf) -> Result<u32> {
416    let mut max_var: u32 = 0;
417    for (&ty, &lit) in circuit.node_type.iter().zip(circuit.lit.iter()) {
418        if ty == XgcfNodeType::Lit {
419            if lit == 0 {
420                return Err(XlogError::Compilation(
421                    "XGCF invariant violation: LIT node has lit=0".to_string(),
422                ));
423            }
424            max_var = max_var.max(lit.unsigned_abs());
425        }
426    }
427    for (&ty, &var) in circuit.node_type.iter().zip(circuit.decision_var.iter()) {
428        if ty == XgcfNodeType::Decision {
429            if var == 0 {
430                return Err(XlogError::Compilation(
431                    "XGCF invariant violation: DECISION node has var=0".to_string(),
432                ));
433            }
434            max_var = max_var.max(var);
435        }
436    }
437    Ok(max_var)
438}
439
440fn merge_union_sorted(a: &[u32], b: &[u32], out: &mut Vec<u32>) {
441    out.clear();
442    let mut i = 0usize;
443    let mut j = 0usize;
444    while i < a.len() && j < b.len() {
445        let va = a[i];
446        let vb = b[j];
447        if va == vb {
448            out.push(va);
449            i += 1;
450            j += 1;
451        } else if va < vb {
452            out.push(va);
453            i += 1;
454        } else {
455            out.push(vb);
456            j += 1;
457        }
458    }
459    if i < a.len() {
460        out.extend_from_slice(&a[i..]);
461    }
462    if j < b.len() {
463        out.extend_from_slice(&b[j..]);
464    }
465}
466
467fn sorted_difference(a: &[u32], b: &[u32], out: &mut Vec<u32>) {
468    out.clear();
469    let mut i = 0usize;
470    let mut j = 0usize;
471    while i < a.len() {
472        let va = a[i];
473        while j < b.len() && b[j] < va {
474            j += 1;
475        }
476        if j < b.len() && b[j] == va {
477            i += 1;
478            j += 1;
479            continue;
480        }
481        out.push(va);
482        i += 1;
483    }
484}
485
486fn insert_sorted_unique(sorted: &mut Vec<u32>, var: u32) {
487    match sorted.binary_search(&var) {
488        Ok(_) => {}
489        Err(pos) => sorted.insert(pos, var),
490    }
491}
492
493fn compute_random_support(circuit: &Xgcf, is_random_var: &[bool]) -> Result<Vec<Vec<u32>>> {
494    let n = circuit.node_type.len();
495    let mut support: Vec<Vec<u32>> = vec![Vec::new(); n];
496
497    let num_levels = circuit.level_offsets.len().saturating_sub(1);
498    for level in 0..num_levels {
499        let start = circuit.level_offsets[level] as usize;
500        let end = circuit.level_offsets[level + 1] as usize;
501        for &node_u32 in &circuit.level_nodes[start..end] {
502            let idx = node_u32 as usize;
503            match circuit.node_type[idx] {
504                XgcfNodeType::Const0 | XgcfNodeType::Const1 => {}
505                XgcfNodeType::Lit => {
506                    let lit = circuit.lit[idx];
507                    if lit == 0 {
508                        return Err(XlogError::Compilation(format!(
509                            "XGCF invariant violation: LIT node {} has lit=0",
510                            idx
511                        )));
512                    }
513                    let var = lit.unsigned_abs() as usize;
514                    if var < is_random_var.len() && is_random_var[var] {
515                        support[idx].push(var as u32);
516                    }
517                }
518                XgcfNodeType::And | XgcfNodeType::Or => {
519                    let c0 = circuit.child_offsets[idx] as usize;
520                    let c1 = circuit.child_offsets[idx + 1] as usize;
521                    let mut acc: Vec<u32> = Vec::new();
522                    let mut tmp: Vec<u32> = Vec::new();
523                    for &child in &circuit.child_indices[c0..c1] {
524                        let child_idx = child as usize;
525                        merge_union_sorted(&acc, &support[child_idx], &mut tmp);
526                        std::mem::swap(&mut acc, &mut tmp);
527                    }
528                    support[idx] = acc;
529                }
530                XgcfNodeType::Decision => {
531                    let var = circuit.decision_var[idx];
532                    if var == 0 {
533                        return Err(XlogError::Compilation(format!(
534                            "XGCF invariant violation: DECISION node {} has var=0",
535                            idx
536                        )));
537                    }
538                    let child_false = circuit.decision_child_false[idx] as usize;
539                    let child_true = circuit.decision_child_true[idx] as usize;
540
541                    let mut acc: Vec<u32> = Vec::new();
542                    merge_union_sorted(&support[child_false], &support[child_true], &mut acc);
543
544                    let var_usize = var as usize;
545                    if var_usize < is_random_var.len() && is_random_var[var_usize] {
546                        insert_sorted_unique(&mut acc, var);
547                    }
548                    support[idx] = acc;
549                }
550            }
551        }
552    }
553
554    Ok(support)
555}
556
557struct XgcfSmoother<'a> {
558    input: &'a Xgcf,
559    is_random_var: &'a [bool],
560    support: Vec<Vec<u32>>,
561}
562
563impl<'a> XgcfSmoother<'a> {
564    fn new(input: &'a Xgcf, is_random_var: &'a [bool]) -> Result<Self> {
565        let n = input.node_type.len();
566        if input.child_offsets.len() != n + 1 {
567            return Err(XlogError::Compilation(format!(
568                "XGCF invariant violation: child_offsets len {} != num_nodes+1 ({})",
569                input.child_offsets.len(),
570                n + 1
571            )));
572        }
573        if input.lit.len() != n
574            || input.decision_var.len() != n
575            || input.decision_child_false.len() != n
576            || input.decision_child_true.len() != n
577        {
578            return Err(XlogError::Compilation(
579                "XGCF invariant violation: per-node arrays length mismatch".to_string(),
580            ));
581        }
582
583        let max_var = max_var_in_circuit(input)?;
584        if is_random_var.len() <= (max_var as usize) {
585            return Err(XlogError::Compilation(format!(
586                "XGCF smoothing expects is_random_var len >= {}, got {}",
587                (max_var as usize) + 1,
588                is_random_var.len()
589            )));
590        }
591
592        let support = compute_random_support(input, is_random_var)?;
593        Ok(Self {
594            input,
595            is_random_var,
596            support,
597        })
598    }
599
600    fn smooth(&self) -> Result<Xgcf> {
601        XgcfSmoothBuilder::new(self.input).smooth(self.is_random_var, &self.support)
602    }
603}
604
605struct XgcfSmoothBuilder<'a> {
606    input: &'a Xgcf,
607    node_type: Vec<XgcfNodeType>,
608    child_offsets: Vec<u32>,
609    child_indices: Vec<u32>,
610    lit: Vec<i32>,
611    decision_var: Vec<u32>,
612    decision_child_false: Vec<u32>,
613    decision_child_true: Vec<u32>,
614    old_to_new: Vec<Option<u32>>,
615    lit_cache: HashMap<i32, u32>,
616    tautology_cache: HashMap<u32, u32>,
617    const0: u32,
618    const1: u32,
619}
620
621impl<'a> XgcfSmoothBuilder<'a> {
622    fn new(input: &'a Xgcf) -> Self {
623        let mut b = Self {
624            input,
625            node_type: Vec::new(),
626            child_offsets: Vec::new(),
627            child_indices: Vec::new(),
628            lit: Vec::new(),
629            decision_var: Vec::new(),
630            decision_child_false: Vec::new(),
631            decision_child_true: Vec::new(),
632            old_to_new: Vec::new(),
633            lit_cache: HashMap::new(),
634            tautology_cache: HashMap::new(),
635            const0: 0,
636            const1: 0,
637        };
638
639        b.const0 = b.push_const(false);
640        b.const1 = b.push_const(true);
641        b
642    }
643
644    fn push_base_node(&mut self, ty: XgcfNodeType) -> u32 {
645        let idx = u32::try_from(self.node_type.len()).expect("XGCF node index overflow");
646        self.node_type.push(ty);
647        self.child_offsets.push(self.child_indices.len() as u32);
648        self.lit.push(0);
649        self.decision_var.push(0);
650        self.decision_child_false.push(0);
651        self.decision_child_true.push(0);
652        idx
653    }
654
655    fn push_const(&mut self, value: bool) -> u32 {
656        self.push_base_node(if value {
657            XgcfNodeType::Const1
658        } else {
659            XgcfNodeType::Const0
660        })
661    }
662
663    fn get_lit_node(&mut self, lit: i32) -> Result<u32> {
664        if lit == 0 {
665            return Err(XlogError::Compilation(
666                "Cannot create XGCF LIT for 0 literal".to_string(),
667            ));
668        }
669        if let Some(&idx) = self.lit_cache.get(&lit) {
670            return Ok(idx);
671        }
672        let idx = self.push_base_node(XgcfNodeType::Lit);
673        self.lit[idx as usize] = lit;
674        self.lit_cache.insert(lit, idx);
675        Ok(idx)
676    }
677
678    fn push_and(&mut self, mut children: Vec<u32>) -> Result<u32> {
679        if children.contains(&self.const0) {
680            return Ok(self.const0);
681        }
682        children.retain(|&c| c != self.const1);
683        children.sort();
684        children.dedup();
685        match children.as_slice() {
686            [] => Ok(self.const1),
687            [only] => Ok(*only),
688            _ => {
689                let idx = self.push_base_node(XgcfNodeType::And);
690                self.child_indices.extend_from_slice(&children);
691                Ok(idx)
692            }
693        }
694    }
695
696    fn push_or(&mut self, mut children: Vec<u32>) -> Result<u32> {
697        children.retain(|&c| c != self.const0);
698        children.sort();
699        children.dedup();
700        match children.as_slice() {
701            [] => Ok(self.const0),
702            [only] => Ok(*only),
703            _ => {
704                let idx = self.push_base_node(XgcfNodeType::Or);
705                self.child_indices.extend_from_slice(&children);
706                Ok(idx)
707            }
708        }
709    }
710
711    fn push_decision(&mut self, var: u32, child_false: u32, child_true: u32) -> Result<u32> {
712        if var == 0 {
713            return Err(XlogError::Compilation(
714                "Cannot create XGCF DECISION with var=0".to_string(),
715            ));
716        }
717        let idx = self.push_base_node(XgcfNodeType::Decision);
718        self.decision_var[idx as usize] = var;
719        self.decision_child_false[idx as usize] = child_false;
720        self.decision_child_true[idx as usize] = child_true;
721        Ok(idx)
722    }
723
724    fn tautology_decision(&mut self, var: u32) -> Result<u32> {
725        if let Some(&idx) = self.tautology_cache.get(&var) {
726            return Ok(idx);
727        }
728        let idx = self.push_decision(var, self.const1, self.const1)?;
729        self.tautology_cache.insert(var, idx);
730        Ok(idx)
731    }
732
733    fn smooth(mut self, is_random_var: &[bool], support: &[Vec<u32>]) -> Result<Xgcf> {
734        let n = self.input.node_type.len();
735        self.old_to_new = vec![None; n];
736
737        let num_levels = self.input.level_offsets.len().saturating_sub(1);
738        for level in 0..num_levels {
739            let start = self.input.level_offsets[level] as usize;
740            let end = self.input.level_offsets[level + 1] as usize;
741            for &node_u32 in &self.input.level_nodes[start..end] {
742                let idx = node_u32 as usize;
743
744                let new_idx = match self.input.node_type[idx] {
745                    XgcfNodeType::Const0 => self.const0,
746                    XgcfNodeType::Const1 => self.const1,
747                    XgcfNodeType::Lit => {
748                        let lit = self.input.lit[idx];
749                        self.get_lit_node(lit)?
750                    }
751                    XgcfNodeType::And => {
752                        let c0 = self.input.child_offsets[idx] as usize;
753                        let c1 = self.input.child_offsets[idx + 1] as usize;
754                        let mut children: Vec<u32> = Vec::with_capacity(c1 - c0);
755                        for &child in &self.input.child_indices[c0..c1] {
756                            let child_idx = child as usize;
757                            let mapped = self.old_to_new[child_idx].ok_or_else(|| {
758                                XlogError::Compilation(format!(
759                                    "XGCF smoothing error: missing mapped child {} for AND node {}",
760                                    child_idx, idx
761                                ))
762                            })?;
763                            children.push(mapped);
764                        }
765                        self.push_and(children)?
766                    }
767                    XgcfNodeType::Or => {
768                        let parent_support = &support[idx];
769                        let c0 = self.input.child_offsets[idx] as usize;
770                        let c1 = self.input.child_offsets[idx + 1] as usize;
771                        let mut wrapped_children: Vec<u32> = Vec::with_capacity(c1 - c0);
772                        let mut missing: Vec<u32> = Vec::new();
773                        for &child in &self.input.child_indices[c0..c1] {
774                            let child_idx = child as usize;
775                            let child_new = self.old_to_new[child_idx].ok_or_else(|| {
776                                XlogError::Compilation(format!(
777                                    "XGCF smoothing error: missing mapped child {} for OR node {}",
778                                    child_idx, idx
779                                ))
780                            })?;
781
782                            let child_support = &support[child_idx];
783                            sorted_difference(parent_support, child_support, &mut missing);
784
785                            if missing.is_empty() {
786                                wrapped_children.push(child_new);
787                            } else {
788                                let mut and_children: Vec<u32> =
789                                    Vec::with_capacity(1 + missing.len());
790                                and_children.push(child_new);
791                                for &var in &missing {
792                                    let var_usize = var as usize;
793                                    if var_usize < is_random_var.len() && is_random_var[var_usize] {
794                                        and_children.push(self.tautology_decision(var)?);
795                                    }
796                                }
797                                wrapped_children.push(self.push_and(and_children)?);
798                            }
799                        }
800                        self.push_or(wrapped_children)?
801                    }
802                    XgcfNodeType::Decision => {
803                        let var = self.input.decision_var[idx];
804                        let child_false_old = self.input.decision_child_false[idx] as usize;
805                        let child_true_old = self.input.decision_child_true[idx] as usize;
806
807                        let child_false_new = self.old_to_new[child_false_old].ok_or_else(|| {
808                            XlogError::Compilation(format!(
809                                "XGCF smoothing error: missing mapped decision false child {} for node {}",
810                                child_false_old, idx
811                            ))
812                        })?;
813                        let child_true_new = self.old_to_new[child_true_old].ok_or_else(|| {
814                            XlogError::Compilation(format!(
815                                "XGCF smoothing error: missing mapped decision true child {} for node {}",
816                                child_true_old, idx
817                            ))
818                        })?;
819
820                        let mut union_children: Vec<u32> = Vec::new();
821                        merge_union_sorted(
822                            &support[child_false_old],
823                            &support[child_true_old],
824                            &mut union_children,
825                        );
826
827                        if union_children.binary_search(&var).is_ok() {
828                            return Err(XlogError::Compilation(format!(
829                                "XGCF smoothing error: decision var {} appears in child support at node {}",
830                                var, idx
831                            )));
832                        }
833
834                        let mut missing: Vec<u32> = Vec::new();
835
836                        sorted_difference(&union_children, &support[child_false_old], &mut missing);
837                        let new_false = if missing.is_empty() {
838                            child_false_new
839                        } else {
840                            let mut and_children: Vec<u32> = Vec::with_capacity(1 + missing.len());
841                            and_children.push(child_false_new);
842                            for &v in &missing {
843                                let v_usize = v as usize;
844                                if v_usize < is_random_var.len() && is_random_var[v_usize] {
845                                    and_children.push(self.tautology_decision(v)?);
846                                }
847                            }
848                            self.push_and(and_children)?
849                        };
850
851                        sorted_difference(&union_children, &support[child_true_old], &mut missing);
852                        let new_true = if missing.is_empty() {
853                            child_true_new
854                        } else {
855                            let mut and_children: Vec<u32> = Vec::with_capacity(1 + missing.len());
856                            and_children.push(child_true_new);
857                            for &v in &missing {
858                                let v_usize = v as usize;
859                                if v_usize < is_random_var.len() && is_random_var[v_usize] {
860                                    and_children.push(self.tautology_decision(v)?);
861                                }
862                            }
863                            self.push_and(and_children)?
864                        };
865
866                        self.push_decision(var, new_false, new_true)?
867                    }
868                };
869
870                self.old_to_new[idx] = Some(new_idx);
871            }
872        }
873
874        // Finalize offsets (sentinel).
875        self.child_offsets.push(self.child_indices.len() as u32);
876
877        let mut roots: Vec<u32> = Vec::with_capacity(self.input.roots.len());
878        for &root in &self.input.roots {
879            let idx = root as usize;
880            let mapped = self.old_to_new[idx].ok_or_else(|| {
881                XlogError::Compilation(format!(
882                    "XGCF smoothing error: missing mapped root node {}",
883                    idx
884                ))
885            })?;
886            roots.push(mapped);
887        }
888
889        let (level_offsets, level_nodes) = XgcfBuilder::levelize(
890            &self.node_type,
891            &self.child_offsets,
892            &self.child_indices,
893            &self.decision_child_false,
894            &self.decision_child_true,
895            &roots,
896        )?;
897
898        Ok(Xgcf {
899            node_type: self.node_type,
900            child_offsets: self.child_offsets,
901            child_indices: self.child_indices,
902            lit: self.lit,
903            decision_var: self.decision_var,
904            decision_child_false: self.decision_child_false,
905            decision_child_true: self.decision_child_true,
906            roots,
907            level_offsets,
908            level_nodes,
909        })
910    }
911}
912
913struct XgcfBuilder<'a> {
914    ddnnf: &'a DecisionDnnf,
915    node_type: Vec<XgcfNodeType>,
916    child_offsets: Vec<u32>,
917    child_indices: Vec<u32>,
918    lit: Vec<i32>,
919    decision_var: Vec<u32>,
920    decision_child_false: Vec<u32>,
921    decision_child_true: Vec<u32>,
922    lit_cache: HashMap<i32, u32>,
923    ddnnf_cache: HashMap<u32, u32>,
924    const0: u32,
925    const1: u32,
926}
927
928impl<'a> XgcfBuilder<'a> {
929    fn new(ddnnf: &'a DecisionDnnf) -> Self {
930        let mut builder = Self {
931            ddnnf,
932            node_type: Vec::new(),
933            child_offsets: Vec::new(),
934            child_indices: Vec::new(),
935            lit: Vec::new(),
936            decision_var: Vec::new(),
937            decision_child_false: Vec::new(),
938            decision_child_true: Vec::new(),
939            lit_cache: HashMap::new(),
940            ddnnf_cache: HashMap::new(),
941            const0: 0,
942            const1: 0,
943        };
944
945        builder.const0 = builder.push_const(false);
946        builder.const1 = builder.push_const(true);
947        builder
948    }
949
950    fn push_base_node(&mut self, ty: XgcfNodeType) -> u32 {
951        let idx = u32::try_from(self.node_type.len()).expect("XGCF node index overflow");
952        self.node_type.push(ty);
953        self.child_offsets.push(self.child_indices.len() as u32);
954        self.lit.push(0);
955        self.decision_var.push(0);
956        self.decision_child_false.push(0);
957        self.decision_child_true.push(0);
958        idx
959    }
960
961    fn push_const(&mut self, value: bool) -> u32 {
962        self.push_base_node(if value {
963            XgcfNodeType::Const1
964        } else {
965            XgcfNodeType::Const0
966        })
967    }
968
969    fn get_lit_node(&mut self, lit: i32) -> Result<u32> {
970        if lit == 0 {
971            return Err(XlogError::Compilation(
972                "Cannot create XGCF LIT for 0 literal".to_string(),
973            ));
974        }
975        if let Some(&idx) = self.lit_cache.get(&lit) {
976            return Ok(idx);
977        }
978        let idx = self.push_base_node(XgcfNodeType::Lit);
979        self.lit[idx as usize] = lit;
980        self.lit_cache.insert(lit, idx);
981        Ok(idx)
982    }
983
984    fn push_and(&mut self, mut children: Vec<u32>) -> Result<u32> {
985        if children.contains(&self.const0) {
986            return Ok(self.const0);
987        }
988        children.retain(|&c| c != self.const1);
989        children.sort();
990        children.dedup();
991        match children.as_slice() {
992            [] => Ok(self.const1),
993            [only] => Ok(*only),
994            _ => {
995                let idx = self.push_base_node(XgcfNodeType::And);
996                self.child_indices.extend_from_slice(&children);
997                Ok(idx)
998            }
999        }
1000    }
1001
1002    fn push_or(&mut self, mut children: Vec<u32>) -> Result<u32> {
1003        children.retain(|&c| c != self.const0);
1004        children.sort();
1005        children.dedup();
1006        match children.as_slice() {
1007            [] => Ok(self.const0),
1008            [only] => Ok(*only),
1009            _ => {
1010                let idx = self.push_base_node(XgcfNodeType::Or);
1011                self.child_indices.extend_from_slice(&children);
1012                Ok(idx)
1013            }
1014        }
1015    }
1016
1017    fn push_decision(&mut self, var: u32, child_false: u32, child_true: u32) -> Result<u32> {
1018        if var == 0 {
1019            return Err(XlogError::Compilation(
1020                "Cannot create XGCF DECISION with var=0".to_string(),
1021            ));
1022        }
1023        let idx = self.push_base_node(XgcfNodeType::Decision);
1024        self.decision_var[idx as usize] = var;
1025        self.decision_child_false[idx as usize] = child_false;
1026        self.decision_child_true[idx as usize] = child_true;
1027        Ok(idx)
1028    }
1029
1030    fn build(mut self) -> Result<Xgcf> {
1031        let root = self.convert_ddnnf_node(self.ddnnf.root())?;
1032
1033        // Finalize offsets (sentinel).
1034        self.child_offsets.push(self.child_indices.len() as u32);
1035
1036        let roots = vec![root];
1037        let (level_offsets, level_nodes) = Self::levelize(
1038            &self.node_type,
1039            &self.child_offsets,
1040            &self.child_indices,
1041            &self.decision_child_false,
1042            &self.decision_child_true,
1043            &roots,
1044        )?;
1045
1046        Ok(Xgcf {
1047            node_type: self.node_type,
1048            child_offsets: self.child_offsets,
1049            child_indices: self.child_indices,
1050            lit: self.lit,
1051            decision_var: self.decision_var,
1052            decision_child_false: self.decision_child_false,
1053            decision_child_true: self.decision_child_true,
1054            roots,
1055            level_offsets,
1056            level_nodes,
1057        })
1058    }
1059
1060    fn convert_ddnnf_node(&mut self, node_id: u32) -> Result<u32> {
1061        if let Some(&idx) = self.ddnnf_cache.get(&node_id) {
1062            return Ok(idx);
1063        }
1064        let kind = self.ddnnf.node_kind(node_id).ok_or_else(|| {
1065            XlogError::Compilation(format!("XGCF build error: unknown DDNNF node {}", node_id))
1066        })?;
1067
1068        let idx = match kind {
1069            DdnnfNodeKind::True => self.const1,
1070            DdnnfNodeKind::False => self.const0,
1071            DdnnfNodeKind::And => {
1072                let out = self.ddnnf.outgoing_edge_indices(node_id).ok_or_else(|| {
1073                    XlogError::Compilation(format!(
1074                        "XGCF build error: AND node {} has no outgoing edges",
1075                        node_id
1076                    ))
1077                })?;
1078                let mut child_nodes: Vec<u32> = Vec::with_capacity(out.len());
1079                for &edge_idx in out {
1080                    child_nodes.push(self.convert_ddnnf_edge_branch(edge_idx, None)?);
1081                }
1082                self.push_and(child_nodes)?
1083            }
1084            DdnnfNodeKind::Or => {
1085                let out = self.ddnnf.outgoing_edge_indices(node_id).ok_or_else(|| {
1086                    XlogError::Compilation(format!(
1087                        "XGCF build error: OR node {} has no outgoing edges",
1088                        node_id
1089                    ))
1090                })?;
1091                if out.len() == 2 {
1092                    let e0 = self.ddnnf.edge(out[0]).ok_or_else(|| {
1093                        XlogError::Compilation(format!("XGCF build error: missing edge {}", out[0]))
1094                    })?;
1095                    let e1 = self.ddnnf.edge(out[1]).ok_or_else(|| {
1096                        XlogError::Compilation(format!("XGCF build error: missing edge {}", out[1]))
1097                    })?;
1098
1099                    if let Some((var, edge_true, edge_false)) = infer_decision_var(e0, e1)? {
1100                        let edge_true = out[edge_true];
1101                        let edge_false = out[edge_false];
1102                        let child_true =
1103                            self.convert_ddnnf_edge_branch(edge_true, Some(var as i32))?;
1104                        let child_false =
1105                            self.convert_ddnnf_edge_branch(edge_false, Some(-(var as i32)))?;
1106                        self.push_decision(var, child_false, child_true)?
1107                    } else {
1108                        let mut child_nodes: Vec<u32> = Vec::with_capacity(out.len());
1109                        for &edge_idx in out {
1110                            child_nodes.push(self.convert_ddnnf_edge_branch(edge_idx, None)?);
1111                        }
1112                        self.push_or(child_nodes)?
1113                    }
1114                } else {
1115                    let mut child_nodes: Vec<u32> = Vec::with_capacity(out.len());
1116                    for &edge_idx in out {
1117                        child_nodes.push(self.convert_ddnnf_edge_branch(edge_idx, None)?);
1118                    }
1119                    self.push_or(child_nodes)?
1120                }
1121            }
1122        };
1123
1124        self.ddnnf_cache.insert(node_id, idx);
1125        Ok(idx)
1126    }
1127
1128    fn convert_ddnnf_edge_branch(&mut self, edge_idx: usize, drop_lit: Option<i32>) -> Result<u32> {
1129        let edge = self.ddnnf.edge(edge_idx).ok_or_else(|| {
1130            XlogError::Compilation(format!("XGCF build error: missing edge {}", edge_idx))
1131        })?;
1132
1133        let child = self.convert_ddnnf_node(edge.to)?;
1134
1135        let mut children: Vec<u32> = Vec::new();
1136        children.push(child);
1137
1138        let mut dropped = false;
1139        for &lit in &edge.lits {
1140            if let Some(dl) = drop_lit {
1141                if !dropped && lit == dl {
1142                    dropped = true;
1143                    continue;
1144                }
1145            }
1146            children.push(self.get_lit_node(lit)?);
1147        }
1148
1149        if let Some(dl) = drop_lit {
1150            if !dropped {
1151                return Err(XlogError::Compilation(format!(
1152                    "XGCF build error: expected to drop literal {} on edge {}->{} but not present",
1153                    dl, edge.from, edge.to
1154                )));
1155            }
1156        }
1157
1158        self.push_and(children)
1159    }
1160
1161    fn levelize(
1162        node_type: &[XgcfNodeType],
1163        child_offsets: &[u32],
1164        child_indices: &[u32],
1165        decision_child_false: &[u32],
1166        decision_child_true: &[u32],
1167        roots: &[u32],
1168    ) -> Result<(Vec<u32>, Vec<u32>)> {
1169        let n = node_type.len();
1170        let mut levels: Vec<Option<u32>> = vec![None; n];
1171        let mut visiting: Vec<bool> = vec![false; n];
1172
1173        #[allow(clippy::too_many_arguments)]
1174        fn level_of(
1175            idx: usize,
1176            node_type: &[XgcfNodeType],
1177            child_offsets: &[u32],
1178            child_indices: &[u32],
1179            decision_child_false: &[u32],
1180            decision_child_true: &[u32],
1181            levels: &mut [Option<u32>],
1182            visiting: &mut [bool],
1183        ) -> Result<u32> {
1184            if let Some(lvl) = levels[idx] {
1185                return Ok(lvl);
1186            }
1187            if visiting[idx] {
1188                return Err(XlogError::Compilation(format!(
1189                    "XGCF levelize error: cycle detected at node {}",
1190                    idx
1191                )));
1192            }
1193            visiting[idx] = true;
1194
1195            let lvl = match node_type[idx] {
1196                XgcfNodeType::Const0 | XgcfNodeType::Const1 | XgcfNodeType::Lit => 0,
1197                XgcfNodeType::And | XgcfNodeType::Or => {
1198                    let c0 = child_offsets[idx] as usize;
1199                    let c1 = child_offsets[idx + 1] as usize;
1200                    let mut max_child = 0u32;
1201                    for &child in &child_indices[c0..c1] {
1202                        max_child = max_child.max(level_of(
1203                            child as usize,
1204                            node_type,
1205                            child_offsets,
1206                            child_indices,
1207                            decision_child_false,
1208                            decision_child_true,
1209                            levels,
1210                            visiting,
1211                        )?);
1212                    }
1213                    max_child + 1
1214                }
1215                XgcfNodeType::Decision => {
1216                    let lf = level_of(
1217                        decision_child_false[idx] as usize,
1218                        node_type,
1219                        child_offsets,
1220                        child_indices,
1221                        decision_child_false,
1222                        decision_child_true,
1223                        levels,
1224                        visiting,
1225                    )?;
1226                    let lt = level_of(
1227                        decision_child_true[idx] as usize,
1228                        node_type,
1229                        child_offsets,
1230                        child_indices,
1231                        decision_child_false,
1232                        decision_child_true,
1233                        levels,
1234                        visiting,
1235                    )?;
1236                    lf.max(lt) + 1
1237                }
1238            };
1239
1240            visiting[idx] = false;
1241            levels[idx] = Some(lvl);
1242            Ok(lvl)
1243        }
1244
1245        for &root in roots {
1246            level_of(
1247                root as usize,
1248                node_type,
1249                child_offsets,
1250                child_indices,
1251                decision_child_false,
1252                decision_child_true,
1253                &mut levels,
1254                &mut visiting,
1255            )?;
1256        }
1257
1258        let max_level = levels.iter().flatten().copied().max().unwrap_or(0);
1259        let mut buckets: Vec<Vec<u32>> = vec![Vec::new(); (max_level as usize) + 1];
1260        for (i, lvl) in levels.iter().enumerate().take(n) {
1261            let Some(lvl) = lvl else {
1262                continue;
1263            };
1264            buckets[*lvl as usize].push(i as u32);
1265        }
1266
1267        let mut level_offsets: Vec<u32> = Vec::with_capacity(buckets.len() + 1);
1268        let mut level_nodes: Vec<u32> = Vec::new();
1269        level_offsets.push(0);
1270        for bucket in buckets {
1271            level_nodes.extend(bucket);
1272            level_offsets.push(level_nodes.len() as u32);
1273        }
1274        Ok((level_offsets, level_nodes))
1275    }
1276}
1277
1278fn infer_decision_var(e0: &DdnnfEdge, e1: &DdnnfEdge) -> Result<Option<(u32, usize, usize)>> {
1279    fn sign_map(lits: &[i32]) -> Result<HashMap<u32, bool>> {
1280        let mut map: HashMap<u32, bool> = HashMap::new();
1281        for &lit in lits {
1282            let var = lit.unsigned_abs();
1283            let sign = lit > 0;
1284            if let Some(prev) = map.insert(var, sign) {
1285                if prev != sign {
1286                    return Err(XlogError::Compilation(format!(
1287                        "XGCF build error: conflicting literals {} and {} in same branch",
1288                        var, lit
1289                    )));
1290                }
1291            }
1292        }
1293        Ok(map)
1294    }
1295
1296    let m0 = sign_map(&e0.lits)?;
1297    let m1 = sign_map(&e1.lits)?;
1298
1299    let mut candidates: Vec<u32> = Vec::new();
1300    for (var, &s0) in &m0 {
1301        if let Some(&s1) = m1.get(var) {
1302            if s0 != s1 {
1303                candidates.push(*var);
1304            }
1305        }
1306    }
1307
1308    if candidates.len() != 1 {
1309        return Ok(None);
1310    }
1311    let var = candidates[0];
1312
1313    let edge0_is_true = m0.get(&var).copied().unwrap_or(false);
1314    let (edge_true, edge_false) = if edge0_is_true {
1315        (0usize, 1usize)
1316    } else {
1317        (1usize, 0usize)
1318    };
1319
1320    Ok(Some((var, edge_true, edge_false)))
1321}
1322
1323#[cfg(test)]
1324mod tests {
1325    use super::*;
1326    use crate::kc::ddnnf::DecisionDnnf;
1327
1328    #[test]
1329    fn test_xgcf_matches_ddnnf_on_single_decision() {
1330        let nnf = r#"
1331o 1 0
1332t 2 0
1333f 3 0
13341 2 1 0
13351 3 -1 0
1336"#;
1337        let ddnnf = DecisionDnnf::parse_str(nnf).unwrap();
1338        let xgcf = Xgcf::from_ddnnf(&ddnnf).unwrap();
1339
1340        let p = 0.37_f64;
1341        let w = |var: u32| match var {
1342            1 => (p.ln(), (1.0 - p).ln()),
1343            _ => panic!("unexpected var {}", var),
1344        };
1345
1346        let a = ddnnf.eval_log_wmc(w).unwrap();
1347        let b = xgcf.eval_log_wmc(w).unwrap();
1348        assert!((a - b).abs() < 1e-9, "ddnnf={} xgcf={}", a, b);
1349    }
1350
1351    #[test]
1352    fn test_xgcf_matches_ddnnf_on_two_stage_decision() {
1353        // Formula: x1 OR x2, represented as a decision on x1, then x2.
1354        let nnf = r#"
1355o 1 0
1356o 2 0
1357t 3 0
1358f 4 0
13591 3 1 0
13601 2 -1 0
13612 3 2 0
13622 4 -2 0
1363"#;
1364        let ddnnf = DecisionDnnf::parse_str(nnf).unwrap();
1365        let xgcf = Xgcf::from_ddnnf(&ddnnf).unwrap();
1366
1367        let p1 = 0.2_f64;
1368        let p2 = 0.6_f64;
1369        let w = |var: u32| match var {
1370            1 => (p1.ln(), (1.0 - p1).ln()),
1371            2 => (p2.ln(), (1.0 - p2).ln()),
1372            _ => panic!("unexpected var {}", var),
1373        };
1374
1375        let a = ddnnf.eval_log_wmc(w).unwrap();
1376        let b = xgcf.eval_log_wmc(w).unwrap();
1377        assert!((a - b).abs() < 1e-9, "ddnnf={} xgcf={}", a, b);
1378    }
1379
1380    #[test]
1381    fn circuit_evaluators_share_logsumexp_contract() {
1382        let nnf = r#"
1383o 1 0
1384t 2 0
1385f 3 0
13861 2 1 0
13871 3 -1 0
1388"#;
1389        let ddnnf = DecisionDnnf::parse_str(nnf).unwrap();
1390        let xgcf = Xgcf::from_ddnnf(&ddnnf).unwrap();
1391
1392        let ddnnf_error = ddnnf
1393            .eval_log_wmc(|_| (f64::NAN, f64::NEG_INFINITY))
1394            .expect_err("Decision-DNNF must reject NaN weights");
1395        let xgcf_error = xgcf
1396            .eval_log_wmc(|_| (f64::NAN, f64::NEG_INFINITY))
1397            .expect_err("XGCF must reject NaN weights");
1398        let weights = vec![(0.0, 0.0), (f64::NAN, f64::NEG_INFINITY)];
1399        let gradient_error = xgcf
1400            .eval_log_wmc_and_grads(&weights)
1401            .expect_err("XGCF gradient evaluation must reject NaN weights");
1402
1403        for (name, error) in [
1404            ("Decision-DNNF", ddnnf_error),
1405            ("XGCF", xgcf_error),
1406            ("XGCF gradients", gradient_error),
1407        ] {
1408            assert!(
1409                matches!(&error, XlogError::Compilation(message) if message.contains("NaN")),
1410                "{name}: unexpected error: {error}"
1411            );
1412        }
1413
1414        let ddnnf_log_z = ddnnf.eval_log_wmc(|_| (f64::INFINITY, -1.0)).unwrap();
1415        let xgcf_log_z = xgcf.eval_log_wmc(|_| (f64::INFINITY, -1.0)).unwrap();
1416        let weights = vec![(0.0, 0.0), (f64::INFINITY, -1.0)];
1417        let gradient_error = xgcf
1418            .eval_log_wmc_and_grads(&weights)
1419            .expect_err("Decision normalization must reject non-finite gradients");
1420
1421        for (name, value) in [("Decision-DNNF", ddnnf_log_z), ("XGCF", xgcf_log_z)] {
1422            assert!(
1423                value.is_infinite() && value.is_sign_positive(),
1424                "{name}: expected positive infinity, got {value}"
1425            );
1426        }
1427        assert!(
1428            matches!(&gradient_error, XlogError::Compilation(message) if message.contains("gradient") && message.contains("non-finite")),
1429            "XGCF gradients: unexpected error: {gradient_error}"
1430        );
1431    }
1432
1433    #[test]
1434    fn deterministic_circuit_evaluators_reject_nan_weights() {
1435        let nnf = r#"
1436a 1 0
1437t 2 0
14381 2 1 0
1439"#;
1440        let ddnnf = DecisionDnnf::parse_str(nnf).unwrap();
1441        let xgcf = Xgcf::from_ddnnf(&ddnnf).unwrap();
1442
1443        let ddnnf_error = ddnnf
1444            .eval_log_wmc(|_| (f64::NAN, 0.0))
1445            .expect_err("deterministic Decision-DNNF must reject NaN weights");
1446        let xgcf_error = xgcf
1447            .eval_log_wmc(|_| (f64::NAN, 0.0))
1448            .expect_err("deterministic XGCF must reject NaN weights");
1449        let gradient_error = xgcf
1450            .eval_log_wmc_and_grads(&[(0.0, 0.0), (f64::NAN, 0.0)])
1451            .expect_err("deterministic XGCF gradients must reject NaN weights");
1452
1453        for (name, error) in [
1454            ("Decision-DNNF", ddnnf_error),
1455            ("XGCF", xgcf_error),
1456            ("XGCF gradients", gradient_error),
1457        ] {
1458            assert!(
1459                matches!(&error, XlogError::Compilation(message) if message.contains("NaN")),
1460                "{name}: unexpected error: {error}"
1461            );
1462        }
1463    }
1464
1465    #[test]
1466    fn deterministic_circuit_evaluators_reject_computed_nan() {
1467        let nnf = r#"
1468a 1 0
1469t 2 0
14701 2 1 -2 0
1471"#;
1472        let ddnnf = DecisionDnnf::parse_str(nnf).unwrap();
1473        let xgcf = Xgcf::from_ddnnf(&ddnnf).unwrap();
1474        let weights = [(0.0, 0.0), (f64::INFINITY, 0.0), (0.0, f64::NEG_INFINITY)];
1475        let weight = |var: u32| weights[var as usize];
1476
1477        let ddnnf_error = ddnnf
1478            .eval_log_wmc(weight)
1479            .expect_err("Decision-DNNF must reject +inf + -inf");
1480        let xgcf_error = xgcf
1481            .eval_log_wmc(weight)
1482            .expect_err("XGCF must reject +inf + -inf");
1483        let gradient_error = xgcf
1484            .eval_log_wmc_and_grads(&weights)
1485            .expect_err("XGCF gradients must reject computed NaN");
1486
1487        for (name, error) in [("Decision-DNNF", ddnnf_error), ("XGCF", xgcf_error)] {
1488            assert!(
1489                matches!(&error, XlogError::Compilation(message) if message.contains("NaN")),
1490                "{name}: unexpected error: {error}"
1491            );
1492        }
1493        assert!(
1494            matches!(&gradient_error, XlogError::Compilation(message) if message.contains("NaN")),
1495            "XGCF gradients: unexpected error: {gradient_error}"
1496        );
1497    }
1498
1499    #[test]
1500    fn gradient_evaluation_rejects_nan_when_another_weight_is_positive_infinity() {
1501        let nnf = r#"
1502a 1 0
1503t 2 0
15041 2 1 2 0
1505"#;
1506        let ddnnf = DecisionDnnf::parse_str(nnf).unwrap();
1507        let xgcf = Xgcf::from_ddnnf(&ddnnf).unwrap();
1508        let weights = [(0.0, 0.0), (f64::INFINITY, 0.0), (f64::NAN, 0.0)];
1509
1510        let error = xgcf
1511            .eval_log_wmc_and_grads(&weights)
1512            .expect_err("NaN must take precedence over positive infinity");
1513        assert!(
1514            matches!(&error, XlogError::Compilation(message) if message.contains("NaN")),
1515            "unexpected error: {error}"
1516        );
1517    }
1518
1519    #[test]
1520    fn reserved_weight_slot_zero_does_not_affect_value_or_gradients() {
1521        let nnf = r#"
1522a 1 0
1523t 2 0
15241 2 1 0
1525"#;
1526        let xgcf = Xgcf::from_ddnnf(&DecisionDnnf::parse_str(nnf).unwrap()).unwrap();
1527        let baseline_weights = [(0.0, 0.0), (-0.25, -0.75)];
1528        let reserved_non_finite = [(f64::NAN, f64::INFINITY), (-0.25, -0.75)];
1529
1530        let baseline_value = xgcf
1531            .eval_log_wmc(|var| baseline_weights[var as usize])
1532            .unwrap();
1533        let reserved_value = xgcf
1534            .eval_log_wmc(|var| reserved_non_finite[var as usize])
1535            .unwrap();
1536        assert_eq!(reserved_value, baseline_value);
1537
1538        let baseline_grads = xgcf.eval_log_wmc_and_grads(&baseline_weights).unwrap();
1539        let reserved_grads = xgcf
1540            .eval_log_wmc_and_grads(&reserved_non_finite)
1541            .expect("reserved slot 0 must not be validated or consumed");
1542        assert_eq!(reserved_grads, baseline_grads);
1543    }
1544
1545    #[test]
1546    fn deterministic_gradient_allows_selected_and_unused_positive_infinity() {
1547        let nnf = r#"
1548a 1 0
1549t 2 0
15501 2 3 0
1551"#;
1552        let xgcf = Xgcf::from_ddnnf(&DecisionDnnf::parse_str(nnf).unwrap()).unwrap();
1553        let weights = [
1554            (0.0, 0.0),
1555            (f64::INFINITY, f64::INFINITY),
1556            (f64::INFINITY, f64::INFINITY),
1557            (f64::INFINITY, f64::INFINITY),
1558        ];
1559
1560        let (log_z, grad_true, grad_false) = xgcf
1561            .eval_log_wmc_and_grads(&weights)
1562            .expect("+inf is valid when deterministic backward outputs stay finite");
1563        assert!(log_z.is_infinite() && log_z.is_sign_positive());
1564        assert_eq!(grad_true, vec![0.0, 0.0, 0.0, 1.0]);
1565        assert_eq!(grad_false, vec![0.0; 4]);
1566    }
1567}