Skip to main content

xlog_cuda_tests/harness/
xgcf.rs

1//! Helpers for testing XGCF circuit CUDA kernels.
2
3use std::ffi::c_void;
4
5use cudarc::driver::{DeviceSlice, LaunchConfig};
6use xlog_core::{Result, XlogError};
7use xlog_cuda::memory::TrackedCudaSlice;
8use xlog_cuda::{circuit_kernels, AsKernelParam, CudaFunction, LaunchAsync, CIRCUIT_MODULE};
9
10use super::TestContext;
11
12#[derive(Debug, Clone)]
13pub struct TinyXgcfSpec {
14    pub num_nodes: usize,
15    pub num_vars: usize,
16    pub root: u32,
17    pub node_type: Vec<u8>,
18    pub child_offsets: Vec<u32>,
19    pub child_indices: Vec<u32>,
20    pub lit: Vec<i32>,
21    pub decision_var: Vec<u32>,
22    pub decision_child_false: Vec<u32>,
23    pub decision_child_true: Vec<u32>,
24    pub level_nodes: Vec<u32>,
25    pub levels: Vec<(u32, u32)>,
26    pub var_log_true: Vec<f64>,
27    pub var_log_false: Vec<f64>,
28    pub expected_values: Vec<f64>,
29    pub expected_grad_true: Vec<f64>,
30    pub expected_grad_false: Vec<f64>,
31}
32
33#[derive(Debug, Clone)]
34pub struct TinyXgcfRun {
35    pub values: Vec<f64>,
36    pub adj: Vec<f64>,
37    pub grad_true: Vec<f64>,
38    pub grad_false: Vec<f64>,
39}
40
41/// Device-resident XGCF circuit + reusable buffers.
42///
43/// This is used by certification categories that validate *transfer efficiency* and *circuit reuse*.
44/// The key property: circuit structure is uploaded once; repeated evaluations reuse device buffers.
45pub struct TinyXgcfDevice {
46    pub num_nodes: usize,
47    pub num_vars: usize,
48    pub root: u32,
49    levels: Vec<(u32, u32)>,
50
51    // Cached kernel handles for performance-sensitive certification categories.
52    forward_fn: CudaFunction,
53    backward_propagate_fn: CudaFunction,
54    backward_decision_grad_fn: CudaFunction,
55    backward_lit_grad_fn: CudaFunction,
56
57    // Circuit structure (device-resident).
58    d_node_type: TrackedCudaSlice<u8>,
59    d_child_offsets: TrackedCudaSlice<u32>,
60    d_child_indices: TrackedCudaSlice<u32>,
61    d_lit: TrackedCudaSlice<i32>,
62    d_decision_var: TrackedCudaSlice<u32>,
63    d_decision_child_false: TrackedCudaSlice<u32>,
64    d_decision_child_true: TrackedCudaSlice<u32>,
65    d_level_nodes: TrackedCudaSlice<u32>,
66    d_level_offsets: TrackedCudaSlice<u32>,
67
68    // Per-evaluation inputs (device-resident).
69    d_var_log_true: TrackedCudaSlice<f64>,
70    d_var_log_false: TrackedCudaSlice<f64>,
71
72    // Per-evaluation outputs / scratch (device-resident).
73    d_values: TrackedCudaSlice<f64>,
74    d_adj: TrackedCudaSlice<f64>,
75    d_grad_true: TrackedCudaSlice<f64>,
76    d_grad_false: TrackedCudaSlice<f64>,
77}
78
79impl TinyXgcfDevice {
80    pub fn upload(ctx: &TestContext, spec: &TinyXgcfSpec) -> Result<Self> {
81        let device = ctx.device.inner();
82
83        let forward_fn = device
84            .get_func(CIRCUIT_MODULE, circuit_kernels::XGCF_FORWARD_LEVEL)
85            .ok_or_else(|| {
86                XlogError::Kernel(format!(
87                    "Kernel {} not found in {}",
88                    circuit_kernels::XGCF_FORWARD_LEVEL,
89                    CIRCUIT_MODULE
90                ))
91            })?;
92        let backward_propagate_fn = device
93            .get_func(
94                CIRCUIT_MODULE,
95                circuit_kernels::XGCF_BACKWARD_LEVEL_PROPAGATE,
96            )
97            .ok_or_else(|| {
98                XlogError::Kernel(format!(
99                    "Kernel {} not found in {}",
100                    circuit_kernels::XGCF_BACKWARD_LEVEL_PROPAGATE,
101                    CIRCUIT_MODULE
102                ))
103            })?;
104        let backward_decision_grad_fn = device
105            .get_func(
106                CIRCUIT_MODULE,
107                circuit_kernels::XGCF_BACKWARD_LEVEL_DECISION_GRAD,
108            )
109            .ok_or_else(|| {
110                XlogError::Kernel(format!(
111                    "Kernel {} not found in {}",
112                    circuit_kernels::XGCF_BACKWARD_LEVEL_DECISION_GRAD,
113                    CIRCUIT_MODULE
114                ))
115            })?;
116        let backward_lit_grad_fn = device
117            .get_func(
118                CIRCUIT_MODULE,
119                circuit_kernels::XGCF_BACKWARD_LEVEL_LIT_GRAD,
120            )
121            .ok_or_else(|| {
122                XlogError::Kernel(format!(
123                    "Kernel {} not found in {}",
124                    circuit_kernels::XGCF_BACKWARD_LEVEL_LIT_GRAD,
125                    CIRCUIT_MODULE
126                ))
127            })?;
128
129        let mut d_node_type = ctx.memory.alloc::<u8>(spec.node_type.len())?;
130        ctx.htod_sync_copy_into(&spec.node_type, &mut d_node_type)
131            .map_err(|e| XlogError::Kernel(format!("Failed to upload node_type: {}", e)))?;
132
133        let mut d_child_offsets = ctx.memory.alloc::<u32>(spec.child_offsets.len())?;
134        ctx.htod_sync_copy_into(&spec.child_offsets, &mut d_child_offsets)
135            .map_err(|e| XlogError::Kernel(format!("Failed to upload child_offsets: {}", e)))?;
136
137        let mut d_child_indices = ctx.memory.alloc::<u32>(spec.child_indices.len())?;
138        ctx.htod_sync_copy_into(&spec.child_indices, &mut d_child_indices)
139            .map_err(|e| XlogError::Kernel(format!("Failed to upload child_indices: {}", e)))?;
140
141        let mut d_lit = ctx.memory.alloc::<i32>(spec.lit.len())?;
142        ctx.htod_sync_copy_into(&spec.lit, &mut d_lit)
143            .map_err(|e| XlogError::Kernel(format!("Failed to upload lit: {}", e)))?;
144
145        let mut d_decision_var = ctx.memory.alloc::<u32>(spec.decision_var.len())?;
146        ctx.htod_sync_copy_into(&spec.decision_var, &mut d_decision_var)
147            .map_err(|e| XlogError::Kernel(format!("Failed to upload decision_var: {}", e)))?;
148
149        let mut d_decision_child_false =
150            ctx.memory.alloc::<u32>(spec.decision_child_false.len())?;
151        ctx.htod_sync_copy_into(&spec.decision_child_false, &mut d_decision_child_false)
152            .map_err(|e| {
153                XlogError::Kernel(format!("Failed to upload decision_child_false: {}", e))
154            })?;
155
156        let mut d_decision_child_true = ctx.memory.alloc::<u32>(spec.decision_child_true.len())?;
157        ctx.htod_sync_copy_into(&spec.decision_child_true, &mut d_decision_child_true)
158            .map_err(|e| {
159                XlogError::Kernel(format!("Failed to upload decision_child_true: {}", e))
160            })?;
161
162        let mut d_level_nodes = ctx.memory.alloc::<u32>(spec.level_nodes.len())?;
163        ctx.htod_sync_copy_into(&spec.level_nodes, &mut d_level_nodes)
164            .map_err(|e| XlogError::Kernel(format!("Failed to upload level_nodes: {}", e)))?;
165
166        // Device-resident level offsets (len = num_levels + 1) for level-aware kernels.
167        if spec.levels.is_empty() {
168            return Err(XlogError::Kernel(
169                "TinyXgcfSpec requires non-empty levels".to_string(),
170            ));
171        }
172        let mut level_offsets: Vec<u32> = Vec::with_capacity(spec.levels.len() + 1);
173        for &(offset, _len) in &spec.levels {
174            level_offsets.push(offset);
175        }
176        let (last_offset, last_len) = *spec.levels.last().unwrap();
177        level_offsets.push(last_offset + last_len);
178
179        if level_offsets[0] != 0 {
180            return Err(XlogError::Kernel(
181                "TinyXgcfSpec level_offsets must start at 0".to_string(),
182            ));
183        }
184        for (i, &(offset, len)) in spec.levels.iter().enumerate() {
185            let expected_next = offset + len;
186            if level_offsets[i] != offset || level_offsets[i + 1] != expected_next {
187                return Err(XlogError::Kernel(
188                    "TinyXgcfSpec levels must be contiguous and match offsets".to_string(),
189                ));
190            }
191        }
192        let total = *level_offsets.last().unwrap() as usize;
193        if total != spec.level_nodes.len() {
194            return Err(XlogError::Kernel(format!(
195                "TinyXgcfSpec level_nodes len {} != level_offsets.last {}",
196                spec.level_nodes.len(),
197                total
198            )));
199        }
200
201        let mut d_level_offsets = ctx.memory.alloc::<u32>(level_offsets.len())?;
202        ctx.htod_sync_copy_into(&level_offsets, &mut d_level_offsets)
203            .map_err(|e| XlogError::Kernel(format!("Failed to upload level_offsets: {}", e)))?;
204
205        let mut d_var_log_true = ctx.memory.alloc::<f64>(spec.var_log_true.len())?;
206        ctx.htod_sync_copy_into(&spec.var_log_true, &mut d_var_log_true)
207            .map_err(|e| XlogError::Kernel(format!("Failed to upload var_log_true: {}", e)))?;
208
209        let mut d_var_log_false = ctx.memory.alloc::<f64>(spec.var_log_false.len())?;
210        ctx.htod_sync_copy_into(&spec.var_log_false, &mut d_var_log_false)
211            .map_err(|e| XlogError::Kernel(format!("Failed to upload var_log_false: {}", e)))?;
212
213        let d_values = ctx.memory.alloc::<f64>(spec.num_nodes)?;
214        let d_adj = ctx.memory.alloc::<f64>(spec.num_nodes)?;
215        let d_grad_true = ctx.memory.alloc::<f64>(spec.num_vars + 1)?;
216        let d_grad_false = ctx.memory.alloc::<f64>(spec.num_vars + 1)?;
217
218        Ok(Self {
219            num_nodes: spec.num_nodes,
220            num_vars: spec.num_vars,
221            root: spec.root,
222            levels: spec.levels.clone(),
223            forward_fn,
224            backward_propagate_fn,
225            backward_decision_grad_fn,
226            backward_lit_grad_fn,
227            d_node_type,
228            d_child_offsets,
229            d_child_indices,
230            d_lit,
231            d_decision_var,
232            d_decision_child_false,
233            d_decision_child_true,
234            d_level_nodes,
235            d_level_offsets,
236            d_var_log_true,
237            d_var_log_false,
238            d_values,
239            d_adj,
240            d_grad_true,
241            d_grad_false,
242        })
243    }
244
245    fn launch_level_cached(
246        kernel: &CudaFunction,
247        num_level_nodes: u32,
248        params: &mut Vec<*mut c_void>,
249    ) -> Result<()> {
250        if num_level_nodes == 0 {
251            return Ok(());
252        }
253        let block_size = 256u32;
254        let num_blocks = (num_level_nodes + block_size - 1) / block_size;
255        let config = LaunchConfig {
256            grid_dim: (num_blocks, 1, 1),
257            block_dim: (block_size, 1, 1),
258            shared_mem_bytes: 0,
259        };
260        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
261        unsafe { kernel.clone().launch(config, params) }
262            .map_err(|e| XlogError::Kernel(format!("Failed to launch level kernel: {}", e)))?;
263        Ok(())
264    }
265
266    pub fn set_weights(
267        &mut self,
268        ctx: &TestContext,
269        log_true: &[f64],
270        log_false: &[f64],
271    ) -> Result<()> {
272        if log_true.len() != self.d_var_log_true.len()
273            || log_false.len() != self.d_var_log_false.len()
274        {
275            return Err(XlogError::Kernel(format!(
276                "Weight length mismatch: got (true={}, false={}), expected (true={}, false={})",
277                log_true.len(),
278                log_false.len(),
279                self.d_var_log_true.len(),
280                self.d_var_log_false.len()
281            )));
282        }
283        ctx.htod_sync_copy_into(log_true, &mut self.d_var_log_true)
284            .map_err(|e| XlogError::Kernel(format!("Failed to upload var_log_true: {}", e)))?;
285        ctx.htod_sync_copy_into(log_false, &mut self.d_var_log_false)
286            .map_err(|e| XlogError::Kernel(format!("Failed to upload var_log_false: {}", e)))?;
287        Ok(())
288    }
289
290    /// Launch forward kernels (no sync, no host transfers).
291    pub fn forward_launch(&mut self, _ctx: &TestContext) -> Result<()> {
292        for (level, &(_offset, len)) in self.levels.iter().enumerate() {
293            let level_u32 = level as u32;
294            let mut params: Vec<*mut c_void> = vec![
295                (&self.d_node_type).as_kernel_param(),
296                (&self.d_child_offsets).as_kernel_param(),
297                (&self.d_child_indices).as_kernel_param(),
298                (&self.d_lit).as_kernel_param(),
299                (&self.d_decision_var).as_kernel_param(),
300                (&self.d_decision_child_false).as_kernel_param(),
301                (&self.d_decision_child_true).as_kernel_param(),
302                (&self.d_level_nodes).as_kernel_param(),
303                (&self.d_level_offsets).as_kernel_param(),
304                level_u32.as_kernel_param(),
305                (&self.d_var_log_true).as_kernel_param(),
306                (&self.d_var_log_false).as_kernel_param(),
307                (&mut self.d_values).as_kernel_param(),
308            ];
309            Self::launch_level_cached(&self.forward_fn, len, &mut params)?;
310        }
311        Ok(())
312    }
313
314    pub fn forward_download_values(&mut self, ctx: &TestContext) -> Result<Vec<f64>> {
315        self.forward_launch(ctx)?;
316        ctx.sync_and_check()?;
317        ctx.dtoh_sync_copy(&self.d_values)
318            .map_err(|e| XlogError::Kernel(format!("Failed to download values: {}", e)))
319    }
320
321    pub fn forward_download_root(&mut self, ctx: &TestContext) -> Result<f64> {
322        self.forward_launch(ctx)?;
323        ctx.sync_and_check()?;
324        let root_idx: usize = self.root as usize;
325        if root_idx >= self.num_nodes {
326            return Err(XlogError::Kernel(format!(
327                "Root {} out of bounds for num_nodes {}",
328                self.root, self.num_nodes
329            )));
330        }
331        let root_view = self.d_values.slice(root_idx..(root_idx + 1));
332        let mut root_host = [0.0f64];
333        ctx.dtoh_sync_copy_into(&root_view, &mut root_host)
334            .map_err(|e| XlogError::Kernel(format!("Failed to download root value: {}", e)))?;
335        Ok(root_host[0])
336    }
337
338    /// Launch backward kernels using existing `d_values` (no sync, no host transfers).
339    pub fn backward_only_launch(&mut self, ctx: &TestContext) -> Result<()> {
340        let device = ctx.device.inner();
341        device
342            .memset_zeros(&mut self.d_adj)
343            .map_err(|e| XlogError::Kernel(format!("Failed to zero adj: {}", e)))?;
344        device
345            .memset_zeros(&mut self.d_grad_true)
346            .map_err(|e| XlogError::Kernel(format!("Failed to zero grad_true: {}", e)))?;
347        device
348            .memset_zeros(&mut self.d_grad_false)
349            .map_err(|e| XlogError::Kernel(format!("Failed to zero grad_false: {}", e)))?;
350
351        let root_idx: usize = self.root as usize;
352        if root_idx >= self.num_nodes {
353            return Err(XlogError::Kernel(format!(
354                "Root {} out of bounds for num_nodes {}",
355                self.root, self.num_nodes
356            )));
357        }
358        let mut root_view = self.d_adj.slice_mut(root_idx..(root_idx + 1));
359        ctx.htod_sync_copy_into(&[1.0f64], &mut root_view)
360            .map_err(|e| XlogError::Kernel(format!("Failed to set root adjoint: {}", e)))?;
361
362        for (level, &(_offset, len)) in self.levels.iter().enumerate().rev() {
363            let level_u32 = level as u32;
364            let mut params: Vec<*mut c_void> = vec![
365                (&self.d_node_type).as_kernel_param(),
366                (&self.d_child_offsets).as_kernel_param(),
367                (&self.d_child_indices).as_kernel_param(),
368                (&self.d_decision_var).as_kernel_param(),
369                (&self.d_decision_child_false).as_kernel_param(),
370                (&self.d_decision_child_true).as_kernel_param(),
371                (&self.d_level_nodes).as_kernel_param(),
372                (&self.d_level_offsets).as_kernel_param(),
373                level_u32.as_kernel_param(),
374                (&self.d_var_log_true).as_kernel_param(),
375                (&self.d_var_log_false).as_kernel_param(),
376                (&self.d_values).as_kernel_param(),
377                (&mut self.d_adj).as_kernel_param(),
378            ];
379            Self::launch_level_cached(&self.backward_propagate_fn, len, &mut params)?;
380        }
381
382        for (level, &(_offset, len)) in self.levels.iter().enumerate().rev() {
383            let level_u32 = level as u32;
384            let mut params: Vec<*mut c_void> = vec![
385                (&self.d_node_type).as_kernel_param(),
386                (&self.d_decision_var).as_kernel_param(),
387                (&self.d_decision_child_false).as_kernel_param(),
388                (&self.d_decision_child_true).as_kernel_param(),
389                (&self.d_level_nodes).as_kernel_param(),
390                (&self.d_level_offsets).as_kernel_param(),
391                level_u32.as_kernel_param(),
392                (&self.d_var_log_true).as_kernel_param(),
393                (&self.d_var_log_false).as_kernel_param(),
394                (&self.d_values).as_kernel_param(),
395                (&self.d_adj).as_kernel_param(),
396                (&mut self.d_grad_true).as_kernel_param(),
397                (&mut self.d_grad_false).as_kernel_param(),
398            ];
399            Self::launch_level_cached(&self.backward_decision_grad_fn, len, &mut params)?;
400        }
401
402        for (level, &(_offset, len)) in self.levels.iter().enumerate().rev() {
403            let level_u32 = level as u32;
404            let mut params: Vec<*mut c_void> = vec![
405                (&self.d_node_type).as_kernel_param(),
406                (&self.d_lit).as_kernel_param(),
407                (&self.d_level_nodes).as_kernel_param(),
408                (&self.d_level_offsets).as_kernel_param(),
409                level_u32.as_kernel_param(),
410                (&self.d_adj).as_kernel_param(),
411                (&mut self.d_grad_true).as_kernel_param(),
412                (&mut self.d_grad_false).as_kernel_param(),
413            ];
414            Self::launch_level_cached(&self.backward_lit_grad_fn, len, &mut params)?;
415        }
416
417        Ok(())
418    }
419
420    /// Convenience helper: forward + backward in one launch sequence (no sync, no host transfers).
421    pub fn forward_then_backward_launch(&mut self, ctx: &TestContext) -> Result<()> {
422        self.forward_launch(ctx)?;
423        self.backward_only_launch(ctx)
424    }
425}
426
427fn logsumexp2(a: f64, b: f64) -> f64 {
428    if a.is_nan() || b.is_nan() {
429        return f64::NAN;
430    }
431    let m = a.max(b);
432    if m.is_infinite() {
433        return m;
434    }
435    m + ((a - m).exp() + (b - m).exp()).ln()
436}
437
438/// Tiny Decision-DNNF-shaped XGCF circuit that exercises CONST/LIT/AND/OR/DECISION nodes.
439pub fn tiny_xgcf_spec() -> TinyXgcfSpec {
440    const CONST0: u8 = 0;
441    const CONST1: u8 = 1;
442    const LIT: u8 = 2;
443    const AND: u8 = 3;
444    const OR: u8 = 4;
445    const DECISION: u8 = 5;
446
447    // Node indices:
448    // 0: CONST1
449    // 1: LIT(+1)
450    // 2: LIT(-2)
451    // 3: AND(1,2)
452    // 4: DECISION(var3, child_f=0, child_t=3)
453    // 5: CONST0
454    // 6: OR(4,5)   (root)
455    let num_nodes = 7;
456    let root = 6u32;
457
458    let node_type: Vec<u8> = vec![CONST1, LIT, LIT, AND, DECISION, CONST0, OR];
459    let lit: Vec<i32> = vec![0, 1, -2, 0, 0, 0, 0];
460    let decision_var: Vec<u32> = vec![0, 0, 0, 0, 3, 0, 0];
461    let decision_child_false: Vec<u32> = vec![0, 0, 0, 0, 0, 0, 0];
462    let decision_child_true: Vec<u32> = vec![0, 0, 0, 0, 3, 0, 0];
463
464    let child_offsets: Vec<u32> = vec![0, 0, 0, 0, 2, 2, 2, 4];
465    let child_indices: Vec<u32> = vec![1, 2, 4, 5];
466
467    // Levels: [0,1,2,5], [3], [4], [6]
468    let level_nodes: Vec<u32> = vec![0, 1, 2, 5, 3, 4, 6];
469    let levels: Vec<(u32, u32)> = vec![(0, 4), (4, 1), (5, 1), (6, 1)];
470
471    let num_vars = 3usize;
472    let var_log_true: Vec<f64> = vec![
473        0.0,
474        0.7f64.ln(), // var1
475        0.2f64.ln(), // var2
476        0.6f64.ln(), // var3
477    ];
478    let var_log_false: Vec<f64> = vec![
479        0.0,
480        0.3f64.ln(), // var1
481        0.8f64.ln(), // var2
482        0.4f64.ln(), // var3
483    ];
484
485    let v0 = 0.0;
486    let v1 = var_log_true[1];
487    let v2 = var_log_false[2];
488    let v3 = v1 + v2;
489    let v4 = logsumexp2(var_log_false[3] + v0, var_log_true[3] + v3);
490    let v5 = f64::NEG_INFINITY;
491    let v6 = logsumexp2(v4, v5);
492
493    let expected_values: Vec<f64> = vec![v0, v1, v2, v3, v4, v5, v6];
494
495    let p_false = (var_log_false[3] + v0 - v4).exp();
496    let p_true = (var_log_true[3] + v3 - v4).exp();
497
498    let mut expected_grad_true = vec![0.0f64; num_vars + 1];
499    let mut expected_grad_false = vec![0.0f64; num_vars + 1];
500    expected_grad_true[1] = p_true; // LIT(+1)
501    expected_grad_false[2] = p_true; // LIT(-2)
502    expected_grad_true[3] = p_true; // DECISION var3
503    expected_grad_false[3] = p_false;
504
505    TinyXgcfSpec {
506        num_nodes,
507        num_vars,
508        root,
509        node_type,
510        child_offsets,
511        child_indices,
512        lit,
513        decision_var,
514        decision_child_false,
515        decision_child_true,
516        level_nodes,
517        levels,
518        var_log_true,
519        var_log_false,
520        expected_values,
521        expected_grad_true,
522        expected_grad_false,
523    }
524}
525
526fn launch_level(
527    ctx: &TestContext,
528    kernel_name: &str,
529    num_level_nodes: u32,
530    params: &mut Vec<*mut c_void>,
531) -> Result<()> {
532    if num_level_nodes == 0 {
533        return Ok(());
534    }
535    let device = ctx.device.inner();
536    let kernel = device
537        .get_func(CIRCUIT_MODULE, kernel_name)
538        .ok_or_else(|| {
539            XlogError::Kernel(format!(
540                "Kernel {} not found in {}",
541                kernel_name, CIRCUIT_MODULE
542            ))
543        })?;
544
545    let block_size = 256u32;
546    let num_blocks = (num_level_nodes + block_size - 1) / block_size;
547    let config = LaunchConfig {
548        grid_dim: (num_blocks, 1, 1),
549        block_dim: (block_size, 1, 1),
550        shared_mem_bytes: 0,
551    };
552
553    // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
554    unsafe { kernel.clone().launch(config, params) }
555        .map_err(|e| XlogError::Kernel(format!("Failed to launch {}: {}", kernel_name, e)))?;
556    Ok(())
557}
558
559pub fn run_tiny_xgcf_forward(ctx: &TestContext, spec: &TinyXgcfSpec) -> Result<Vec<f64>> {
560    let mut d_node_type = ctx.memory.alloc::<u8>(spec.node_type.len())?;
561    ctx.htod_sync_copy_into(&spec.node_type, &mut d_node_type)
562        .map_err(|e| XlogError::Kernel(format!("Failed to upload node_type: {}", e)))?;
563
564    let mut d_child_offsets = ctx.memory.alloc::<u32>(spec.child_offsets.len())?;
565    ctx.htod_sync_copy_into(&spec.child_offsets, &mut d_child_offsets)
566        .map_err(|e| XlogError::Kernel(format!("Failed to upload child_offsets: {}", e)))?;
567
568    let mut d_child_indices = ctx.memory.alloc::<u32>(spec.child_indices.len())?;
569    ctx.htod_sync_copy_into(&spec.child_indices, &mut d_child_indices)
570        .map_err(|e| XlogError::Kernel(format!("Failed to upload child_indices: {}", e)))?;
571
572    let mut d_lit = ctx.memory.alloc::<i32>(spec.lit.len())?;
573    ctx.htod_sync_copy_into(&spec.lit, &mut d_lit)
574        .map_err(|e| XlogError::Kernel(format!("Failed to upload lit: {}", e)))?;
575
576    let mut d_decision_var = ctx.memory.alloc::<u32>(spec.decision_var.len())?;
577    ctx.htod_sync_copy_into(&spec.decision_var, &mut d_decision_var)
578        .map_err(|e| XlogError::Kernel(format!("Failed to upload decision_var: {}", e)))?;
579
580    let mut d_decision_child_false = ctx.memory.alloc::<u32>(spec.decision_child_false.len())?;
581    ctx.htod_sync_copy_into(&spec.decision_child_false, &mut d_decision_child_false)
582        .map_err(|e| XlogError::Kernel(format!("Failed to upload decision_child_false: {}", e)))?;
583
584    let mut d_decision_child_true = ctx.memory.alloc::<u32>(spec.decision_child_true.len())?;
585    ctx.htod_sync_copy_into(&spec.decision_child_true, &mut d_decision_child_true)
586        .map_err(|e| XlogError::Kernel(format!("Failed to upload decision_child_true: {}", e)))?;
587
588    let mut d_level_nodes = ctx.memory.alloc::<u32>(spec.level_nodes.len())?;
589    ctx.htod_sync_copy_into(&spec.level_nodes, &mut d_level_nodes)
590        .map_err(|e| XlogError::Kernel(format!("Failed to upload level_nodes: {}", e)))?;
591
592    let mut level_offsets: Vec<u32> = Vec::with_capacity(spec.levels.len() + 1);
593    for &(offset, _len) in &spec.levels {
594        level_offsets.push(offset);
595    }
596    let (last_offset, last_len) = *spec
597        .levels
598        .last()
599        .ok_or_else(|| XlogError::Kernel("TinyXgcfSpec requires non-empty levels".to_string()))?;
600    level_offsets.push(last_offset + last_len);
601    if level_offsets[0] != 0 {
602        return Err(XlogError::Kernel(
603            "TinyXgcfSpec level_offsets must start at 0".to_string(),
604        ));
605    }
606    let mut d_level_offsets = ctx.memory.alloc::<u32>(level_offsets.len())?;
607    ctx.htod_sync_copy_into(&level_offsets, &mut d_level_offsets)
608        .map_err(|e| XlogError::Kernel(format!("Failed to upload level_offsets: {}", e)))?;
609
610    let mut level_offsets: Vec<u32> = Vec::with_capacity(spec.levels.len() + 1);
611    for &(offset, _len) in &spec.levels {
612        level_offsets.push(offset);
613    }
614    let (last_offset, last_len) = *spec
615        .levels
616        .last()
617        .ok_or_else(|| XlogError::Kernel("TinyXgcfSpec requires non-empty levels".to_string()))?;
618    level_offsets.push(last_offset + last_len);
619    if level_offsets[0] != 0 {
620        return Err(XlogError::Kernel(
621            "TinyXgcfSpec level_offsets must start at 0".to_string(),
622        ));
623    }
624    let mut d_level_offsets = ctx.memory.alloc::<u32>(level_offsets.len())?;
625    ctx.htod_sync_copy_into(&level_offsets, &mut d_level_offsets)
626        .map_err(|e| XlogError::Kernel(format!("Failed to upload level_offsets: {}", e)))?;
627
628    let mut d_var_log_true = ctx.memory.alloc::<f64>(spec.var_log_true.len())?;
629    ctx.htod_sync_copy_into(&spec.var_log_true, &mut d_var_log_true)
630        .map_err(|e| XlogError::Kernel(format!("Failed to upload var_log_true: {}", e)))?;
631
632    let mut d_var_log_false = ctx.memory.alloc::<f64>(spec.var_log_false.len())?;
633    ctx.htod_sync_copy_into(&spec.var_log_false, &mut d_var_log_false)
634        .map_err(|e| XlogError::Kernel(format!("Failed to upload var_log_false: {}", e)))?;
635
636    let mut d_values = ctx.memory.alloc::<f64>(spec.num_nodes)?;
637    let init_values = vec![0.0f64; spec.num_nodes];
638    ctx.htod_sync_copy_into(&init_values, &mut d_values)
639        .map_err(|e| XlogError::Kernel(format!("Failed to init values: {}", e)))?;
640
641    for (level, &(_offset, len)) in spec.levels.iter().enumerate() {
642        let level_u32 = level as u32;
643        let mut params: Vec<*mut c_void> = vec![
644            (&d_node_type).as_kernel_param(),
645            (&d_child_offsets).as_kernel_param(),
646            (&d_child_indices).as_kernel_param(),
647            (&d_lit).as_kernel_param(),
648            (&d_decision_var).as_kernel_param(),
649            (&d_decision_child_false).as_kernel_param(),
650            (&d_decision_child_true).as_kernel_param(),
651            (&d_level_nodes).as_kernel_param(),
652            (&d_level_offsets).as_kernel_param(),
653            level_u32.as_kernel_param(),
654            (&d_var_log_true).as_kernel_param(),
655            (&d_var_log_false).as_kernel_param(),
656            (&mut d_values).as_kernel_param(),
657        ];
658        launch_level(ctx, circuit_kernels::XGCF_FORWARD_LEVEL, len, &mut params)?;
659    }
660
661    ctx.sync_and_check()?;
662
663    ctx.dtoh_sync_copy(&d_values)
664        .map_err(|e| XlogError::Kernel(format!("Failed to download values: {}", e)))
665}
666
667pub fn run_tiny_xgcf_backward(ctx: &TestContext, spec: &TinyXgcfSpec) -> Result<TinyXgcfRun> {
668    let mut d_node_type = ctx.memory.alloc::<u8>(spec.node_type.len())?;
669    ctx.htod_sync_copy_into(&spec.node_type, &mut d_node_type)
670        .map_err(|e| XlogError::Kernel(format!("Failed to upload node_type: {}", e)))?;
671
672    let mut d_child_offsets = ctx.memory.alloc::<u32>(spec.child_offsets.len())?;
673    ctx.htod_sync_copy_into(&spec.child_offsets, &mut d_child_offsets)
674        .map_err(|e| XlogError::Kernel(format!("Failed to upload child_offsets: {}", e)))?;
675
676    let mut d_child_indices = ctx.memory.alloc::<u32>(spec.child_indices.len())?;
677    ctx.htod_sync_copy_into(&spec.child_indices, &mut d_child_indices)
678        .map_err(|e| XlogError::Kernel(format!("Failed to upload child_indices: {}", e)))?;
679
680    let mut d_lit = ctx.memory.alloc::<i32>(spec.lit.len())?;
681    ctx.htod_sync_copy_into(&spec.lit, &mut d_lit)
682        .map_err(|e| XlogError::Kernel(format!("Failed to upload lit: {}", e)))?;
683
684    let mut d_decision_var = ctx.memory.alloc::<u32>(spec.decision_var.len())?;
685    ctx.htod_sync_copy_into(&spec.decision_var, &mut d_decision_var)
686        .map_err(|e| XlogError::Kernel(format!("Failed to upload decision_var: {}", e)))?;
687
688    let mut d_decision_child_false = ctx.memory.alloc::<u32>(spec.decision_child_false.len())?;
689    ctx.htod_sync_copy_into(&spec.decision_child_false, &mut d_decision_child_false)
690        .map_err(|e| XlogError::Kernel(format!("Failed to upload decision_child_false: {}", e)))?;
691
692    let mut d_decision_child_true = ctx.memory.alloc::<u32>(spec.decision_child_true.len())?;
693    ctx.htod_sync_copy_into(&spec.decision_child_true, &mut d_decision_child_true)
694        .map_err(|e| XlogError::Kernel(format!("Failed to upload decision_child_true: {}", e)))?;
695
696    let mut d_level_nodes = ctx.memory.alloc::<u32>(spec.level_nodes.len())?;
697    ctx.htod_sync_copy_into(&spec.level_nodes, &mut d_level_nodes)
698        .map_err(|e| XlogError::Kernel(format!("Failed to upload level_nodes: {}", e)))?;
699
700    let mut level_offsets: Vec<u32> = Vec::with_capacity(spec.levels.len() + 1);
701    for &(offset, _len) in &spec.levels {
702        level_offsets.push(offset);
703    }
704    let (last_offset, last_len) = *spec
705        .levels
706        .last()
707        .ok_or_else(|| XlogError::Kernel("TinyXgcfSpec requires non-empty levels".to_string()))?;
708    level_offsets.push(last_offset + last_len);
709    if level_offsets[0] != 0 {
710        return Err(XlogError::Kernel(
711            "TinyXgcfSpec level_offsets must start at 0".to_string(),
712        ));
713    }
714    let mut d_level_offsets = ctx.memory.alloc::<u32>(level_offsets.len())?;
715    ctx.htod_sync_copy_into(&level_offsets, &mut d_level_offsets)
716        .map_err(|e| XlogError::Kernel(format!("Failed to upload level_offsets: {}", e)))?;
717
718    let mut d_var_log_true = ctx.memory.alloc::<f64>(spec.var_log_true.len())?;
719    ctx.htod_sync_copy_into(&spec.var_log_true, &mut d_var_log_true)
720        .map_err(|e| XlogError::Kernel(format!("Failed to upload var_log_true: {}", e)))?;
721
722    let mut d_var_log_false = ctx.memory.alloc::<f64>(spec.var_log_false.len())?;
723    ctx.htod_sync_copy_into(&spec.var_log_false, &mut d_var_log_false)
724        .map_err(|e| XlogError::Kernel(format!("Failed to upload var_log_false: {}", e)))?;
725
726    let mut d_values = ctx.memory.alloc::<f64>(spec.num_nodes)?;
727    let init_values = vec![0.0f64; spec.num_nodes];
728    ctx.htod_sync_copy_into(&init_values, &mut d_values)
729        .map_err(|e| XlogError::Kernel(format!("Failed to init values: {}", e)))?;
730
731    for (level, &(_offset, len)) in spec.levels.iter().enumerate() {
732        let level_u32 = level as u32;
733        let mut params: Vec<*mut c_void> = vec![
734            (&d_node_type).as_kernel_param(),
735            (&d_child_offsets).as_kernel_param(),
736            (&d_child_indices).as_kernel_param(),
737            (&d_lit).as_kernel_param(),
738            (&d_decision_var).as_kernel_param(),
739            (&d_decision_child_false).as_kernel_param(),
740            (&d_decision_child_true).as_kernel_param(),
741            (&d_level_nodes).as_kernel_param(),
742            (&d_level_offsets).as_kernel_param(),
743            level_u32.as_kernel_param(),
744            (&d_var_log_true).as_kernel_param(),
745            (&d_var_log_false).as_kernel_param(),
746            (&mut d_values).as_kernel_param(),
747        ];
748        launch_level(ctx, circuit_kernels::XGCF_FORWARD_LEVEL, len, &mut params)?;
749    }
750
751    // adj[root] = 1, others 0
752    let mut adj_init = vec![0.0f64; spec.num_nodes];
753    let root_idx: usize = spec.root as usize;
754    if root_idx >= adj_init.len() {
755        return Err(XlogError::Kernel(format!(
756            "Root {} out of bounds for num_nodes {}",
757            spec.root, spec.num_nodes
758        )));
759    }
760    adj_init[root_idx] = 1.0;
761    let mut d_adj = ctx.memory.alloc::<f64>(spec.num_nodes)?;
762    ctx.htod_sync_copy_into(&adj_init, &mut d_adj)
763        .map_err(|e| XlogError::Kernel(format!("Failed to init adj: {}", e)))?;
764
765    let mut d_grad_true = ctx.memory.alloc::<f64>(spec.num_vars + 1)?;
766    let mut d_grad_false = ctx.memory.alloc::<f64>(spec.num_vars + 1)?;
767    let grad_init = vec![0.0f64; spec.num_vars + 1];
768    ctx.htod_sync_copy_into(&grad_init, &mut d_grad_true)
769        .map_err(|e| XlogError::Kernel(format!("Failed to init grad_true: {}", e)))?;
770    ctx.htod_sync_copy_into(&grad_init, &mut d_grad_false)
771        .map_err(|e| XlogError::Kernel(format!("Failed to init grad_false: {}", e)))?;
772
773    for (level, &(_offset, len)) in spec.levels.iter().enumerate().rev() {
774        let level_u32 = level as u32;
775        let mut params: Vec<*mut c_void> = vec![
776            (&d_node_type).as_kernel_param(),
777            (&d_child_offsets).as_kernel_param(),
778            (&d_child_indices).as_kernel_param(),
779            (&d_decision_var).as_kernel_param(),
780            (&d_decision_child_false).as_kernel_param(),
781            (&d_decision_child_true).as_kernel_param(),
782            (&d_level_nodes).as_kernel_param(),
783            (&d_level_offsets).as_kernel_param(),
784            level_u32.as_kernel_param(),
785            (&d_var_log_true).as_kernel_param(),
786            (&d_var_log_false).as_kernel_param(),
787            (&d_values).as_kernel_param(),
788            (&mut d_adj).as_kernel_param(),
789        ];
790        launch_level(
791            ctx,
792            circuit_kernels::XGCF_BACKWARD_LEVEL_PROPAGATE,
793            len,
794            &mut params,
795        )?;
796    }
797
798    for (level, &(_offset, len)) in spec.levels.iter().enumerate().rev() {
799        let level_u32 = level as u32;
800        let mut params: Vec<*mut c_void> = vec![
801            (&d_node_type).as_kernel_param(),
802            (&d_decision_var).as_kernel_param(),
803            (&d_decision_child_false).as_kernel_param(),
804            (&d_decision_child_true).as_kernel_param(),
805            (&d_level_nodes).as_kernel_param(),
806            (&d_level_offsets).as_kernel_param(),
807            level_u32.as_kernel_param(),
808            (&d_var_log_true).as_kernel_param(),
809            (&d_var_log_false).as_kernel_param(),
810            (&d_values).as_kernel_param(),
811            (&d_adj).as_kernel_param(),
812            (&mut d_grad_true).as_kernel_param(),
813            (&mut d_grad_false).as_kernel_param(),
814        ];
815        launch_level(
816            ctx,
817            circuit_kernels::XGCF_BACKWARD_LEVEL_DECISION_GRAD,
818            len,
819            &mut params,
820        )?;
821    }
822
823    for (level, &(_offset, len)) in spec.levels.iter().enumerate().rev() {
824        let level_u32 = level as u32;
825        let mut params: Vec<*mut c_void> = vec![
826            (&d_node_type).as_kernel_param(),
827            (&d_lit).as_kernel_param(),
828            (&d_level_nodes).as_kernel_param(),
829            (&d_level_offsets).as_kernel_param(),
830            level_u32.as_kernel_param(),
831            (&d_adj).as_kernel_param(),
832            (&mut d_grad_true).as_kernel_param(),
833            (&mut d_grad_false).as_kernel_param(),
834        ];
835        launch_level(
836            ctx,
837            circuit_kernels::XGCF_BACKWARD_LEVEL_LIT_GRAD,
838            len,
839            &mut params,
840        )?;
841    }
842
843    ctx.sync_and_check()?;
844
845    let values = ctx
846        .dtoh_sync_copy(&d_values)
847        .map_err(|e| XlogError::Kernel(format!("Failed to download values: {}", e)))?;
848    let adj = ctx
849        .dtoh_sync_copy(&d_adj)
850        .map_err(|e| XlogError::Kernel(format!("Failed to download adj: {}", e)))?;
851    let grad_true = ctx
852        .dtoh_sync_copy(&d_grad_true)
853        .map_err(|e| XlogError::Kernel(format!("Failed to download grad_true: {}", e)))?;
854    let grad_false = ctx
855        .dtoh_sync_copy(&d_grad_false)
856        .map_err(|e| XlogError::Kernel(format!("Failed to download grad_false: {}", e)))?;
857
858    Ok(TinyXgcfRun {
859        values,
860        adj,
861        grad_true,
862        grad_false,
863    })
864}
865
866/// Generate a single-literal circuit: root = Lit(+var)
867pub fn gen_single_lit_circuit(var: u32) -> TinyXgcfSpec {
868    const LIT: u8 = 2;
869
870    let num_nodes = 1;
871    let num_vars = var as usize;
872    let root = 0;
873
874    let node_type = vec![LIT];
875    let child_offsets = vec![0, 0];
876    let child_indices = vec![];
877    let lit = vec![var as i32];
878    let decision_var = vec![0];
879    let decision_child_false = vec![0];
880    let decision_child_true = vec![0];
881    let level_nodes = vec![0];
882    let levels = vec![(0, 1)];
883
884    let mut var_log_true = vec![0.0; num_vars + 1];
885    let mut var_log_false = vec![0.0; num_vars + 1];
886    var_log_true[var as usize] = 0.7_f64.ln();
887    var_log_false[var as usize] = 0.3_f64.ln();
888
889    let expected_values = vec![var_log_true[var as usize]];
890    let mut expected_grad_true = vec![0.0; num_vars + 1];
891    let expected_grad_false = vec![0.0; num_vars + 1];
892    expected_grad_true[var as usize] = 1.0;
893
894    TinyXgcfSpec {
895        num_nodes,
896        num_vars,
897        root,
898        node_type,
899        child_offsets,
900        child_indices,
901        lit,
902        decision_var,
903        decision_child_false,
904        decision_child_true,
905        level_nodes,
906        levels,
907        var_log_true,
908        var_log_false,
909        expected_values,
910        expected_grad_true,
911        expected_grad_false,
912    }
913}
914
915/// Generate an AND circuit: root = AND(Lit(+1), Lit(+2))
916pub fn gen_and_circuit() -> TinyXgcfSpec {
917    const LIT: u8 = 2;
918    const AND: u8 = 3;
919
920    let num_nodes = 3;
921    let num_vars = 2;
922    let root = 2;
923
924    let node_type = vec![LIT, LIT, AND];
925    let child_offsets = vec![0, 0, 0, 2];
926    let child_indices = vec![0, 1];
927    let lit = vec![1, 2, 0];
928    let decision_var = vec![0, 0, 0];
929    let decision_child_false = vec![0, 0, 0];
930    let decision_child_true = vec![0, 0, 0];
931    let level_nodes = vec![0, 1, 2];
932    let levels = vec![(0, 2), (2, 1)];
933
934    let p1 = 0.7_f64;
935    let p2 = 0.6_f64;
936    let var_log_true = vec![0.0, p1.ln(), p2.ln()];
937    let var_log_false = vec![0.0, (1.0 - p1).ln(), (1.0 - p2).ln()];
938
939    let v0 = var_log_true[1];
940    let v1 = var_log_true[2];
941    let v2 = v0 + v1;
942    let expected_values = vec![v0, v1, v2];
943    let expected_grad_true = vec![0.0, 1.0, 1.0];
944    let expected_grad_false = vec![0.0, 0.0, 0.0];
945
946    TinyXgcfSpec {
947        num_nodes,
948        num_vars,
949        root,
950        node_type,
951        child_offsets,
952        child_indices,
953        lit,
954        decision_var,
955        decision_child_false,
956        decision_child_true,
957        level_nodes,
958        levels,
959        var_log_true,
960        var_log_false,
961        expected_values,
962        expected_grad_true,
963        expected_grad_false,
964    }
965}
966
967/// Generate an OR circuit: root = OR(Lit(+1), Lit(+2))
968pub fn gen_or_circuit() -> TinyXgcfSpec {
969    const LIT: u8 = 2;
970    const OR: u8 = 4;
971
972    let num_nodes = 3;
973    let num_vars = 2;
974    let root = 2;
975
976    let node_type = vec![LIT, LIT, OR];
977    let child_offsets = vec![0, 0, 0, 2];
978    let child_indices = vec![0, 1];
979    let lit = vec![1, 2, 0];
980    let decision_var = vec![0, 0, 0];
981    let decision_child_false = vec![0, 0, 0];
982    let decision_child_true = vec![0, 0, 0];
983    let level_nodes = vec![0, 1, 2];
984    let levels = vec![(0, 2), (2, 1)];
985
986    let p1 = 0.7_f64;
987    let p2 = 0.6_f64;
988    let var_log_true = vec![0.0, p1.ln(), p2.ln()];
989    let var_log_false = vec![0.0, (1.0 - p1).ln(), (1.0 - p2).ln()];
990
991    let v0 = var_log_true[1];
992    let v1 = var_log_true[2];
993    let v2 = logsumexp2(v0, v1);
994    let expected_values = vec![v0, v1, v2];
995
996    let p_child0 = (v0 - v2).exp();
997    let p_child1 = (v1 - v2).exp();
998    let expected_grad_true = vec![0.0, p_child0, p_child1];
999    let expected_grad_false = vec![0.0, 0.0, 0.0];
1000
1001    TinyXgcfSpec {
1002        num_nodes,
1003        num_vars,
1004        root,
1005        node_type,
1006        child_offsets,
1007        child_indices,
1008        lit,
1009        decision_var,
1010        decision_child_false,
1011        decision_child_true,
1012        level_nodes,
1013        levels,
1014        var_log_true,
1015        var_log_false,
1016        expected_values,
1017        expected_grad_true,
1018        expected_grad_false,
1019    }
1020}
1021
1022/// Generate a Decision circuit: root = Decision(var, false_child=Const1, true_child=Lit(+1))
1023pub fn gen_decision_circuit() -> TinyXgcfSpec {
1024    const CONST1: u8 = 1;
1025    const LIT: u8 = 2;
1026    const DECISION: u8 = 5;
1027
1028    let num_nodes = 3;
1029    let num_vars = 2;
1030    let root = 2;
1031
1032    let node_type = vec![CONST1, LIT, DECISION];
1033    let child_offsets = vec![0, 0, 0, 0];
1034    let child_indices = vec![];
1035    let lit = vec![0, 1, 0];
1036    let decision_var = vec![0, 0, 2];
1037    let decision_child_false = vec![0, 0, 0];
1038    let decision_child_true = vec![0, 0, 1];
1039    let level_nodes = vec![0, 1, 2];
1040    let levels = vec![(0, 2), (2, 1)];
1041
1042    let p1 = 0.7_f64;
1043    let p2 = 0.6_f64;
1044    let var_log_true = vec![0.0, p1.ln(), p2.ln()];
1045    let var_log_false = vec![0.0, (1.0 - p1).ln(), (1.0 - p2).ln()];
1046
1047    let v0 = 0.0;
1048    let v1 = var_log_true[1];
1049    let v2 = logsumexp2(var_log_false[2] + v0, var_log_true[2] + v1);
1050    let expected_values = vec![v0, v1, v2];
1051
1052    let p_false = (var_log_false[2] + v0 - v2).exp();
1053    let p_true = (var_log_true[2] + v1 - v2).exp();
1054
1055    let expected_grad_true = vec![0.0, p_true, p_true];
1056    let expected_grad_false = vec![0.0, 0.0, p_false];
1057
1058    TinyXgcfSpec {
1059        num_nodes,
1060        num_vars,
1061        root,
1062        node_type,
1063        child_offsets,
1064        child_indices,
1065        lit,
1066        decision_var,
1067        decision_child_false,
1068        decision_child_true,
1069        level_nodes,
1070        levels,
1071        var_log_true,
1072        var_log_false,
1073        expected_values,
1074        expected_grad_true,
1075        expected_grad_false,
1076    }
1077}
1078
1079/// Generate a large circuit with N parallel literals under an OR node
1080pub fn gen_large_or_circuit(num_vars: usize) -> TinyXgcfSpec {
1081    const LIT: u8 = 2;
1082    const OR: u8 = 4;
1083
1084    let num_nodes = num_vars + 1;
1085    let root = num_vars as u32;
1086
1087    let mut node_type = vec![LIT; num_vars];
1088    node_type.push(OR);
1089
1090    let mut child_offsets: Vec<u32> = (0..=num_vars).map(|_| 0).collect();
1091    child_offsets.push(num_vars as u32);
1092
1093    let child_indices: Vec<u32> = (0..num_vars as u32).collect();
1094
1095    let mut lit: Vec<i32> = (1..=num_vars as i32).collect();
1096    lit.push(0);
1097
1098    let decision_var = vec![0; num_nodes];
1099    let decision_child_false = vec![0; num_nodes];
1100    let decision_child_true = vec![0; num_nodes];
1101
1102    let mut level_nodes: Vec<u32> = (0..num_vars as u32).collect();
1103    level_nodes.push(root);
1104    let levels = vec![(0, num_vars as u32), (num_vars as u32, 1)];
1105
1106    let p = 0.5_f64;
1107    let mut var_log_true = vec![0.0; num_vars + 1];
1108    let mut var_log_false = vec![0.0; num_vars + 1];
1109    for i in 1..=num_vars {
1110        var_log_true[i] = p.ln();
1111        var_log_false[i] = (1.0 - p).ln();
1112    }
1113
1114    let lit_val = p.ln();
1115    let mut expected_values = vec![lit_val; num_vars];
1116    let or_val = lit_val + (num_vars as f64).ln();
1117    expected_values.push(or_val);
1118
1119    let grad_per_lit = 1.0 / num_vars as f64;
1120    let mut expected_grad_true = vec![0.0; num_vars + 1];
1121    for i in 1..=num_vars {
1122        expected_grad_true[i] = grad_per_lit;
1123    }
1124    let expected_grad_false = vec![0.0; num_vars + 1];
1125
1126    TinyXgcfSpec {
1127        num_nodes,
1128        num_vars,
1129        root,
1130        node_type,
1131        child_offsets,
1132        child_indices,
1133        lit,
1134        decision_var,
1135        decision_child_false,
1136        decision_child_true,
1137        level_nodes,
1138        levels,
1139        var_log_true,
1140        var_log_false,
1141        expected_values,
1142        expected_grad_true,
1143        expected_grad_false,
1144    }
1145}
1146
1147/// Generate a deep chain circuit: AND(AND(AND(...Lit(1)...)))
1148pub fn gen_deep_chain_circuit(depth: usize) -> TinyXgcfSpec {
1149    const LIT: u8 = 2;
1150    const AND: u8 = 3;
1151
1152    let num_nodes = depth + 1;
1153    let num_vars = 1;
1154    let root = depth as u32;
1155
1156    let mut node_type = vec![LIT];
1157    for _ in 0..depth {
1158        node_type.push(AND);
1159    }
1160
1161    let mut child_offsets: Vec<u32> = vec![0];
1162    let mut child_indices: Vec<u32> = vec![];
1163    for i in 0..depth {
1164        child_offsets.push(child_indices.len() as u32);
1165        child_indices.push(i as u32);
1166    }
1167    child_offsets.push(child_indices.len() as u32);
1168
1169    let mut lit = vec![1i32];
1170    lit.extend(vec![0i32; depth]);
1171
1172    let decision_var = vec![0; num_nodes];
1173    let decision_child_false = vec![0; num_nodes];
1174    let decision_child_true = vec![0; num_nodes];
1175
1176    let level_nodes: Vec<u32> = (0..num_nodes as u32).collect();
1177    let levels: Vec<(u32, u32)> = (0..num_nodes).map(|i| (i as u32, 1)).collect();
1178
1179    let p = 0.7_f64;
1180    let var_log_true = vec![0.0, p.ln()];
1181    let var_log_false = vec![0.0, (1.0 - p).ln()];
1182
1183    let lit_val = p.ln();
1184    let expected_values = vec![lit_val; num_nodes];
1185
1186    let expected_grad_true = vec![0.0, 1.0];
1187    let expected_grad_false = vec![0.0, 0.0];
1188
1189    TinyXgcfSpec {
1190        num_nodes,
1191        num_vars,
1192        root,
1193        node_type,
1194        child_offsets,
1195        child_indices,
1196        lit,
1197        decision_var,
1198        decision_child_false,
1199        decision_child_true,
1200        level_nodes,
1201        levels,
1202        var_log_true,
1203        var_log_false,
1204        expected_values,
1205        expected_grad_true,
1206        expected_grad_false,
1207    }
1208}
1209
1210/// Compute numerical gradient for verification
1211pub fn numerical_gradient(
1212    ctx: &TestContext,
1213    spec: &TinyXgcfSpec,
1214    var: usize,
1215    eps: f64,
1216) -> xlog_core::Result<(f64, f64)> {
1217    let mut spec_plus = spec.clone();
1218    let mut spec_minus = spec.clone();
1219    spec_plus.var_log_true[var] += eps;
1220    spec_minus.var_log_true[var] -= eps;
1221
1222    let values_plus = run_tiny_xgcf_forward(ctx, &spec_plus)?;
1223    let values_minus = run_tiny_xgcf_forward(ctx, &spec_minus)?;
1224
1225    let grad_true =
1226        (values_plus[spec.root as usize] - values_minus[spec.root as usize]) / (2.0 * eps);
1227
1228    let mut spec_plus = spec.clone();
1229    let mut spec_minus = spec.clone();
1230    spec_plus.var_log_false[var] += eps;
1231    spec_minus.var_log_false[var] -= eps;
1232
1233    let values_plus = run_tiny_xgcf_forward(ctx, &spec_plus)?;
1234    let values_minus = run_tiny_xgcf_forward(ctx, &spec_minus)?;
1235
1236    let grad_false =
1237        (values_plus[spec.root as usize] - values_minus[spec.root as usize]) / (2.0 * eps);
1238
1239    Ok((grad_true, grad_false))
1240}
1241
1242#[cfg(test)]
1243mod tests {
1244    use super::logsumexp2;
1245
1246    #[test]
1247    fn circuit_logsumexp2_edge_policy() {
1248        for (a, b) in [
1249            (f64::INFINITY, -2.0),
1250            (-2.0, f64::INFINITY),
1251            (f64::INFINITY, f64::NEG_INFINITY),
1252        ] {
1253            let value = logsumexp2(a, b);
1254            assert!(
1255                value.is_infinite() && value.is_sign_positive(),
1256                "logsumexp2({a}, {b}) expected +inf, got {value}"
1257            );
1258        }
1259
1260        for (a, b) in [(f64::NAN, 0.0), (0.0, f64::NAN), (f64::INFINITY, f64::NAN)] {
1261            assert!(
1262                logsumexp2(a, b).is_nan(),
1263                "logsumexp2({a}, {b}) must give NaN precedence"
1264            );
1265        }
1266    }
1267}