1use 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 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 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 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 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}