Skip to main content

xlog_prob/
gpu.rs

1//! GPU evaluator for XGCF circuits (CUDA).
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::provider::{
9    arith_kernels, d4_kernels, filter_kernels, ARITH_MODULE, D4_MODULE, FILTER_MODULE,
10};
11use xlog_cuda::{circuit_kernels, AsKernelParam, CudaKernelProvider, LaunchAsync, CIRCUIT_MODULE};
12
13use crate::compilation::gpu_d4::exclusive_scan_u32_inplace;
14use crate::logsumexp::validate_circuit_log_weight_pair;
15#[cfg(feature = "host-io")]
16use crate::logsumexp::{validate_circuit_gradient_values, validate_circuit_value};
17use crate::xgcf::{Xgcf, XgcfNodeType};
18
19/// Device-resident circuit buffers produced by the GPU compiler.
20///
21/// This matches the XGCF node layout used by `kernels/circuit.cu` and the SAT verifier CNF encoder
22/// in `kernels/sat.cu`.
23pub struct GpuCircuitBuilder {
24    pub node_type: TrackedCudaSlice<u8>,
25    pub child_offsets: TrackedCudaSlice<u32>,
26    pub child_indices: TrackedCudaSlice<u32>,
27    pub lit: TrackedCudaSlice<i32>,
28    pub decision_var: TrackedCudaSlice<u32>,
29    pub decision_child_false: TrackedCudaSlice<u32>,
30    pub decision_child_true: TrackedCudaSlice<u32>,
31}
32
33/// Device layout metadata for XGCF construction.
34pub struct GpuCircuitLayout {
35    pub num_nodes: u32,
36    pub num_edges: u32,
37    pub num_levels: u32,
38    pub level_offsets: TrackedCudaSlice<u32>,
39    pub level_nodes: TrackedCudaSlice<u32>,
40    pub root: u32,
41    pub max_var: u32,
42    pub num_nodes_device: Option<TrackedCudaSlice<u32>>,
43    pub num_edges_device: Option<TrackedCudaSlice<u32>>,
44}
45
46pub struct GpuXgcf {
47    node_type: TrackedCudaSlice<u8>,
48    child_offsets: TrackedCudaSlice<u32>,
49    child_indices: TrackedCudaSlice<u32>,
50    lit: TrackedCudaSlice<i32>,
51    decision_var: TrackedCudaSlice<u32>,
52    decision_child_false: TrackedCudaSlice<u32>,
53    decision_child_true: TrackedCudaSlice<u32>,
54    level_nodes: TrackedCudaSlice<u32>,
55    level_offsets: TrackedCudaSlice<u32>,
56    /// Optional host mirror for efficient per-level launch sizing when the circuit was uploaded
57    /// from host (`GpuXgcf::upload`). GPU-native compilation paths do not populate this.
58    level_offsets_host: Option<Vec<u32>>,
59    node_cap: u32,
60    edge_cap: u32,
61    num_levels: u32,
62    root: u32,
63    max_var: u32,
64    meta_num_nodes: TrackedCudaSlice<u32>,
65    meta_num_edges: TrackedCudaSlice<u32>,
66    var_log_true: TrackedCudaSlice<f64>,
67    var_log_false: TrackedCudaSlice<f64>,
68    values: TrackedCudaSlice<f64>,
69    adj: TrackedCudaSlice<f64>,
70    grad_true: TrackedCudaSlice<f64>,
71    grad_false: TrackedCudaSlice<f64>,
72    free_var_mask: Option<TrackedCudaSlice<u8>>,
73}
74
75fn checked_gpu_u32_len(context: &str, len: usize) -> Result<u32> {
76    u32::try_from(len)
77        .map_err(|_| XlogError::Compilation(format!("{context} exceeds u32::MAX: {len}")))
78}
79
80fn checked_gpu_len_add_one(context: &str, len: usize) -> Result<usize> {
81    len.checked_add(1)
82        .ok_or_else(|| XlogError::Compilation(format!("{context} length overflow")))
83}
84
85fn checked_gpu_launch_blocks(context: &str, item_count: usize, block_size: u32) -> Result<u32> {
86    let item_count = u32::try_from(item_count).map_err(|_| {
87        XlogError::Kernel(format!(
88            "{context} launch item count exceeds u32::MAX: {item_count}"
89        ))
90    })?;
91    item_count
92        .checked_add(block_size - 1)
93        .map(|rounded| rounded / block_size)
94        .ok_or_else(|| XlogError::Kernel(format!("{context} launch grid overflow")))
95}
96
97fn checked_host_level_width(level_offsets: &[u32], level: usize) -> Result<usize> {
98    let start = level_offsets[level];
99    let end = level_offsets[level + 1];
100    if end < start {
101        return Err(XlogError::Compilation(format!(
102            "XGCF invariant violation: level_offsets decrease at level {} ({} > {})",
103            level, start, end
104        )));
105    }
106    Ok((end - start) as usize)
107}
108
109fn validate_xgcf_for_gpu_upload(circuit: &Xgcf) -> Result<(u32, u32, u32)> {
110    let n = circuit.node_type.len();
111    if n == 0 {
112        return Err(XlogError::Compilation(
113            "GPU XGCF upload requires at least one node".to_string(),
114        ));
115    }
116    let node_count = checked_gpu_u32_len("GPU XGCF node count", n)?;
117    let child_offsets_len = checked_gpu_len_add_one("GPU XGCF child_offsets", n)?;
118    if circuit.child_offsets.len() != child_offsets_len {
119        return Err(XlogError::Compilation(format!(
120            "XGCF invariant violation: child_offsets len {} != num_nodes+1 ({})",
121            circuit.child_offsets.len(),
122            child_offsets_len
123        )));
124    }
125    if circuit.lit.len() != n
126        || circuit.decision_var.len() != n
127        || circuit.decision_child_false.len() != n
128        || circuit.decision_child_true.len() != n
129    {
130        return Err(XlogError::Compilation(
131            "XGCF invariant violation: per-node arrays length mismatch".to_string(),
132        ));
133    }
134
135    let edge_count = checked_gpu_u32_len("GPU XGCF edge count", circuit.child_indices.len())?;
136    let mut previous_offset = 0u32;
137    for (idx, &offset) in circuit.child_offsets.iter().enumerate() {
138        if offset < previous_offset {
139            return Err(XlogError::Compilation(format!(
140                "XGCF invariant violation: child_offsets decrease at index {} ({} > {})",
141                idx, previous_offset, offset
142            )));
143        }
144        if offset > edge_count {
145            return Err(XlogError::Compilation(format!(
146                "XGCF invariant violation: child_offsets[{}] {} exceeds child_indices len {}",
147                idx, offset, edge_count
148            )));
149        }
150        previous_offset = offset;
151    }
152    if previous_offset != edge_count {
153        return Err(XlogError::Compilation(format!(
154            "XGCF invariant violation: final child offset {} != child_indices len {}",
155            previous_offset, edge_count
156        )));
157    }
158    for (edge, &child) in circuit.child_indices.iter().enumerate() {
159        if child >= node_count {
160            return Err(XlogError::Compilation(format!(
161                "XGCF invariant violation: child_indices[{}] {} out of bounds (num_nodes={})",
162                edge, child, node_count
163            )));
164        }
165    }
166
167    for (idx, &ty) in circuit.node_type.iter().enumerate() {
168        match ty {
169            XgcfNodeType::Const0 | XgcfNodeType::Const1 => {}
170            XgcfNodeType::Lit => {
171                if circuit.lit[idx] == 0 {
172                    return Err(XlogError::Compilation(format!(
173                        "XGCF invariant violation: LIT node {} has lit=0",
174                        idx
175                    )));
176                }
177            }
178            XgcfNodeType::And | XgcfNodeType::Or => {
179                if circuit.child_offsets[idx] == circuit.child_offsets[idx + 1] {
180                    return Err(XlogError::Compilation(format!(
181                        "XGCF invariant violation: {:?} node {} has no children",
182                        ty, idx
183                    )));
184                }
185            }
186            XgcfNodeType::Decision => {
187                if circuit.decision_var[idx] == 0 {
188                    return Err(XlogError::Compilation(format!(
189                        "XGCF invariant violation: DECISION node {} has var=0",
190                        idx
191                    )));
192                }
193                if circuit.decision_child_false[idx] >= node_count {
194                    return Err(XlogError::Compilation(format!(
195                        "XGCF invariant violation: DECISION node {} false child {} out of bounds",
196                        idx, circuit.decision_child_false[idx]
197                    )));
198                }
199                if circuit.decision_child_true[idx] >= node_count {
200                    return Err(XlogError::Compilation(format!(
201                        "XGCF invariant violation: DECISION node {} true child {} out of bounds",
202                        idx, circuit.decision_child_true[idx]
203                    )));
204                }
205            }
206        }
207    }
208
209    if circuit.level_offsets.is_empty() || circuit.level_offsets[0] != 0 {
210        return Err(XlogError::Compilation(
211            "XGCF invariant violation: level_offsets must start at 0".to_string(),
212        ));
213    }
214    let level_nodes_len =
215        checked_gpu_u32_len("GPU XGCF level_nodes len", circuit.level_nodes.len())?;
216    let mut previous_level_offset = 0u32;
217    for (idx, &offset) in circuit.level_offsets.iter().enumerate() {
218        if offset < previous_level_offset {
219            return Err(XlogError::Compilation(format!(
220                "XGCF invariant violation: level_offsets decrease at index {} ({} > {})",
221                idx, previous_level_offset, offset
222            )));
223        }
224        if offset > level_nodes_len {
225            return Err(XlogError::Compilation(format!(
226                "XGCF invariant violation: level_offsets[{}] {} exceeds level_nodes len {}",
227                idx, offset, level_nodes_len
228            )));
229        }
230        previous_level_offset = offset;
231    }
232    if previous_level_offset != level_nodes_len {
233        return Err(XlogError::Compilation(format!(
234            "XGCF invariant violation: level_offsets last {} != level_nodes.len {}",
235            previous_level_offset, level_nodes_len
236        )));
237    }
238    for (idx, &node) in circuit.level_nodes.iter().enumerate() {
239        if node >= node_count {
240            return Err(XlogError::Compilation(format!(
241                "XGCF invariant violation: level_nodes[{}] {} out of bounds (num_nodes={})",
242                idx, node, node_count
243            )));
244        }
245    }
246    let num_levels_usize = circuit.level_offsets.len() - 1;
247    let num_levels = checked_gpu_u32_len("GPU XGCF level count", num_levels_usize)?;
248    if num_levels == 0 {
249        return Err(XlogError::Compilation(
250            "GPU XGCF upload requires at least one level".to_string(),
251        ));
252    }
253
254    if circuit.roots.len() != 1 {
255        return Err(XlogError::Compilation(format!(
256            "GPU XGCF eval expects exactly 1 root, got {}",
257            circuit.roots.len()
258        )));
259    }
260    if circuit.roots[0] >= node_count {
261        return Err(XlogError::Compilation(format!(
262            "XGCF invariant violation: root {} out of bounds (num_nodes={})",
263            circuit.roots[0], node_count
264        )));
265    }
266
267    Ok((node_count, edge_count, num_levels))
268}
269
270impl GpuXgcf {
271    pub fn from_device(
272        builder: GpuCircuitBuilder,
273        layout: GpuCircuitLayout,
274        provider: &CudaKernelProvider,
275    ) -> Result<GpuXgcf> {
276        if layout.num_nodes == 0 {
277            return Err(XlogError::Compilation(
278                "GpuXgcf::from_device requires num_nodes > 0".to_string(),
279            ));
280        }
281        if layout.root >= layout.num_nodes {
282            return Err(XlogError::Compilation(format!(
283                "GpuXgcf::from_device: root {} out of bounds (num_nodes={})",
284                layout.root, layout.num_nodes
285            )));
286        }
287        if layout.num_levels == 0 {
288            return Err(XlogError::Compilation(
289                "GpuXgcf::from_device requires num_levels > 0".to_string(),
290            ));
291        }
292
293        let num_nodes = layout.num_nodes as usize;
294        let num_edges = layout.num_edges as usize;
295        let node_cap = builder.node_type.len();
296        if num_nodes == 0 || num_nodes > node_cap {
297            return Err(XlogError::Compilation(
298                "GpuXgcf::from_device: num_nodes out of bounds".to_string(),
299            ));
300        }
301        let child_offsets_len =
302            checked_gpu_len_add_one("GpuXgcf::from_device child_offsets", node_cap)?;
303        if builder.child_offsets.len() != child_offsets_len
304            || builder.lit.len() != node_cap
305            || builder.decision_var.len() != node_cap
306            || builder.decision_child_false.len() != node_cap
307            || builder.decision_child_true.len() != node_cap
308        {
309            return Err(XlogError::Compilation(
310                "GpuXgcf::from_device: circuit buffer length mismatch".to_string(),
311            ));
312        }
313        if num_edges > builder.child_indices.len() {
314            return Err(XlogError::Compilation(
315                "GpuXgcf::from_device: num_edges out of bounds".to_string(),
316            ));
317        }
318
319        let num_levels = layout.num_levels as usize;
320        let level_offsets_len =
321            checked_gpu_len_add_one("GpuXgcf::from_device level_offsets", num_levels)?;
322        if layout.level_offsets.len() != level_offsets_len {
323            return Err(XlogError::Compilation(format!(
324                "GpuXgcf::from_device: level_offsets len {} != num_levels+1 ({})",
325                layout.level_offsets.len(),
326                level_offsets_len
327            )));
328        }
329        if layout.level_nodes.len() < num_nodes {
330            return Err(XlogError::Compilation(format!(
331                "GpuXgcf::from_device: level_nodes len {} < num_nodes ({})",
332                layout.level_nodes.len(),
333                num_nodes
334            )));
335        }
336
337        let memory = provider.memory();
338
339        let weights_len = (layout.max_var as usize) + 1;
340        let var_log_true = memory.alloc::<f64>(weights_len)?;
341        let var_log_false = memory.alloc::<f64>(weights_len)?;
342        let values = memory.alloc::<f64>(num_nodes)?;
343        let adj = memory.alloc::<f64>(num_nodes)?;
344        let grad_true = memory.alloc::<f64>(weights_len)?;
345        let grad_false = memory.alloc::<f64>(weights_len)?;
346
347        let meta_num_nodes = match layout.num_nodes_device {
348            Some(meta) => meta,
349            None => {
350                let mut meta = memory.alloc::<u32>(1)?;
351                provider
352                    .htod_launch_metadata_sync_copy_into(&[layout.num_nodes], &mut meta)
353                    .map_err(|e| {
354                        XlogError::Kernel(format!("Failed to upload num_nodes meta: {}", e))
355                    })?;
356                meta
357            }
358        };
359        let meta_num_edges = match layout.num_edges_device {
360            Some(meta) => meta,
361            None => {
362                let mut meta = memory.alloc::<u32>(1)?;
363                provider
364                    .htod_launch_metadata_sync_copy_into(&[layout.num_edges], &mut meta)
365                    .map_err(|e| {
366                        XlogError::Kernel(format!("Failed to upload num_edges meta: {}", e))
367                    })?;
368                meta
369            }
370        };
371
372        Ok(Self {
373            node_type: builder.node_type,
374            child_offsets: builder.child_offsets,
375            child_indices: builder.child_indices,
376            lit: builder.lit,
377            decision_var: builder.decision_var,
378            decision_child_false: builder.decision_child_false,
379            decision_child_true: builder.decision_child_true,
380            level_nodes: layout.level_nodes,
381            level_offsets: layout.level_offsets,
382            level_offsets_host: None,
383            node_cap: layout.num_nodes,
384            edge_cap: layout.num_edges,
385            num_levels: layout.num_levels,
386            root: layout.root,
387            max_var: layout.max_var,
388            meta_num_nodes,
389            meta_num_edges,
390            var_log_true,
391            var_log_false,
392            values,
393            adj,
394            grad_true,
395            grad_false,
396            free_var_mask: None,
397        })
398    }
399
400    /// GPU-native smoothing pass for random variables.
401    ///
402    /// Returns a new device-resident circuit that is smooth w.r.t. `random_var_list`.
403    /// This method performs no device->host data-plane transfers and traps on capacity overflow.
404    pub fn smooth_random_vars_device(
405        &self,
406        provider: &CudaKernelProvider,
407        random_var_list: &TrackedCudaSlice<u32>,
408        random_var_count: u32,
409        smooth_node_cap: u32,
410        smooth_edge_cap: u32,
411    ) -> Result<GpuXgcf> {
412        if smooth_node_cap == 0 || smooth_edge_cap == 0 {
413            return Err(XlogError::Compilation(
414                "GPU smoothing requires non-zero node/edge caps".to_string(),
415            ));
416        }
417
418        let num_nodes = self.node_cap;
419        if num_nodes == 0 {
420            return Err(XlogError::Compilation(
421                "GPU smoothing: num_nodes must be > 0".to_string(),
422            ));
423        }
424        if self.child_offsets.len() < (num_nodes as usize + 1) {
425            return Err(XlogError::Compilation(
426                "GPU smoothing: child_offsets len mismatch".to_string(),
427            ));
428        }
429        let num_edges = self.edge_cap;
430        if num_edges == 0 {
431            return Err(XlogError::Compilation(
432                "GPU smoothing: num_edges must be > 0".to_string(),
433            ));
434        }
435
436        let list_len = u32::try_from(random_var_list.len()).map_err(|_| {
437            XlogError::Compilation("GPU smoothing: random var list len exceeds u32".to_string())
438        })?;
439        let num_random_vars = random_var_count;
440        if num_random_vars > list_len {
441            return Err(XlogError::Compilation(format!(
442                "GPU smoothing: random var count {} exceeds list len {}",
443                num_random_vars, list_len
444            )));
445        }
446
447        let base_node = 2u32.checked_add(num_random_vars).ok_or_else(|| {
448            XlogError::Compilation("GPU smoothing: base node overflow".to_string())
449        })?;
450        let base_nodes = (base_node as u64)
451            .checked_add(num_nodes as u64)
452            .ok_or_else(|| {
453                XlogError::Compilation("GPU smoothing: base node overflow".to_string())
454            })?;
455        if base_nodes > smooth_node_cap as u64 {
456            return Err(XlogError::Compilation(format!(
457                "GPU smoothing: base nodes {} exceed smooth_node_cap {}",
458                base_nodes, smooth_node_cap
459            )));
460        }
461
462        let words_per_support = num_random_vars.div_ceil(32).max(1);
463
464        let support_len = (num_nodes as u64)
465            .checked_mul(words_per_support as u64)
466            .and_then(|v| usize::try_from(v).ok())
467            .ok_or_else(|| {
468                XlogError::Compilation("GPU smoothing: support size overflow".to_string())
469            })?;
470
471        let dec_entries = (num_nodes as u64)
472            .checked_mul(2)
473            .and_then(|v| usize::try_from(v).ok())
474            .ok_or_else(|| {
475                XlogError::Compilation("GPU smoothing: decision array overflow".to_string())
476            })?;
477        let dec_entries_u32 = u32::try_from(dec_entries).map_err(|_| {
478            XlogError::Compilation("GPU smoothing: decision entries exceed u32".to_string())
479        })?;
480
481        let device = provider.device().inner();
482        let memory = provider.memory();
483        let block_size: u32 = 256;
484
485        let map_len = (self.max_var as usize)
486            .checked_add(1)
487            .ok_or_else(|| XlogError::Compilation("GPU smoothing: max_var overflow".to_string()))?;
488        let map_len_u32 = u32::try_from(map_len).map_err(|_| {
489            XlogError::Compilation("GPU smoothing: random map len exceeds u32".to_string())
490        })?;
491        let mut d_random_map = memory.alloc::<u32>(map_len)?;
492        if map_len > 0 {
493            let fill_const = device
494                .get_func(FILTER_MODULE, filter_kernels::FILL_U32_CONST)
495                .ok_or_else(|| XlogError::Kernel("fill_u32_const kernel not found".to_string()))?;
496            let grid = map_len_u32.div_ceil(block_size);
497            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
498            unsafe {
499                fill_const.clone().launch(
500                    LaunchConfig {
501                        grid_dim: (grid, 1, 1),
502                        block_dim: (block_size, 1, 1),
503                        shared_mem_bytes: 0,
504                    },
505                    (&mut d_random_map, map_len_u32, u32::MAX),
506                )
507            }
508            .map_err(|e| XlogError::Kernel(format!("fill_u32_const failed: {}", e)))?;
509        }
510        if num_random_vars > 0 {
511            let map_kernel = device
512                .get_func(FILTER_MODULE, filter_kernels::RANDOM_VAR_TO_BIT_FROM_LIST)
513                .ok_or_else(|| {
514                    XlogError::Kernel("random_var_to_bit_from_list kernel not found".to_string())
515                })?;
516            let grid = num_random_vars.div_ceil(block_size);
517            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
518            unsafe {
519                map_kernel.clone().launch(
520                    LaunchConfig {
521                        grid_dim: (grid, 1, 1),
522                        block_dim: (block_size, 1, 1),
523                        shared_mem_bytes: 0,
524                    },
525                    (
526                        random_var_list,
527                        num_random_vars,
528                        map_len_u32,
529                        &mut d_random_map,
530                    ),
531                )
532            }
533            .map_err(|e| XlogError::Kernel(format!("random_var_to_bit_from_list failed: {}", e)))?;
534        }
535
536        let mut support = memory.alloc::<u32>(support_len)?;
537        device
538            .memset_zeros(&mut support)
539            .map_err(|e| XlogError::Kernel(format!("Failed to zero support: {}", e)))?;
540
541        let support_kernel = device
542            .get_func(D4_MODULE, d4_kernels::D4_SUPPORT_LEVEL)
543            .ok_or_else(|| XlogError::Kernel("d4_support_level kernel not found".to_string()))?;
544
545        let num_levels = self.num_levels as usize;
546        let random_map_len = map_len_u32;
547        for level in 0..num_levels {
548            let num_level_nodes: usize = match self.level_offsets_host.as_ref() {
549                Some(off) => checked_host_level_width(off, level)?,
550                None => self.level_nodes.len(),
551            };
552            if num_level_nodes == 0 {
553                continue;
554            }
555            let num_blocks =
556                checked_gpu_launch_blocks("d4_support_level", num_level_nodes, block_size)?;
557            let config = LaunchConfig {
558                grid_dim: (num_blocks, 1, 1),
559                block_dim: (block_size, 1, 1),
560                shared_mem_bytes: 0,
561            };
562            let level_u32 = level as u32;
563            let mut params: Vec<*mut c_void> = vec![
564                (&self.node_type).as_kernel_param(),
565                (&self.child_offsets).as_kernel_param(),
566                (&self.child_indices).as_kernel_param(),
567                (&self.lit).as_kernel_param(),
568                (&self.decision_var).as_kernel_param(),
569                (&self.decision_child_false).as_kernel_param(),
570                (&self.decision_child_true).as_kernel_param(),
571                (&self.level_nodes).as_kernel_param(),
572                (&self.level_offsets).as_kernel_param(),
573                level_u32.as_kernel_param(),
574                (&d_random_map).as_kernel_param(),
575                random_map_len.as_kernel_param(),
576                words_per_support.as_kernel_param(),
577                (&support).as_kernel_param(),
578            ];
579            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
580            unsafe { support_kernel.clone().launch(config, &mut params) }
581                .map_err(|e| XlogError::Kernel(format!("d4_support_level failed: {}", e)))?;
582        }
583
584        if num_random_vars > 0 {
585            let root_kernel = device
586                .get_func(D4_MODULE, d4_kernels::D4_SUPPORT_SET_ROOT_BITS)
587                .ok_or_else(|| {
588                    XlogError::Kernel("d4_support_set_root_bits kernel not found".to_string())
589                })?;
590            let num_words = num_random_vars.div_ceil(32);
591            let grid = num_words.div_ceil(block_size);
592            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
593            unsafe {
594                root_kernel.clone().launch(
595                    LaunchConfig {
596                        grid_dim: (grid, 1, 1),
597                        block_dim: (block_size, 1, 1),
598                        shared_mem_bytes: 0,
599                    },
600                    (self.root, num_random_vars, words_per_support, &mut support),
601                )
602            }
603            .map_err(|e| XlogError::Kernel(format!("d4_support_set_root_bits failed: {}", e)))?;
604        }
605
606        let mut wrap_prefix_or = memory.alloc::<u32>(num_edges as usize)?;
607        let mut wrap_missing_or = memory.alloc::<u32>(num_edges as usize)?;
608        let mut wrap_prefix_dec = memory.alloc::<u32>(dec_entries)?;
609        let mut wrap_missing_dec = memory.alloc::<u32>(dec_entries)?;
610
611        device
612            .memset_zeros(&mut wrap_prefix_or)
613            .map_err(|e| XlogError::Kernel(format!("Failed to zero wrap_prefix_or: {}", e)))?;
614        device
615            .memset_zeros(&mut wrap_missing_or)
616            .map_err(|e| XlogError::Kernel(format!("Failed to zero wrap_missing_or: {}", e)))?;
617        device
618            .memset_zeros(&mut wrap_prefix_dec)
619            .map_err(|e| XlogError::Kernel(format!("Failed to zero wrap_prefix_dec: {}", e)))?;
620        device
621            .memset_zeros(&mut wrap_missing_dec)
622            .map_err(|e| XlogError::Kernel(format!("Failed to zero wrap_missing_dec: {}", e)))?;
623
624        let mut out_edge_counts = memory.alloc::<u32>(smooth_node_cap as usize)?;
625        device
626            .memset_zeros(&mut out_edge_counts)
627            .map_err(|e| XlogError::Kernel(format!("Failed to zero edge_counts: {}", e)))?;
628
629        let count_kernel = device
630            .get_func(D4_MODULE, d4_kernels::D4_SMOOTH_COUNT)
631            .ok_or_else(|| XlogError::Kernel("d4_smooth_count kernel not found".to_string()))?;
632        let num_blocks = num_nodes.div_ceil(block_size);
633        let mut params: Vec<*mut c_void> = vec![
634            (&self.node_type).as_kernel_param(),
635            (&self.child_offsets).as_kernel_param(),
636            (&self.child_indices).as_kernel_param(),
637            (&self.decision_var).as_kernel_param(),
638            (&self.decision_child_false).as_kernel_param(),
639            (&self.decision_child_true).as_kernel_param(),
640            (&self.meta_num_nodes).as_kernel_param(),
641            (&support).as_kernel_param(),
642            words_per_support.as_kernel_param(),
643            (&d_random_map).as_kernel_param(),
644            random_map_len.as_kernel_param(),
645            (&wrap_prefix_or).as_kernel_param(),
646            (&wrap_missing_or).as_kernel_param(),
647            (&wrap_prefix_dec).as_kernel_param(),
648            (&wrap_missing_dec).as_kernel_param(),
649            (&out_edge_counts).as_kernel_param(),
650            base_node.as_kernel_param(),
651            smooth_node_cap.as_kernel_param(),
652        ];
653        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
654        unsafe {
655            count_kernel.clone().launch(
656                LaunchConfig {
657                    grid_dim: (num_blocks, 1, 1),
658                    block_dim: (block_size, 1, 1),
659                    shared_mem_bytes: 0,
660                },
661                &mut params,
662            )
663        }
664        .map_err(|e| XlogError::Kernel(format!("d4_smooth_count failed: {}", e)))?;
665
666        exclusive_scan_u32_inplace(provider, &mut wrap_prefix_or, num_edges)?;
667        exclusive_scan_u32_inplace(provider, &mut wrap_prefix_dec, dec_entries_u32)?;
668
669        let mut wrap_counts = memory.alloc::<u32>(3)?;
670        device
671            .memset_zeros(&mut wrap_counts)
672            .map_err(|e| XlogError::Kernel(format!("Failed to zero wrap_counts: {}", e)))?;
673
674        let counts_kernel = device
675            .get_func(D4_MODULE, d4_kernels::D4_SMOOTH_WRAPPER_COUNTS)
676            .ok_or_else(|| {
677                XlogError::Kernel("d4_smooth_wrapper_counts kernel not found".to_string())
678            })?;
679        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
680        unsafe {
681            counts_kernel.clone().launch(
682                LaunchConfig {
683                    grid_dim: (1, 1, 1),
684                    block_dim: (1, 1, 1),
685                    shared_mem_bytes: 0,
686                },
687                (
688                    &wrap_prefix_or,
689                    &wrap_missing_or,
690                    num_edges,
691                    &wrap_prefix_dec,
692                    &wrap_missing_dec,
693                    dec_entries_u32,
694                    base_node,
695                    &self.meta_num_nodes,
696                    u32::MAX,
697                    &mut wrap_counts,
698                ),
699            )
700        }
701        .map_err(|e| XlogError::Kernel(format!("d4_smooth_wrapper_counts failed: {}", e)))?;
702
703        let wrap_or_kernel = device
704            .get_func(D4_MODULE, d4_kernels::D4_SMOOTH_WRAPPER_EDGE_COUNTS_OR)
705            .ok_or_else(|| {
706                XlogError::Kernel("d4_smooth_wrapper_edge_counts_or kernel not found".to_string())
707            })?;
708        if num_edges > 0 {
709            let num_blocks = num_edges.div_ceil(block_size);
710            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
711            unsafe {
712                wrap_or_kernel.clone().launch(
713                    LaunchConfig {
714                        grid_dim: (num_blocks, 1, 1),
715                        block_dim: (block_size, 1, 1),
716                        shared_mem_bytes: 0,
717                    },
718                    (
719                        &wrap_prefix_or,
720                        &wrap_missing_or,
721                        num_edges,
722                        base_node,
723                        &self.meta_num_nodes,
724                        smooth_node_cap,
725                        &mut out_edge_counts,
726                    ),
727                )
728            }
729            .map_err(|e| {
730                XlogError::Kernel(format!("d4_smooth_wrapper_edge_counts_or failed: {}", e))
731            })?;
732        }
733
734        let wrap_dec_kernel = device
735            .get_func(D4_MODULE, d4_kernels::D4_SMOOTH_WRAPPER_EDGE_COUNTS_DEC)
736            .ok_or_else(|| {
737                XlogError::Kernel("d4_smooth_wrapper_edge_counts_dec kernel not found".to_string())
738            })?;
739        if dec_entries > 0 {
740            let num_blocks = dec_entries_u32.div_ceil(block_size);
741            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
742            unsafe {
743                wrap_dec_kernel.clone().launch(
744                    LaunchConfig {
745                        grid_dim: (num_blocks, 1, 1),
746                        block_dim: (block_size, 1, 1),
747                        shared_mem_bytes: 0,
748                    },
749                    (
750                        &wrap_prefix_dec,
751                        &wrap_missing_dec,
752                        dec_entries_u32,
753                        base_node,
754                        &self.meta_num_nodes,
755                        &wrap_counts,
756                        smooth_node_cap,
757                        &mut out_edge_counts,
758                    ),
759                )
760            }
761            .map_err(|e| {
762                XlogError::Kernel(format!("d4_smooth_wrapper_edge_counts_dec failed: {}", e))
763            })?;
764        }
765
766        let mut out_child_offsets = memory.alloc::<u32>((smooth_node_cap as usize) + 1)?;
767        device
768            .memset_zeros(&mut out_child_offsets)
769            .map_err(|e| XlogError::Kernel(format!("Failed to zero child_offsets: {}", e)))?;
770        if smooth_node_cap > 0 {
771            device
772                .dtod_copy(
773                    &out_edge_counts,
774                    &mut out_child_offsets.slice_mut(0..smooth_node_cap as usize),
775                )
776                .map_err(|e| XlogError::Kernel(format!("Failed to copy edge_counts: {}", e)))?;
777        }
778        let child_scan_len = smooth_node_cap.checked_add(1).ok_or_else(|| {
779            XlogError::Compilation("GPU smoothing: child offset scan overflow".to_string())
780        })?;
781        exclusive_scan_u32_inplace(provider, &mut out_child_offsets, child_scan_len)?;
782
783        let edge_cap_check = device
784            .get_func(D4_MODULE, d4_kernels::D4_SMOOTH_CHECK_EDGE_CAP)
785            .ok_or_else(|| {
786                XlogError::Kernel("d4_smooth_check_edge_cap kernel not found".to_string())
787            })?;
788        let mut meta_num_nodes = memory.alloc::<u32>(1)?;
789        let mut meta_num_edges = memory.alloc::<u32>(1)?;
790        device
791            .memset_zeros(&mut meta_num_nodes)
792            .map_err(|e| XlogError::Kernel(format!("Failed to zero smooth num_nodes: {}", e)))?;
793        device
794            .memset_zeros(&mut meta_num_edges)
795            .map_err(|e| XlogError::Kernel(format!("Failed to zero smooth num_edges: {}", e)))?;
796        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
797        unsafe {
798            edge_cap_check.clone().launch(
799                LaunchConfig {
800                    grid_dim: (1, 1, 1),
801                    block_dim: (1, 1, 1),
802                    shared_mem_bytes: 0,
803                },
804                (
805                    &out_child_offsets,
806                    smooth_node_cap,
807                    smooth_edge_cap,
808                    &wrap_counts,
809                    &mut meta_num_nodes,
810                    &mut meta_num_edges,
811                ),
812            )
813        }
814        .map_err(|e| XlogError::Kernel(format!("d4_smooth_check_edge_cap failed: {}", e)))?;
815
816        let mut out_node_type = memory.alloc::<u8>(smooth_node_cap as usize)?;
817        let mut out_child_indices = memory.alloc::<u32>(smooth_edge_cap as usize)?;
818        let mut out_lit = memory.alloc::<i32>(smooth_node_cap as usize)?;
819        let mut out_decision_var = memory.alloc::<u32>(smooth_node_cap as usize)?;
820        let mut out_decision_child_false = memory.alloc::<u32>(smooth_node_cap as usize)?;
821        let mut out_decision_child_true = memory.alloc::<u32>(smooth_node_cap as usize)?;
822        let mut out_node_level = memory.alloc::<u32>(smooth_node_cap as usize)?;
823
824        device
825            .memset_zeros(&mut out_node_type)
826            .map_err(|e| XlogError::Kernel(format!("Failed to zero node_type: {}", e)))?;
827        device
828            .memset_zeros(&mut out_child_indices)
829            .map_err(|e| XlogError::Kernel(format!("Failed to zero child_indices: {}", e)))?;
830        device
831            .memset_zeros(&mut out_lit)
832            .map_err(|e| XlogError::Kernel(format!("Failed to zero lit: {}", e)))?;
833        device
834            .memset_zeros(&mut out_decision_var)
835            .map_err(|e| XlogError::Kernel(format!("Failed to zero decision_var: {}", e)))?;
836        device
837            .memset_zeros(&mut out_decision_child_false)
838            .map_err(|e| {
839                XlogError::Kernel(format!("Failed to zero decision_child_false: {}", e))
840            })?;
841        device
842            .memset_zeros(&mut out_decision_child_true)
843            .map_err(|e| XlogError::Kernel(format!("Failed to zero decision_child_true: {}", e)))?;
844        device
845            .memset_zeros(&mut out_node_level)
846            .map_err(|e| XlogError::Kernel(format!("Failed to zero node_level: {}", e)))?;
847
848        let init_kernel = device
849            .get_func(D4_MODULE, d4_kernels::D4_SMOOTH_INIT_NODES)
850            .ok_or_else(|| {
851                XlogError::Kernel("d4_smooth_init_nodes kernel not found".to_string())
852            })?;
853        let init_blocks = checked_gpu_launch_blocks(
854            "d4_smooth_init_nodes",
855            num_random_vars.max(1) as usize,
856            block_size,
857        )?;
858        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
859        unsafe {
860            init_kernel.clone().launch(
861                LaunchConfig {
862                    grid_dim: (init_blocks, 1, 1),
863                    block_dim: (block_size, 1, 1),
864                    shared_mem_bytes: 0,
865                },
866                (
867                    random_var_list,
868                    num_random_vars,
869                    smooth_node_cap,
870                    &mut out_node_type,
871                    &mut out_lit,
872                    &mut out_decision_var,
873                    &mut out_decision_child_false,
874                    &mut out_decision_child_true,
875                    &mut out_node_level,
876                ),
877            )
878        }
879        .map_err(|e| XlogError::Kernel(format!("d4_smooth_init_nodes failed: {}", e)))?;
880
881        let num_levels_out = self
882            .num_levels
883            .checked_mul(2)
884            .and_then(|levels| levels.checked_add(4))
885            .ok_or_else(|| {
886                XlogError::Compilation("GPU smoothing output level count overflow".to_string())
887            })?;
888        let num_levels_out_usize = num_levels_out as usize;
889        let level_offsets_len =
890            checked_gpu_len_add_one("GPU smoothing level offsets", num_levels_out_usize)?;
891
892        let emit_kernel = device
893            .get_func(D4_MODULE, d4_kernels::D4_SMOOTH_EMIT_LEVEL)
894            .ok_or_else(|| {
895                XlogError::Kernel("d4_smooth_emit_level kernel not found".to_string())
896            })?;
897        for level in 0..num_levels {
898            let num_level_nodes: usize = match self.level_offsets_host.as_ref() {
899                Some(off) => checked_host_level_width(off, level)?,
900                None => self.level_nodes.len(),
901            };
902            if num_level_nodes == 0 {
903                continue;
904            }
905            let num_blocks =
906                checked_gpu_launch_blocks("xgcf_smooth_forward", num_level_nodes, block_size)?;
907            let level_u32 = level as u32;
908            let mut params: Vec<*mut c_void> = vec![
909                (&self.node_type).as_kernel_param(),
910                (&self.child_offsets).as_kernel_param(),
911                (&self.child_indices).as_kernel_param(),
912                (&self.lit).as_kernel_param(),
913                (&self.decision_var).as_kernel_param(),
914                (&self.decision_child_false).as_kernel_param(),
915                (&self.decision_child_true).as_kernel_param(),
916                (&self.level_nodes).as_kernel_param(),
917                (&self.level_offsets).as_kernel_param(),
918                level_u32.as_kernel_param(),
919                (&support).as_kernel_param(),
920                words_per_support.as_kernel_param(),
921                (&wrap_prefix_or).as_kernel_param(),
922                (&wrap_missing_or).as_kernel_param(),
923                (&wrap_prefix_dec).as_kernel_param(),
924                (&wrap_missing_dec).as_kernel_param(),
925                base_node.as_kernel_param(),
926                (&self.meta_num_nodes).as_kernel_param(),
927                (&wrap_counts).as_kernel_param(),
928                num_random_vars.as_kernel_param(),
929                num_levels_out.as_kernel_param(),
930                (&out_node_type).as_kernel_param(),
931                (&out_child_offsets).as_kernel_param(),
932                (&out_child_indices).as_kernel_param(),
933                (&out_lit).as_kernel_param(),
934                (&out_decision_var).as_kernel_param(),
935                (&out_decision_child_false).as_kernel_param(),
936                (&out_decision_child_true).as_kernel_param(),
937                (&out_node_level).as_kernel_param(),
938            ];
939            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
940            unsafe {
941                emit_kernel.clone().launch(
942                    LaunchConfig {
943                        grid_dim: (num_blocks, 1, 1),
944                        block_dim: (block_size, 1, 1),
945                        shared_mem_bytes: 0,
946                    },
947                    &mut params,
948                )
949            }
950            .map_err(|e| XlogError::Kernel(format!("d4_smooth_emit_level failed: {}", e)))?;
951        }
952
953        let mut level_counts = memory.alloc::<u32>(num_levels_out_usize)?;
954        let mut level_offsets = memory.alloc::<u32>(level_offsets_len)?;
955        let mut level_cursors = memory.alloc::<u32>(num_levels_out_usize)?;
956        let mut level_nodes = memory.alloc::<u32>(smooth_node_cap as usize)?;
957
958        device
959            .memset_zeros(&mut level_counts)
960            .map_err(|e| XlogError::Kernel(format!("Failed to zero level_counts: {}", e)))?;
961        device
962            .memset_zeros(&mut level_offsets)
963            .map_err(|e| XlogError::Kernel(format!("Failed to zero level_offsets: {}", e)))?;
964        device
965            .memset_zeros(&mut level_cursors)
966            .map_err(|e| XlogError::Kernel(format!("Failed to zero level_cursors: {}", e)))?;
967        device
968            .memset_zeros(&mut level_nodes)
969            .map_err(|e| XlogError::Kernel(format!("Failed to zero level_nodes: {}", e)))?;
970
971        let mut compile_needed = memory.alloc::<u32>(1)?;
972        provider
973            .htod_launch_metadata_sync_copy_into(&[1u32], &mut compile_needed)
974            .map_err(|e| XlogError::Kernel(format!("Failed to upload compile_needed: {}", e)))?;
975
976        let levelize_counts = device
977            .get_func(D4_MODULE, d4_kernels::D4_LEVELIZE_COUNTS)
978            .ok_or_else(|| XlogError::Kernel("d4_levelize_counts kernel not found".to_string()))?;
979        let num_blocks =
980            checked_gpu_launch_blocks("d4_smooth_levelize", smooth_node_cap as usize, block_size)?;
981        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
982        unsafe {
983            levelize_counts.clone().launch(
984                LaunchConfig {
985                    grid_dim: (num_blocks, 1, 1),
986                    block_dim: (block_size, 1, 1),
987                    shared_mem_bytes: 0,
988                },
989                (
990                    &compile_needed,
991                    &out_node_level,
992                    &meta_num_nodes,
993                    num_levels_out,
994                    &mut level_counts,
995                ),
996            )
997        }
998        .map_err(|e| XlogError::Kernel(format!("d4_levelize_counts failed: {}", e)))?;
999
1000        device
1001            .dtod_copy(
1002                &level_counts,
1003                &mut level_offsets.slice_mut(0..num_levels_out_usize),
1004            )
1005            .map_err(|e| XlogError::Kernel(format!("Failed to copy level_counts: {}", e)))?;
1006        let level_scan_len = num_levels_out.checked_add(1).ok_or_else(|| {
1007            XlogError::Compilation("GPU smoothing: level offset scan overflow".to_string())
1008        })?;
1009        exclusive_scan_u32_inplace(provider, &mut level_offsets, level_scan_len)?;
1010
1011        let levelize_emit = device
1012            .get_func(D4_MODULE, d4_kernels::D4_LEVELIZE_EMIT)
1013            .ok_or_else(|| XlogError::Kernel("d4_levelize_emit kernel not found".to_string()))?;
1014        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1015        unsafe {
1016            levelize_emit.clone().launch(
1017                LaunchConfig {
1018                    grid_dim: (num_blocks, 1, 1),
1019                    block_dim: (block_size, 1, 1),
1020                    shared_mem_bytes: 0,
1021                },
1022                (
1023                    &compile_needed,
1024                    &out_node_level,
1025                    &meta_num_nodes,
1026                    num_levels_out,
1027                    &level_offsets,
1028                    &mut level_cursors,
1029                    &mut level_nodes,
1030                ),
1031            )
1032        }
1033        .map_err(|e| XlogError::Kernel(format!("d4_levelize_emit failed: {}", e)))?;
1034
1035        // No device synchronize: result is device buffers used by subsequent GPU ops.
1036        let builder = GpuCircuitBuilder {
1037            node_type: out_node_type,
1038            child_offsets: out_child_offsets,
1039            child_indices: out_child_indices,
1040            lit: out_lit,
1041            decision_var: out_decision_var,
1042            decision_child_false: out_decision_child_false,
1043            decision_child_true: out_decision_child_true,
1044        };
1045        let layout = GpuCircuitLayout {
1046            num_nodes: smooth_node_cap,
1047            num_edges: smooth_edge_cap,
1048            num_levels: num_levels_out,
1049            level_offsets,
1050            level_nodes,
1051            root: base_node + self.root,
1052            max_var: self.max_var,
1053            num_nodes_device: Some(meta_num_nodes),
1054            num_edges_device: Some(meta_num_edges),
1055        };
1056
1057        GpuXgcf::from_device(builder, layout, provider)
1058    }
1059
1060    pub fn upload(provider: &CudaKernelProvider, circuit: &Xgcf) -> Result<Self> {
1061        let (node_cap, edge_cap, num_levels) = validate_xgcf_for_gpu_upload(circuit)?;
1062
1063        let memory = provider.memory().clone();
1064
1065        let n = circuit.node_type.len();
1066        let mut host_node_type: Vec<u8> = Vec::with_capacity(n);
1067        for &ty in &circuit.node_type {
1068            host_node_type.push(ty as u8);
1069        }
1070
1071        let mut max_var: u32 = 0;
1072        for (&ty, &lit) in circuit.node_type.iter().zip(circuit.lit.iter()) {
1073            if ty == XgcfNodeType::Lit && lit != 0 {
1074                max_var = max_var.max(lit.unsigned_abs());
1075            }
1076        }
1077        for &var in &circuit.decision_var {
1078            max_var = max_var.max(var);
1079        }
1080
1081        let mut d_node_type = memory.alloc::<u8>(n)?;
1082        provider
1083            .htod_sync_copy_into_tracked(&host_node_type, &mut d_node_type)
1084            .map_err(|e| XlogError::Kernel(format!("Failed to upload circuit node_type: {}", e)))?;
1085
1086        let mut d_child_offsets = memory.alloc::<u32>(circuit.child_offsets.len())?;
1087        provider
1088            .htod_sync_copy_into_tracked(&circuit.child_offsets, &mut d_child_offsets)
1089            .map_err(|e| {
1090                XlogError::Kernel(format!("Failed to upload circuit child_offsets: {}", e))
1091            })?;
1092
1093        let mut d_child_indices = memory.alloc::<u32>(circuit.child_indices.len())?;
1094        provider
1095            .htod_sync_copy_into_tracked(&circuit.child_indices, &mut d_child_indices)
1096            .map_err(|e| {
1097                XlogError::Kernel(format!("Failed to upload circuit child_indices: {}", e))
1098            })?;
1099
1100        let mut d_lit = memory.alloc::<i32>(circuit.lit.len())?;
1101        provider
1102            .htod_sync_copy_into_tracked(&circuit.lit, &mut d_lit)
1103            .map_err(|e| XlogError::Kernel(format!("Failed to upload circuit lit: {}", e)))?;
1104
1105        let mut d_decision_var = memory.alloc::<u32>(circuit.decision_var.len())?;
1106        provider
1107            .htod_sync_copy_into_tracked(&circuit.decision_var, &mut d_decision_var)
1108            .map_err(|e| {
1109                XlogError::Kernel(format!("Failed to upload circuit decision_var: {}", e))
1110            })?;
1111
1112        let mut d_decision_child_false = memory.alloc::<u32>(circuit.decision_child_false.len())?;
1113        provider
1114            .htod_sync_copy_into_tracked(&circuit.decision_child_false, &mut d_decision_child_false)
1115            .map_err(|e| {
1116                XlogError::Kernel(format!(
1117                    "Failed to upload circuit decision_child_false: {}",
1118                    e
1119                ))
1120            })?;
1121
1122        let mut d_decision_child_true = memory.alloc::<u32>(circuit.decision_child_true.len())?;
1123        provider
1124            .htod_sync_copy_into_tracked(&circuit.decision_child_true, &mut d_decision_child_true)
1125            .map_err(|e| {
1126                XlogError::Kernel(format!(
1127                    "Failed to upload circuit decision_child_true: {}",
1128                    e
1129                ))
1130            })?;
1131
1132        let mut d_level_nodes = memory.alloc::<u32>(circuit.level_nodes.len())?;
1133        provider
1134            .htod_sync_copy_into_tracked(&circuit.level_nodes, &mut d_level_nodes)
1135            .map_err(|e| {
1136                XlogError::Kernel(format!("Failed to upload circuit level_nodes: {}", e))
1137            })?;
1138
1139        let mut d_level_offsets = memory.alloc::<u32>(circuit.level_offsets.len())?;
1140        provider
1141            .htod_sync_copy_into_tracked(&circuit.level_offsets, &mut d_level_offsets)
1142            .map_err(|e| {
1143                XlogError::Kernel(format!("Failed to upload circuit level_offsets: {}", e))
1144            })?;
1145
1146        let weights_len = (max_var as usize) + 1;
1147        let var_log_true = memory.alloc::<f64>(weights_len)?;
1148        let var_log_false = memory.alloc::<f64>(weights_len)?;
1149        let values = memory.alloc::<f64>(n)?;
1150        let adj = memory.alloc::<f64>(n)?;
1151        let grad_true = memory.alloc::<f64>(weights_len)?;
1152        let grad_false = memory.alloc::<f64>(weights_len)?;
1153        let mut meta_num_nodes = memory.alloc::<u32>(1)?;
1154        provider
1155            .htod_launch_metadata_sync_copy_into(&[node_cap], &mut meta_num_nodes)
1156            .map_err(|e| XlogError::Kernel(format!("Failed to upload num_nodes meta: {}", e)))?;
1157        let mut meta_num_edges = memory.alloc::<u32>(1)?;
1158        provider
1159            .htod_launch_metadata_sync_copy_into(&[edge_cap], &mut meta_num_edges)
1160            .map_err(|e| XlogError::Kernel(format!("Failed to upload num_edges meta: {}", e)))?;
1161
1162        Ok(Self {
1163            node_type: d_node_type,
1164            child_offsets: d_child_offsets,
1165            child_indices: d_child_indices,
1166            lit: d_lit,
1167            decision_var: d_decision_var,
1168            decision_child_false: d_decision_child_false,
1169            decision_child_true: d_decision_child_true,
1170            level_nodes: d_level_nodes,
1171            level_offsets: d_level_offsets,
1172            level_offsets_host: Some(circuit.level_offsets.clone()),
1173            node_cap,
1174            edge_cap,
1175            num_levels,
1176            root: circuit.roots[0],
1177            max_var,
1178            meta_num_nodes,
1179            meta_num_edges,
1180            var_log_true,
1181            var_log_false,
1182            values,
1183            adj,
1184            grad_true,
1185            grad_false,
1186            free_var_mask: None,
1187        })
1188    }
1189
1190    pub fn max_var(&self) -> u32 {
1191        self.max_var
1192    }
1193
1194    /// Root node id of the circuit (XGCF requires exactly one root for evaluation/verification).
1195    pub fn root(&self) -> u32 {
1196        self.root
1197    }
1198
1199    /// Capacity (upper bound) for XGCF nodes in the circuit buffers.
1200    pub fn num_nodes(&self) -> usize {
1201        self.node_cap as usize
1202    }
1203
1204    /// Capacity (upper bound) for XGCF edges in the circuit buffers.
1205    pub fn num_edges(&self) -> usize {
1206        self.edge_cap as usize
1207    }
1208
1209    /// Number of topological levels in the circuit.
1210    pub fn num_levels(&self) -> u32 {
1211        self.num_levels
1212    }
1213
1214    /// Device-resident actual node count (len = 1).
1215    pub fn num_nodes_device(&self) -> &TrackedCudaSlice<u32> {
1216        &self.meta_num_nodes
1217    }
1218
1219    /// Device-resident actual edge count (len = 1).
1220    pub fn num_edges_device(&self) -> &TrackedCudaSlice<u32> {
1221        &self.meta_num_edges
1222    }
1223
1224    /// Device-resident level -> node index mapping (len = num_nodes).
1225    pub fn level_nodes(&self) -> &TrackedCudaSlice<u32> {
1226        &self.level_nodes
1227    }
1228
1229    /// Device-resident offsets for each level (len = num_levels + 1).
1230    pub fn level_offsets(&self) -> &TrackedCudaSlice<u32> {
1231        &self.level_offsets
1232    }
1233
1234    /// Device-resident node type tags (see `XgcfNodeType`).
1235    pub fn node_type(&self) -> &TrackedCudaSlice<u8> {
1236        &self.node_type
1237    }
1238
1239    /// Device-resident CSR child offsets for AND/OR nodes (len = num_nodes + 1).
1240    pub fn child_offsets(&self) -> &TrackedCudaSlice<u32> {
1241        &self.child_offsets
1242    }
1243
1244    /// Device-resident CSR child indices for AND/OR nodes.
1245    pub fn child_indices(&self) -> &TrackedCudaSlice<u32> {
1246        &self.child_indices
1247    }
1248
1249    /// Device-resident literals for LIT nodes (signed DIMACS, 1-based var ids).
1250    pub fn lit(&self) -> &TrackedCudaSlice<i32> {
1251        &self.lit
1252    }
1253
1254    /// Device-resident decision var ids for DECISION nodes (0 for non-decision).
1255    pub fn decision_var(&self) -> &TrackedCudaSlice<u32> {
1256        &self.decision_var
1257    }
1258
1259    pub fn decision_child_false(&self) -> &TrackedCudaSlice<u32> {
1260        &self.decision_child_false
1261    }
1262
1263    pub fn decision_child_true(&self) -> &TrackedCudaSlice<u32> {
1264        &self.decision_child_true
1265    }
1266
1267    /// Device-resident per-node values buffer (log-space). Written by forward pass.
1268    pub fn values(&self) -> &TrackedCudaSlice<f64> {
1269        &self.values
1270    }
1271
1272    /// Device-resident gradient buffer for ln(true-weight) per CNF variable.
1273    pub fn grad_true(&self) -> &TrackedCudaSlice<f64> {
1274        &self.grad_true
1275    }
1276
1277    /// Device-resident gradient buffer for ln(false-weight) per CNF variable.
1278    pub fn grad_false(&self) -> &TrackedCudaSlice<f64> {
1279        &self.grad_false
1280    }
1281
1282    /// Device-resident log(true-weight) table.
1283    pub fn var_log_true(&self) -> &TrackedCudaSlice<f64> {
1284        &self.var_log_true
1285    }
1286
1287    /// Device-resident log(false-weight) table.
1288    pub fn var_log_false(&self) -> &TrackedCudaSlice<f64> {
1289        &self.var_log_false
1290    }
1291
1292    /// Mutable access to both log-weight tables (true/false) at once.
1293    ///
1294    /// This is useful when passing both slices to a single CUDA kernel launch, avoiding
1295    /// overlapping mutable borrows of `self`.
1296    pub fn var_log_weights_mut(
1297        &mut self,
1298    ) -> (&mut TrackedCudaSlice<f64>, &mut TrackedCudaSlice<f64>) {
1299        (&mut self.var_log_true, &mut self.var_log_false)
1300    }
1301
1302    /// Attach a device-resident free-variable mask (length = max_var + 1).
1303    pub fn set_free_var_mask_device(&mut self, mask: TrackedCudaSlice<u8>) -> Result<()> {
1304        if mask.len() != self.var_log_true.len() {
1305            return Err(XlogError::Compilation(format!(
1306                "GPU free-var mask len {} != weights len {}",
1307                mask.len(),
1308                self.var_log_true.len()
1309            )));
1310        }
1311        self.free_var_mask = Some(mask);
1312        Ok(())
1313    }
1314
1315    /// Upload a host free-variable mask (length = max_var + 1).
1316    #[allow(dead_code)] // reserved: host-side mask upload for testing/diagnostics
1317    pub(crate) fn set_free_var_mask_from_host(
1318        &mut self,
1319        provider: &CudaKernelProvider,
1320        mask: &[u8],
1321    ) -> Result<()> {
1322        if mask.len() != self.var_log_true.len() {
1323            return Err(XlogError::Compilation(format!(
1324                "GPU free-var mask len {} != weights len {}",
1325                mask.len(),
1326                self.var_log_true.len()
1327            )));
1328        }
1329        let memory = provider.memory();
1330        let mut d_mask = memory.alloc::<u8>(mask.len())?;
1331        provider
1332            .htod_sync_copy_into_tracked(mask, &mut d_mask)
1333            .map_err(|e| XlogError::Kernel(format!("Failed to upload free_var_mask: {}", e)))?;
1334        self.free_var_mask = Some(d_mask);
1335        Ok(())
1336    }
1337
1338    /// Upload a host weight table into the device-resident `var_log_true/var_log_false` buffers.
1339    ///
1340    /// This is intended for one-time initialization of static weights (evidence + non-neural facts).
1341    /// Neural fast-path updates should overwrite only the relevant subset on GPU.
1342    /// The 1-based table prefix (1 through `max_var`) is validated before either upload: NaN is
1343    /// rejected with a typed compilation error, while positive and negative infinity remain valid
1344    /// for value evaluation. Reserved slot 0 is uploaded unchanged but is never consumed.
1345    pub fn set_base_weights(
1346        &mut self,
1347        provider: &CudaKernelProvider,
1348        var_log_weights: &[(f64, f64)],
1349    ) -> Result<()> {
1350        let weights_len = (self.max_var as usize) + 1;
1351        if var_log_weights.len() < weights_len {
1352            return Err(XlogError::Compilation(format!(
1353                "GPU XGCF weights init expects weight table len >= {}, got {}",
1354                weights_len,
1355                var_log_weights.len()
1356            )));
1357        }
1358        for &weights in &var_log_weights[1..weights_len] {
1359            validate_circuit_log_weight_pair(weights)?;
1360        }
1361
1362        let mut host_true: Vec<f64> = Vec::with_capacity(weights_len);
1363        let mut host_false: Vec<f64> = Vec::with_capacity(weights_len);
1364        for &(t, f) in &var_log_weights[..weights_len] {
1365            host_true.push(t);
1366            host_false.push(f);
1367        }
1368
1369        provider
1370            .htod_sync_copy_into_tracked(&host_true, &mut self.var_log_true)
1371            .map_err(|e| XlogError::Kernel(format!("Failed to upload log_true weights: {}", e)))?;
1372        provider
1373            .htod_sync_copy_into_tracked(&host_false, &mut self.var_log_false)
1374            .map_err(|e| XlogError::Kernel(format!("Failed to upload log_false weights: {}", e)))?;
1375
1376        Ok(())
1377    }
1378
1379    /// Evaluate logZ on the device using the currently loaded weights and write it into `out_log_z`.
1380    ///
1381    /// This method performs no device->host transfers. Callers that mutate the device weight
1382    /// buffers directly are responsible for numeric validity. NaN or undefined arithmetic that
1383    /// reaches circuit evaluation emits a NaN sentinel in `out_log_z`; a host boundary must
1384    /// validate that scalar after readback. Positive infinity is supported for value-only
1385    /// log-sum-exp evaluation.
1386    pub fn eval_log_wmc_device_inplace(
1387        &mut self,
1388        provider: &CudaKernelProvider,
1389        out_log_z: &mut TrackedCudaSlice<f64>,
1390    ) -> Result<()> {
1391        if out_log_z.len() != 1 {
1392            return Err(XlogError::Compilation(format!(
1393                "GPU device logZ output len {} != 1",
1394                out_log_z.len()
1395            )));
1396        }
1397
1398        let device = provider.device().inner();
1399        let func = device
1400            .get_func(CIRCUIT_MODULE, circuit_kernels::XGCF_FORWARD_LEVEL)
1401            .ok_or_else(|| XlogError::Kernel("xgcf_forward_level kernel not found".to_string()))?;
1402
1403        let block_size: u32 = 256;
1404        let num_levels: usize = self.num_levels as usize;
1405        for level in 0..num_levels {
1406            let num_level_nodes: usize = match self.level_offsets_host.as_ref() {
1407                Some(off) => checked_host_level_width(off, level)?,
1408                None => self.level_nodes.len(),
1409            };
1410            if num_level_nodes == 0 {
1411                continue;
1412            }
1413
1414            let num_blocks =
1415                checked_gpu_launch_blocks("xgcf_forward_level", num_level_nodes, block_size)?;
1416            let config = LaunchConfig {
1417                grid_dim: (num_blocks, 1, 1),
1418                block_dim: (block_size, 1, 1),
1419                shared_mem_bytes: 0,
1420            };
1421            let level_u32: u32 = level as u32;
1422
1423            let mut params: Vec<*mut c_void> = vec![
1424                (&self.node_type).as_kernel_param(),
1425                (&self.child_offsets).as_kernel_param(),
1426                (&self.child_indices).as_kernel_param(),
1427                (&self.lit).as_kernel_param(),
1428                (&self.decision_var).as_kernel_param(),
1429                (&self.decision_child_false).as_kernel_param(),
1430                (&self.decision_child_true).as_kernel_param(),
1431                (&self.level_nodes).as_kernel_param(),
1432                (&self.level_offsets).as_kernel_param(),
1433                level_u32.as_kernel_param(),
1434                (&self.var_log_true).as_kernel_param(),
1435                (&self.var_log_false).as_kernel_param(),
1436                (&self.values).as_kernel_param(),
1437            ];
1438
1439            // SAFETY: xgcf_forward_level(...) writes values for the provided level nodes.
1440            unsafe { func.clone().launch(config, &mut params) }
1441                .map_err(|e| XlogError::Kernel(format!("xgcf_forward_level failed: {}", e)))?;
1442        }
1443
1444        self.apply_free_var_correction(provider, true, false)?;
1445
1446        let root_idx = self.root as usize;
1447        let root_view = self.values.slice(root_idx..(root_idx + 1));
1448        device
1449            .dtod_copy(&root_view, out_log_z)
1450            .map_err(|e| XlogError::Kernel(format!("Failed to copy device logZ: {}", e)))?;
1451
1452        // No device synchronize: callers read back with a synchronous host copy
1453        // or pass the result to subsequent GPU operations (same-stream ordering).
1454        Ok(())
1455    }
1456
1457    /// Evaluate logZ on the device and write it into `out_log_z` (uploads weights from host).
1458    ///
1459    /// Host weights are validated before launch. Undefined arithmetic can still produce a NaN
1460    /// sentinel in `out_log_z`, which the eventual host readback boundary must reject.
1461    pub fn eval_log_wmc_device_into(
1462        &mut self,
1463        provider: &CudaKernelProvider,
1464        var_log_weights: &[(f64, f64)],
1465        out_log_z: &mut TrackedCudaSlice<f64>,
1466    ) -> Result<()> {
1467        self.set_base_weights(provider, var_log_weights)?;
1468        self.eval_log_wmc_device_inplace(provider, out_log_z)
1469    }
1470
1471    /// Evaluate logZ on the device and return a device-resident scalar (uploads weights from host).
1472    ///
1473    /// The returned scalar uses NaN as the device-only sentinel for undefined circuit arithmetic;
1474    /// callers that read it back must convert that sentinel to a typed error.
1475    pub fn eval_log_wmc_device(
1476        &mut self,
1477        provider: &CudaKernelProvider,
1478        var_log_weights: &[(f64, f64)],
1479    ) -> Result<TrackedCudaSlice<f64>> {
1480        let memory = provider.memory();
1481        let mut out_log_z = memory.alloc::<f64>(1)?;
1482        self.eval_log_wmc_device_into(provider, var_log_weights, &mut out_log_z)?;
1483        Ok(out_log_z)
1484    }
1485
1486    fn apply_free_var_correction(
1487        &mut self,
1488        provider: &CudaKernelProvider,
1489        apply_log_z: bool,
1490        apply_grads: bool,
1491    ) -> Result<()> {
1492        let Some(mask) = self.free_var_mask.as_ref() else {
1493            return Ok(());
1494        };
1495
1496        if mask.len() != self.var_log_true.len() {
1497            return Err(XlogError::Compilation(format!(
1498                "GPU free-var mask len {} != weights len {}",
1499                mask.len(),
1500                self.var_log_true.len()
1501            )));
1502        }
1503
1504        let n = u32::try_from(mask.len())
1505            .map_err(|_| XlogError::Compilation("GPU free-var mask length overflow".to_string()))?;
1506        if n == 0 {
1507            return Ok(());
1508        }
1509
1510        let device = provider.device().inner();
1511        let block_dim = 256u32;
1512        let grid_dim = n.div_ceil(block_dim);
1513
1514        if apply_grads {
1515            let apply_grad = device
1516                .get_func(CIRCUIT_MODULE, circuit_kernels::XGCF_FREE_VAR_APPLY_GRAD)
1517                .ok_or_else(|| {
1518                    XlogError::Kernel("xgcf_free_var_apply_grad kernel not found".to_string())
1519                })?;
1520            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1521            unsafe {
1522                apply_grad.clone().launch(
1523                    LaunchConfig {
1524                        grid_dim: (grid_dim, 1, 1),
1525                        block_dim: (block_dim, 1, 1),
1526                        shared_mem_bytes: 0,
1527                    },
1528                    (
1529                        mask,
1530                        &self.var_log_true,
1531                        &self.var_log_false,
1532                        n,
1533                        &mut self.grad_true,
1534                        &mut self.grad_false,
1535                    ),
1536                )
1537            }
1538            .map_err(|e| XlogError::Kernel(format!("xgcf_free_var_apply_grad failed: {}", e)))?;
1539        }
1540
1541        if apply_log_z {
1542            let reduce_stage = device
1543                .get_func(CIRCUIT_MODULE, circuit_kernels::XGCF_FREE_VAR_REDUCE_STAGE)
1544                .ok_or_else(|| {
1545                    XlogError::Kernel("xgcf_free_var_reduce_stage kernel not found".to_string())
1546                })?;
1547            let add_scalar = device
1548                .get_func(CIRCUIT_MODULE, circuit_kernels::XGCF_ADD_SCALAR)
1549                .ok_or_else(|| XlogError::Kernel("xgcf_add_scalar kernel not found".to_string()))?;
1550
1551            let memory = provider.memory();
1552            let mut buf_a = memory.alloc::<f64>(mask.len())?;
1553            let mut buf_b = memory.alloc::<f64>(mask.len())?;
1554
1555            let mut stage_n = n;
1556            let mut stage0 = true;
1557            let mut output_is_a = true;
1558            loop {
1559                let out_len = stage_n.div_ceil(2);
1560                let stage_grid = out_len.div_ceil(block_dim);
1561
1562                let (in_buf, out_buf): (&TrackedCudaSlice<f64>, &mut TrackedCudaSlice<f64>) =
1563                    if output_is_a {
1564                        (&buf_b, &mut buf_a)
1565                    } else {
1566                        (&buf_a, &mut buf_b)
1567                    };
1568                let mode = if stage0 { 0u32 } else { 1u32 };
1569
1570                // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1571                unsafe {
1572                    reduce_stage.clone().launch(
1573                        LaunchConfig {
1574                            grid_dim: (stage_grid, 1, 1),
1575                            block_dim: (block_dim, 1, 1),
1576                            shared_mem_bytes: 0,
1577                        },
1578                        (
1579                            mask,
1580                            &self.var_log_true,
1581                            &self.var_log_false,
1582                            in_buf,
1583                            stage_n,
1584                            mode,
1585                            out_buf,
1586                        ),
1587                    )
1588                }
1589                .map_err(|e| {
1590                    XlogError::Kernel(format!("xgcf_free_var_reduce_stage failed: {}", e))
1591                })?;
1592
1593                if out_len == 1 {
1594                    let result_buf = if output_is_a { &buf_a } else { &buf_b };
1595                    // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1596                    unsafe {
1597                        add_scalar.clone().launch(
1598                            LaunchConfig {
1599                                grid_dim: (1, 1, 1),
1600                                block_dim: (1, 1, 1),
1601                                shared_mem_bytes: 0,
1602                            },
1603                            (&mut self.values, self.root, result_buf),
1604                        )
1605                    }
1606                    .map_err(|e| XlogError::Kernel(format!("xgcf_add_scalar failed: {}", e)))?;
1607                    break;
1608                }
1609
1610                stage_n = out_len;
1611                stage0 = false;
1612                output_is_a = !output_is_a;
1613            }
1614        }
1615
1616        Ok(())
1617    }
1618
1619    /// Evaluate the circuit and populate `grad_true/grad_false` on the device (no host reads).
1620    ///
1621    /// Preconditions:
1622    /// - `var_log_true/var_log_false` contain the current weights on device. Callers must validate
1623    ///   numeric outputs: NaN or undefined normalization that reaches evaluation is represented by
1624    ///   a non-finite device result because this API cannot return a host-side numeric error.
1625    /// - Caller may read back results for testing/debugging, but this API performs no dtoh transfers.
1626    pub fn eval_grads_inplace(&mut self, provider: &CudaKernelProvider) -> Result<()> {
1627        let device = provider.device().inner();
1628
1629        // Forward pass (identical to eval_log_wmc, minus weight upload and root readback).
1630        let func = device
1631            .get_func(CIRCUIT_MODULE, circuit_kernels::XGCF_FORWARD_LEVEL)
1632            .ok_or_else(|| XlogError::Kernel("xgcf_forward_level kernel not found".to_string()))?;
1633
1634        let block_size: u32 = 256;
1635        let num_levels: usize = self.num_levels as usize;
1636        for level in 0..num_levels {
1637            let num_level_nodes: usize = match self.level_offsets_host.as_ref() {
1638                Some(off) => checked_host_level_width(off, level)?,
1639                None => self.level_nodes.len(),
1640            };
1641            if num_level_nodes == 0 {
1642                continue;
1643            }
1644
1645            let num_blocks =
1646                checked_gpu_launch_blocks("xgcf_forward_level", num_level_nodes, block_size)?;
1647            let config = LaunchConfig {
1648                grid_dim: (num_blocks, 1, 1),
1649                block_dim: (block_size, 1, 1),
1650                shared_mem_bytes: 0,
1651            };
1652            let level_u32: u32 = level as u32;
1653
1654            let mut params: Vec<*mut c_void> = vec![
1655                (&self.node_type).as_kernel_param(),
1656                (&self.child_offsets).as_kernel_param(),
1657                (&self.child_indices).as_kernel_param(),
1658                (&self.lit).as_kernel_param(),
1659                (&self.decision_var).as_kernel_param(),
1660                (&self.decision_child_false).as_kernel_param(),
1661                (&self.decision_child_true).as_kernel_param(),
1662                (&self.level_nodes).as_kernel_param(),
1663                (&self.level_offsets).as_kernel_param(),
1664                level_u32.as_kernel_param(),
1665                (&self.var_log_true).as_kernel_param(),
1666                (&self.var_log_false).as_kernel_param(),
1667                (&self.values).as_kernel_param(),
1668            ];
1669
1670            // SAFETY: xgcf_forward_level(...) writes values for the provided level nodes.
1671            unsafe { func.clone().launch(config, &mut params) }
1672                .map_err(|e| XlogError::Kernel(format!("xgcf_forward_level failed: {}", e)))?;
1673        }
1674
1675        // Backward pass buffers.
1676        device
1677            .memset_zeros(&mut self.adj)
1678            .map_err(|e| XlogError::Kernel(format!("Failed to zero adj buffer: {}", e)))?;
1679        device
1680            .memset_zeros(&mut self.grad_true)
1681            .map_err(|e| XlogError::Kernel(format!("Failed to zero grad_true buffer: {}", e)))?;
1682        device
1683            .memset_zeros(&mut self.grad_false)
1684            .map_err(|e| XlogError::Kernel(format!("Failed to zero grad_false buffer: {}", e)))?;
1685
1686        // Set root adjoint to 1.0 via GPU kernel (avoid host copy).
1687        let root_idx = self.root as usize;
1688        let mut root_adj_view = self.adj.slice_mut(root_idx..(root_idx + 1));
1689        let fill_const = device
1690            .get_func(ARITH_MODULE, arith_kernels::ARITH_FILL_CONST_F64)
1691            .ok_or_else(|| {
1692                XlogError::Kernel("arith_fill_const_f64 kernel not found".to_string())
1693            })?;
1694        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1695        unsafe {
1696            fill_const.clone().launch(
1697                LaunchConfig {
1698                    grid_dim: (1, 1, 1),
1699                    block_dim: (1, 1, 1),
1700                    shared_mem_bytes: 0,
1701                },
1702                (1.0_f64, 1u32, &mut root_adj_view),
1703            )
1704        }
1705        .map_err(|e| XlogError::Kernel(format!("arith_fill_const_f64 failed: {}", e)))?;
1706
1707        let propagate = device
1708            .get_func(
1709                CIRCUIT_MODULE,
1710                circuit_kernels::XGCF_BACKWARD_LEVEL_PROPAGATE,
1711            )
1712            .ok_or_else(|| {
1713                XlogError::Kernel("xgcf_backward_level_propagate kernel not found".to_string())
1714            })?;
1715        let decision_grad = device
1716            .get_func(
1717                CIRCUIT_MODULE,
1718                circuit_kernels::XGCF_BACKWARD_LEVEL_DECISION_GRAD,
1719            )
1720            .ok_or_else(|| {
1721                XlogError::Kernel("xgcf_backward_level_decision_grad kernel not found".to_string())
1722            })?;
1723        let lit_grad = device
1724            .get_func(
1725                CIRCUIT_MODULE,
1726                circuit_kernels::XGCF_BACKWARD_LEVEL_LIT_GRAD,
1727            )
1728            .ok_or_else(|| {
1729                XlogError::Kernel("xgcf_backward_level_lit_grad kernel not found".to_string())
1730            })?;
1731
1732        let num_levels: usize = self.num_levels as usize;
1733        for level in (0..num_levels).rev() {
1734            let num_level_nodes: usize = match self.level_offsets_host.as_ref() {
1735                Some(off) => checked_host_level_width(off, level)?,
1736                None => self.level_nodes.len(),
1737            };
1738            if num_level_nodes == 0 {
1739                continue;
1740            }
1741
1742            let num_blocks =
1743                checked_gpu_launch_blocks("xgcf_backward_level", num_level_nodes, block_size)?;
1744            let config = LaunchConfig {
1745                grid_dim: (num_blocks, 1, 1),
1746                block_dim: (block_size, 1, 1),
1747                shared_mem_bytes: 0,
1748            };
1749            let level_u32: u32 = level as u32;
1750
1751            let mut params: Vec<*mut c_void> = vec![
1752                (&self.node_type).as_kernel_param(),
1753                (&self.child_offsets).as_kernel_param(),
1754                (&self.child_indices).as_kernel_param(),
1755                (&self.decision_var).as_kernel_param(),
1756                (&self.decision_child_false).as_kernel_param(),
1757                (&self.decision_child_true).as_kernel_param(),
1758                (&self.level_nodes).as_kernel_param(),
1759                (&self.level_offsets).as_kernel_param(),
1760                level_u32.as_kernel_param(),
1761                (&self.var_log_true).as_kernel_param(),
1762                (&self.var_log_false).as_kernel_param(),
1763                (&self.values).as_kernel_param(),
1764                (&self.adj).as_kernel_param(),
1765            ];
1766
1767            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1768            unsafe { propagate.clone().launch(config, &mut params) }.map_err(|e| {
1769                XlogError::Kernel(format!("xgcf_backward_level_propagate failed: {}", e))
1770            })?;
1771
1772            let mut params: Vec<*mut c_void> = vec![
1773                (&self.node_type).as_kernel_param(),
1774                (&self.decision_var).as_kernel_param(),
1775                (&self.decision_child_false).as_kernel_param(),
1776                (&self.decision_child_true).as_kernel_param(),
1777                (&self.level_nodes).as_kernel_param(),
1778                (&self.level_offsets).as_kernel_param(),
1779                level_u32.as_kernel_param(),
1780                (&self.var_log_true).as_kernel_param(),
1781                (&self.var_log_false).as_kernel_param(),
1782                (&self.values).as_kernel_param(),
1783                (&self.adj).as_kernel_param(),
1784                (&self.grad_true).as_kernel_param(),
1785                (&self.grad_false).as_kernel_param(),
1786            ];
1787
1788            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1789            unsafe { decision_grad.clone().launch(config, &mut params) }.map_err(|e| {
1790                XlogError::Kernel(format!("xgcf_backward_level_decision_grad failed: {}", e))
1791            })?;
1792
1793            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1794            unsafe {
1795                lit_grad.clone().launch(
1796                    config,
1797                    (
1798                        &self.node_type,
1799                        &self.lit,
1800                        &self.level_nodes,
1801                        &self.level_offsets,
1802                        level_u32,
1803                        &self.adj,
1804                        &self.grad_true,
1805                        &self.grad_false,
1806                    ),
1807                )
1808            }
1809            .map_err(|e| {
1810                XlogError::Kernel(format!("xgcf_backward_level_lit_grad failed: {}", e))
1811            })?;
1812        }
1813
1814        self.apply_free_var_correction(provider, true, true)?;
1815        // No device synchronize: callers batch multiple eval/backward calls
1816        // before syncing at the query boundary.
1817        Ok(())
1818    }
1819
1820    #[cfg(feature = "host-io")]
1821    pub fn eval_log_wmc(
1822        &mut self,
1823        provider: &CudaKernelProvider,
1824        var_log_weights: &[(f64, f64)],
1825    ) -> Result<f64> {
1826        let device = provider.device().inner();
1827        let mut out_log_z = provider.memory().alloc::<f64>(1)?;
1828        self.eval_log_wmc_device_into(provider, var_log_weights, &mut out_log_z)?;
1829
1830        let mut host = [0.0_f64];
1831        device
1832            .dtoh_sync_copy_into(&out_log_z, &mut host)
1833            .map_err(|e| XlogError::Kernel(format!("Failed to read circuit root value: {}", e)))?;
1834        validate_circuit_value(host[0])
1835    }
1836
1837    #[cfg(feature = "host-io")]
1838    pub fn eval_log_wmc_and_grads(
1839        &mut self,
1840        provider: &CudaKernelProvider,
1841        var_log_weights: &[(f64, f64)],
1842    ) -> Result<(f64, Vec<f64>, Vec<f64>)> {
1843        let weights_len = (self.max_var as usize) + 1;
1844        if var_log_weights.len() < weights_len {
1845            return Err(XlogError::Compilation(format!(
1846                "GPU XGCF weights init expects weight table len >= {}, got {}",
1847                weights_len,
1848                var_log_weights.len()
1849            )));
1850        }
1851        self.set_base_weights(provider, var_log_weights)?;
1852        self.eval_grads_inplace(provider)?;
1853
1854        let device = provider.device().inner();
1855
1856        let mut host_grad_true: Vec<f64> = vec![0.0; weights_len];
1857        let mut host_grad_false: Vec<f64> = vec![0.0; weights_len];
1858
1859        let root_idx = self.root as usize;
1860        let root_view = self.values.slice(root_idx..(root_idx + 1));
1861        let mut log_z = [0.0_f64];
1862        device
1863            .dtoh_sync_copy_into(&root_view, &mut log_z)
1864            .map_err(|e| XlogError::Kernel(format!("Failed to read circuit root value: {}", e)))?;
1865        let log_z = validate_circuit_value(log_z[0])?;
1866
1867        device
1868            .dtoh_sync_copy_into(&self.grad_true, &mut host_grad_true)
1869            .map_err(|e| XlogError::Kernel(format!("Failed to download grad_true: {}", e)))?;
1870        device
1871            .dtoh_sync_copy_into(&self.grad_false, &mut host_grad_false)
1872            .map_err(|e| XlogError::Kernel(format!("Failed to download grad_false: {}", e)))?;
1873        validate_circuit_gradient_values(&host_grad_true, &host_grad_false)?;
1874
1875        Ok((log_z, host_grad_true, host_grad_false))
1876    }
1877}