Skip to main content

xlog_prob/compilation/
gpu_cache.rs

1//! GPU-resident circuit cache helpers.
2
3use std::sync::Arc;
4
5use cudarc::driver::{DeviceSlice, LaunchConfig};
6use xlog_core::{Result, XlogError};
7use xlog_cuda::memory::TrackedCudaSlice;
8use xlog_cuda::provider::{cache_kernels, CACHE_MODULE};
9use xlog_cuda::{AsKernelParam, CudaKernelProvider, LaunchAsync};
10use xlog_solve::GpuCnf;
11
12use super::disk_cache;
13use crate::gpu::GpuXgcf;
14
15/// Configuration for the GPU-resident circuit cache.
16///
17/// Controls the number of cached XGCF circuit slots and the per-slot capacity
18/// limits for nodes, edges, levels, and variables. Production callers should
19/// use [`crate::exact::default_cache_config`] which sizes caps from the CNF
20/// and compile config.
21#[derive(Debug, Clone, Copy)]
22#[non_exhaustive]
23pub struct GpuCircuitCacheConfig {
24    /// Number of circuit slots kept resident on the GPU.
25    pub num_slots: u32,
26    /// Hash table size for the circuit lookup (should be >= 2 * num_slots).
27    pub table_size: u32,
28    /// Maximum nodes per cached circuit.
29    pub node_cap: u32,
30    /// Maximum edges per cached circuit.
31    pub edge_cap: u32,
32    /// Maximum levels (BFS depth) per cached circuit.
33    pub level_cap: u32,
34    /// Maximum CNF variable id (1-based, DIMACS) across all cached circuits.
35    pub var_cap: u32,
36}
37
38impl Default for GpuCircuitCacheConfig {
39    /// Conservative defaults for small CNFs (< 64 variables).
40    ///
41    /// Production callers should derive caps from the actual CNF dimensions.
42    fn default() -> Self {
43        Self {
44            num_slots: 4,
45            table_size: 8,
46            node_cap: 65_536,
47            edge_cap: 131_072,
48            level_cap: 65_536,
49            var_cap: 128,
50        }
51    }
52}
53
54fn cache_grid_dim_for_u32_count(context: &str, count: u32, block_dim: u32) -> Result<u32> {
55    if count == 0 {
56        return Ok(0);
57    }
58    if block_dim == 0 {
59        return Err(XlogError::Compilation(format!(
60            "{context}: GPU cache block size must be nonzero"
61        )));
62    }
63    let padded = count
64        .checked_add(block_dim - 1)
65        .ok_or_else(|| XlogError::Compilation(format!("{context}: GPU cache grid overflow")))?;
66    Ok(padded / block_dim)
67}
68
69fn cache_grid_dim_for_u64_count(context: &str, count: u64, block_dim: u32) -> Result<u32> {
70    if count == 0 {
71        return Ok(0);
72    }
73    if block_dim == 0 {
74        return Err(XlogError::Compilation(format!(
75            "{context}: GPU cache block size must be nonzero"
76        )));
77    }
78    let block = block_dim as u64;
79    let grid = count
80        .checked_add(block - 1)
81        .map(|padded| padded / block)
82        .ok_or_else(|| XlogError::Compilation(format!("{context}: GPU cache grid overflow")))?;
83    u32::try_from(grid)
84        .map_err(|_| XlogError::Compilation(format!("{context}: GPU cache grid exceeds u32")))
85}
86
87pub struct GpuCircuitCache {
88    provider: Arc<CudaKernelProvider>,
89    table_size: u32,
90    num_slots: u32,
91    node_cap: u32,
92    edge_cap: u32,
93    level_cap: u32,
94    var_cap: u32,
95    keys: TrackedCudaSlice<u64>,
96    slots: TrackedCudaSlice<u32>,
97    state: TrackedCudaSlice<u32>,
98    last_used: TrackedCudaSlice<u64>,
99    slot_states: TrackedCudaSlice<u32>,
100    clock: TrackedCudaSlice<u64>,
101    node_type: TrackedCudaSlice<u8>,
102    child_offsets: TrackedCudaSlice<u32>,
103    child_indices: TrackedCudaSlice<u32>,
104    lit: TrackedCudaSlice<i32>,
105    decision_var: TrackedCudaSlice<u32>,
106    decision_child_false: TrackedCudaSlice<u32>,
107    decision_child_true: TrackedCudaSlice<u32>,
108    level_nodes: TrackedCudaSlice<u32>,
109    level_offsets: TrackedCudaSlice<u32>,
110    var_log_true: TrackedCudaSlice<f64>,
111    var_log_false: TrackedCudaSlice<f64>,
112    values: TrackedCudaSlice<f64>,
113    adj: TrackedCudaSlice<f64>,
114    grad_true: TrackedCudaSlice<f64>,
115    grad_false: TrackedCudaSlice<f64>,
116    meta_num_nodes: TrackedCudaSlice<u32>,
117    meta_num_levels: TrackedCudaSlice<u32>,
118    meta_root: TrackedCudaSlice<u32>,
119    meta_max_var: TrackedCudaSlice<u32>,
120    always_on: TrackedCudaSlice<u32>,
121    zero_f64: TrackedCudaSlice<f64>,
122    one_f64: TrackedCudaSlice<f64>,
123    free_var_mask: TrackedCudaSlice<u8>,
124    has_free_var_mask: Vec<bool>,
125}
126
127pub struct GpuCacheLookup {
128    provider: Arc<CudaKernelProvider>,
129    slot: TrackedCudaSlice<u32>,
130    compile_needed: TrackedCudaSlice<u32>,
131}
132
133impl GpuCacheLookup {
134    pub fn slot_device(&self) -> &TrackedCudaSlice<u32> {
135        &self.slot
136    }
137
138    pub fn compile_needed_device(&self) -> &TrackedCudaSlice<u32> {
139        &self.compile_needed
140    }
141
142    pub fn provider(&self) -> &Arc<CudaKernelProvider> {
143        &self.provider
144    }
145
146    pub fn into_handle(self) -> Result<GpuCircuitCacheHandle> {
147        let slot_host_vec: Vec<u32> = self
148            .provider
149            .device()
150            .inner()
151            .dtoh_sync_copy(&self.slot)
152            .map_err(|e| XlogError::Kernel(format!("dtoh slot index: {}", e)))?;
153        Ok(GpuCircuitCacheHandle {
154            provider: self.provider,
155            slot: self.slot,
156            compile_needed: self.compile_needed,
157            slot_host: slot_host_vec[0],
158            num_nodes: 0,
159            num_levels: 0,
160            root: 0,
161            max_var: 0,
162        })
163    }
164}
165
166pub struct GpuCircuitCacheHandle {
167    provider: Arc<CudaKernelProvider>,
168    slot: TrackedCudaSlice<u32>,
169    compile_needed: TrackedCudaSlice<u32>,
170    slot_host: u32,
171    num_nodes: u32,
172    num_levels: u32,
173    root: u32,
174    max_var: u32,
175}
176
177impl GpuCircuitCacheHandle {
178    pub fn slot_device(&self) -> &TrackedCudaSlice<u32> {
179        &self.slot
180    }
181
182    pub fn compile_needed_device(&self) -> &TrackedCudaSlice<u32> {
183        &self.compile_needed
184    }
185
186    pub fn provider(&self) -> &Arc<CudaKernelProvider> {
187        &self.provider
188    }
189
190    pub fn num_nodes(&self) -> u32 {
191        self.num_nodes
192    }
193
194    pub fn num_levels(&self) -> u32 {
195        self.num_levels
196    }
197
198    pub fn root(&self) -> u32 {
199        self.root
200    }
201
202    pub fn max_var(&self) -> u32 {
203        self.max_var
204    }
205
206    pub(crate) fn slot_index(&self) -> u32 {
207        self.slot_host
208    }
209}
210
211/// Compute a deterministic CNF hash on the GPU.
212///
213/// Hash input order matches the cache kernel: num_vars, num_clauses, num_lits,
214/// clause_offsets[0..num_clauses], literals[0..num_lits-1].
215pub fn hash_cnf_gpu(
216    cnf: &GpuCnf,
217    provider: &Arc<CudaKernelProvider>,
218) -> Result<TrackedCudaSlice<u64>> {
219    let memory = provider.memory();
220    let mut out_hash = memory.alloc::<u64>(1)?;
221
222    let func = provider
223        .device()
224        .inner()
225        .get_func(CACHE_MODULE, cache_kernels::CACHE_CNF_HASH)
226        .ok_or_else(|| XlogError::Kernel("cache_cnf_hash kernel not found".to_string()))?;
227
228    // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
229    unsafe {
230        func.clone().launch(
231            LaunchConfig {
232                grid_dim: (1, 1, 1),
233                block_dim: (1, 1, 1),
234                shared_mem_bytes: 0,
235            },
236            (
237                &cnf.num_vars,
238                &cnf.num_clauses,
239                &cnf.num_lits,
240                &cnf.clause_offsets,
241                &cnf.literals,
242                &mut out_hash,
243            ),
244        )
245    }
246    .map_err(|e| XlogError::Kernel(format!("cache_cnf_hash launch failed: {}", e)))?;
247    // No device synchronize: hash stays device-resident for lookup kernel; same-stream ordering suffices.
248    Ok(out_hash)
249}
250
251impl GpuCircuitCache {
252    pub fn provider(&self) -> &Arc<CudaKernelProvider> {
253        &self.provider
254    }
255
256    /// Returns mutable device-resident log-weight tables.
257    ///
258    /// Device-only callers are responsible for validating numeric inputs and outputs. Positive
259    /// infinity is valid when it does not enter an undefined normalization; NaN or undefined
260    /// arithmetic that reaches evaluation is represented by a non-finite device result.
261    pub fn var_log_weights_mut(
262        &mut self,
263    ) -> (&mut TrackedCudaSlice<f64>, &mut TrackedCudaSlice<f64>) {
264        (&mut self.var_log_true, &mut self.var_log_false)
265    }
266
267    pub fn grad_true(&self) -> &TrackedCudaSlice<f64> {
268        &self.grad_true
269    }
270
271    pub fn grad_false(&self) -> &TrackedCudaSlice<f64> {
272        &self.grad_false
273    }
274
275    pub fn values(&self) -> &TrackedCudaSlice<f64> {
276        &self.values
277    }
278
279    pub fn meta_num_nodes_device(&self) -> &TrackedCudaSlice<u32> {
280        &self.meta_num_nodes
281    }
282
283    pub fn meta_num_levels_device(&self) -> &TrackedCudaSlice<u32> {
284        &self.meta_num_levels
285    }
286
287    pub fn meta_root_device(&self) -> &TrackedCudaSlice<u32> {
288        &self.meta_root
289    }
290
291    pub fn meta_max_var_device(&self) -> &TrackedCudaSlice<u32> {
292        &self.meta_max_var
293    }
294
295    pub fn num_slots(&self) -> u32 {
296        self.num_slots
297    }
298
299    pub(crate) fn has_any_free_var_mask(&self) -> bool {
300        self.has_free_var_mask.iter().any(|&v| v)
301    }
302
303    pub(crate) fn has_free_var_mask_for_slot(&self, slot: u32) -> bool {
304        self.has_free_var_mask
305            .get(slot as usize)
306            .copied()
307            .unwrap_or(false)
308    }
309
310    pub(crate) fn var_stride(&self) -> Result<u32> {
311        self.var_cap
312            .checked_add(1)
313            .ok_or_else(|| XlogError::Compilation("GpuCircuitCache var_cap overflow".to_string()))
314    }
315
316    pub(crate) fn node_stride(&self) -> u32 {
317        self.node_cap
318    }
319
320    pub(crate) fn copy_slot_weights_to_batch(
321        &mut self,
322        handle: &GpuCircuitCacheHandle,
323        out_true_batch: &mut TrackedCudaSlice<f64>,
324        out_false_batch: &mut TrackedCudaSlice<f64>,
325        batch_size: u32,
326    ) -> Result<()> {
327        if batch_size == 0 {
328            return Ok(());
329        }
330        let var_stride = self.var_stride()?;
331        let expected = (batch_size as usize)
332            .checked_mul(var_stride as usize)
333            .ok_or_else(|| {
334                XlogError::Compilation("GpuCircuitCache batch weight size overflow".to_string())
335            })?;
336        if out_true_batch.len() != expected || out_false_batch.len() != expected {
337            return Err(XlogError::Compilation(format!(
338                "GpuCircuitCache batched weight buffers must both have len {}, got {} and {}",
339                expected,
340                out_true_batch.len(),
341                out_false_batch.len()
342            )));
343        }
344
345        let device = self.provider.device().inner();
346        let func = device
347            .get_func(
348                xlog_cuda::provider::WEIGHTS_MODULE,
349                xlog_cuda::provider::weights_kernels::WEIGHTS_COPY_SLOT_TO_BATCH,
350            )
351            .ok_or_else(|| {
352                XlogError::Kernel("weights_copy_slot_to_batch kernel not found".to_string())
353            })?;
354
355        let block_dim = 256u32;
356        let total = (batch_size as u64)
357            .checked_mul(var_stride as u64)
358            .ok_or_else(|| {
359                XlogError::Compilation("GpuCircuitCache batch copy overflow".to_string())
360            })?;
361        let grid_dim =
362            cache_grid_dim_for_u64_count("GpuCircuitCache batch weight copy", total, block_dim)?;
363        if grid_dim == 0 {
364            return Ok(());
365        }
366
367        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
368        unsafe {
369            func.clone().launch(
370                LaunchConfig {
371                    grid_dim: (grid_dim, 1, 1),
372                    block_dim: (block_dim, 1, 1),
373                    shared_mem_bytes: 0,
374                },
375                (
376                    handle.slot_device(),
377                    self.var_cap,
378                    &self.var_log_true,
379                    &self.var_log_false,
380                    out_true_batch,
381                    out_false_batch,
382                    var_stride,
383                    batch_size,
384                ),
385            )
386        }
387        .map_err(|e| XlogError::Kernel(format!("weights_copy_slot_to_batch failed: {}", e)))?;
388
389        Ok(())
390    }
391
392    #[allow(clippy::too_many_arguments)]
393    pub(crate) fn eval_grads_inplace_fused_batched(
394        &mut self,
395        handle: &GpuCircuitCacheHandle,
396        var_log_true_batch: &TrackedCudaSlice<f64>,
397        var_log_false_batch: &TrackedCudaSlice<f64>,
398        values_batch: &mut TrackedCudaSlice<f64>,
399        adj_batch: &mut TrackedCudaSlice<f64>,
400        grad_true_batch: &mut TrackedCudaSlice<f64>,
401        grad_false_batch: &mut TrackedCudaSlice<f64>,
402        batch_size: u32,
403    ) -> Result<()> {
404        if batch_size == 0 {
405            return Ok(());
406        }
407        if self.has_free_var_mask_for_slot(handle.slot_index()) {
408            return Err(XlogError::Execution(
409                "Batched fused eval currently does not support free-var correction".to_string(),
410            ));
411        }
412
413        let var_stride = self.var_stride()?;
414        let node_stride = self.node_stride();
415        let expected_var = (batch_size as usize)
416            .checked_mul(var_stride as usize)
417            .ok_or_else(|| {
418                XlogError::Compilation("GpuCircuitCache batched var buffer overflow".to_string())
419            })?;
420        let expected_node = (batch_size as usize)
421            .checked_mul(node_stride as usize)
422            .ok_or_else(|| {
423                XlogError::Compilation("GpuCircuitCache batched node buffer overflow".to_string())
424            })?;
425
426        if var_log_true_batch.len() != expected_var
427            || var_log_false_batch.len() != expected_var
428            || grad_true_batch.len() != expected_var
429            || grad_false_batch.len() != expected_var
430        {
431            return Err(XlogError::Compilation(format!(
432                "GpuCircuitCache batched var buffers must have len {}",
433                expected_var
434            )));
435        }
436        if values_batch.len() != expected_node || adj_batch.len() != expected_node {
437            return Err(XlogError::Compilation(format!(
438                "GpuCircuitCache batched node buffers must have len {}",
439                expected_node
440            )));
441        }
442
443        let device = self.provider.device().inner();
444        device
445            .memset_zeros(adj_batch)
446            .map_err(|e| XlogError::Kernel(format!("Failed to zero batched adj: {}", e)))?;
447        device
448            .memset_zeros(grad_true_batch)
449            .map_err(|e| XlogError::Kernel(format!("Failed to zero batched grad_true: {}", e)))?;
450        device
451            .memset_zeros(grad_false_batch)
452            .map_err(|e| XlogError::Kernel(format!("Failed to zero batched grad_false: {}", e)))?;
453
454        let eval_all = device
455            .get_func(
456                xlog_cuda::CIRCUIT_MODULE,
457                xlog_cuda::circuit_kernels::XGCF_EVAL_ALL_LEVELS_CACHED_BATCHED,
458            )
459            .ok_or_else(|| {
460                XlogError::Kernel("xgcf_eval_all_levels_cached_batched not found".to_string())
461            })?;
462        let set_root_adj = device
463            .get_func(
464                xlog_cuda::CIRCUIT_MODULE,
465                xlog_cuda::circuit_kernels::XGCF_SET_ROOT_ADJ_CACHED_BATCHED,
466            )
467            .ok_or_else(|| {
468                XlogError::Kernel("xgcf_set_root_adj_cached_batched not found".to_string())
469            })?;
470        let backward_all = device
471            .get_func(
472                xlog_cuda::CIRCUIT_MODULE,
473                xlog_cuda::circuit_kernels::XGCF_BACKWARD_ALL_LEVELS_CACHED_BATCHED,
474            )
475            .ok_or_else(|| {
476                XlogError::Kernel("xgcf_backward_all_levels_cached_batched not found".to_string())
477            })?;
478
479        let block_size = 256u32;
480        let mut eval_params: Vec<*mut std::ffi::c_void> = vec![
481            handle.slot_device().as_kernel_param(),
482            self.node_cap.as_kernel_param(),
483            self.edge_cap.as_kernel_param(),
484            self.level_cap.as_kernel_param(),
485            self.var_cap.as_kernel_param(),
486            (&self.node_type).as_kernel_param(),
487            (&self.child_offsets).as_kernel_param(),
488            (&self.child_indices).as_kernel_param(),
489            (&self.lit).as_kernel_param(),
490            (&self.decision_var).as_kernel_param(),
491            (&self.decision_child_false).as_kernel_param(),
492            (&self.decision_child_true).as_kernel_param(),
493            (&self.level_nodes).as_kernel_param(),
494            (&self.level_offsets).as_kernel_param(),
495            (&self.meta_num_levels).as_kernel_param(),
496            var_log_true_batch.as_kernel_param(),
497            var_log_false_batch.as_kernel_param(),
498            var_stride.as_kernel_param(),
499            values_batch.as_kernel_param(),
500            node_stride.as_kernel_param(),
501            batch_size.as_kernel_param(),
502        ];
503        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
504        unsafe {
505            eval_all.clone().launch(
506                LaunchConfig {
507                    grid_dim: (batch_size, 1, 1),
508                    block_dim: (block_size, 1, 1),
509                    shared_mem_bytes: 0,
510                },
511                &mut eval_params,
512            )
513        }
514        .map_err(|e| {
515            XlogError::Kernel(format!("xgcf_eval_all_levels_cached_batched failed: {}", e))
516        })?;
517
518        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
519        unsafe {
520            set_root_adj.clone().launch(
521                LaunchConfig {
522                    grid_dim: (batch_size, 1, 1),
523                    block_dim: (1, 1, 1),
524                    shared_mem_bytes: 0,
525                },
526                (
527                    handle.slot_device(),
528                    self.node_cap,
529                    &self.meta_root,
530                    &mut *adj_batch,
531                    node_stride,
532                    batch_size,
533                ),
534            )
535        }
536        .map_err(|e| {
537            XlogError::Kernel(format!("xgcf_set_root_adj_cached_batched failed: {}", e))
538        })?;
539
540        let mut backward_params: Vec<*mut std::ffi::c_void> = vec![
541            handle.slot_device().as_kernel_param(),
542            self.node_cap.as_kernel_param(),
543            self.edge_cap.as_kernel_param(),
544            self.level_cap.as_kernel_param(),
545            self.var_cap.as_kernel_param(),
546            (&self.node_type).as_kernel_param(),
547            (&self.child_offsets).as_kernel_param(),
548            (&self.child_indices).as_kernel_param(),
549            (&self.decision_var).as_kernel_param(),
550            (&self.decision_child_false).as_kernel_param(),
551            (&self.decision_child_true).as_kernel_param(),
552            (&self.lit).as_kernel_param(),
553            (&self.level_nodes).as_kernel_param(),
554            (&self.level_offsets).as_kernel_param(),
555            (&self.meta_num_levels).as_kernel_param(),
556            var_log_true_batch.as_kernel_param(),
557            var_log_false_batch.as_kernel_param(),
558            var_stride.as_kernel_param(),
559            values_batch.as_kernel_param(),
560            node_stride.as_kernel_param(),
561            adj_batch.as_kernel_param(),
562            node_stride.as_kernel_param(),
563            grad_true_batch.as_kernel_param(),
564            grad_false_batch.as_kernel_param(),
565            var_stride.as_kernel_param(),
566            batch_size.as_kernel_param(),
567        ];
568        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
569        unsafe {
570            backward_all.clone().launch(
571                LaunchConfig {
572                    grid_dim: (batch_size, 1, 1),
573                    block_dim: (block_size, 1, 1),
574                    shared_mem_bytes: 0,
575                },
576                &mut backward_params,
577            )
578        }
579        .map_err(|e| {
580            XlogError::Kernel(format!(
581                "xgcf_backward_all_levels_cached_batched failed: {}",
582                e
583            ))
584        })?;
585
586        Ok(())
587    }
588
589    pub(crate) fn copy_root_batched_from_values(
590        &self,
591        handle: &GpuCircuitCacheHandle,
592        values_batch: &TrackedCudaSlice<f64>,
593        out_roots: &mut TrackedCudaSlice<f64>,
594        batch_size: u32,
595    ) -> Result<()> {
596        if batch_size == 0 {
597            return Ok(());
598        }
599        let node_stride = self.node_stride();
600        let expected_values = (batch_size as usize)
601            .checked_mul(node_stride as usize)
602            .ok_or_else(|| {
603                XlogError::Compilation("GpuCircuitCache batched values overflow".to_string())
604            })?;
605        if values_batch.len() != expected_values || out_roots.len() != batch_size as usize {
606            return Err(XlogError::Compilation(format!(
607                "GpuCircuitCache root copy expects values len {} and roots len {}, got {} and {}",
608                expected_values,
609                batch_size,
610                values_batch.len(),
611                out_roots.len()
612            )));
613        }
614
615        let device = self.provider.device().inner();
616        let copy_root = device
617            .get_func(
618                xlog_cuda::CIRCUIT_MODULE,
619                xlog_cuda::circuit_kernels::XGCF_COPY_ROOT_CACHED_META_BATCHED,
620            )
621            .ok_or_else(|| {
622                XlogError::Kernel("xgcf_copy_root_cached_meta_batched not found".to_string())
623            })?;
624        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
625        unsafe {
626            copy_root.clone().launch(
627                LaunchConfig {
628                    grid_dim: (batch_size, 1, 1),
629                    block_dim: (1, 1, 1),
630                    shared_mem_bytes: 0,
631                },
632                (
633                    handle.slot_device(),
634                    self.node_cap,
635                    &self.meta_root,
636                    values_batch,
637                    node_stride,
638                    out_roots,
639                    batch_size,
640                ),
641            )
642        }
643        .map_err(|e| {
644            XlogError::Kernel(format!("xgcf_copy_root_cached_meta_batched failed: {}", e))
645        })?;
646        Ok(())
647    }
648
649    pub fn new(provider: &Arc<CudaKernelProvider>, config: GpuCircuitCacheConfig) -> Result<Self> {
650        if config.num_slots == 0 {
651            return Err(XlogError::Compilation(
652                "GpuCircuitCache requires num_slots > 0".to_string(),
653            ));
654        }
655        if config.table_size == 0 {
656            return Err(XlogError::Compilation(
657                "GpuCircuitCache requires table_size > 0".to_string(),
658            ));
659        }
660        if config.table_size < config.num_slots {
661            return Err(XlogError::Compilation(format!(
662                "GpuCircuitCache table_size {} < num_slots {}",
663                config.table_size, config.num_slots
664            )));
665        }
666        if config.node_cap == 0
667            || config.edge_cap == 0
668            || config.level_cap == 0
669            || config.var_cap == 0
670        {
671            return Err(XlogError::Compilation(
672                "GpuCircuitCache requires non-zero caps".to_string(),
673            ));
674        }
675
676        let memory = provider.memory();
677        let device = provider.device().inner();
678
679        let table_len = usize::try_from(config.table_size).map_err(|_| {
680            XlogError::Compilation("GpuCircuitCache table_size overflow".to_string())
681        })?;
682        let slot_len = usize::try_from(config.num_slots).map_err(|_| {
683            XlogError::Compilation("GpuCircuitCache num_slots overflow".to_string())
684        })?;
685
686        let node_cap = usize::try_from(config.node_cap)
687            .map_err(|_| XlogError::Compilation("GpuCircuitCache node_cap overflow".to_string()))?;
688        let edge_cap = usize::try_from(config.edge_cap)
689            .map_err(|_| XlogError::Compilation("GpuCircuitCache edge_cap overflow".to_string()))?;
690        let level_cap = usize::try_from(config.level_cap).map_err(|_| {
691            XlogError::Compilation("GpuCircuitCache level_cap overflow".to_string())
692        })?;
693        let var_cap = usize::try_from(config.var_cap)
694            .map_err(|_| XlogError::Compilation("GpuCircuitCache var_cap overflow".to_string()))?;
695
696        let node_slots = slot_len.checked_mul(node_cap).ok_or_else(|| {
697            XlogError::Compilation("GpuCircuitCache node slots overflow".to_string())
698        })?;
699        let edge_slots = slot_len.checked_mul(edge_cap).ok_or_else(|| {
700            XlogError::Compilation("GpuCircuitCache edge slots overflow".to_string())
701        })?;
702        let var_slots = slot_len.checked_mul(var_cap + 1).ok_or_else(|| {
703            XlogError::Compilation("GpuCircuitCache var slots overflow".to_string())
704        })?;
705        let node_offsets = slot_len.checked_mul(node_cap + 1).ok_or_else(|| {
706            XlogError::Compilation("GpuCircuitCache offset slots overflow".to_string())
707        })?;
708        let level_offsets = slot_len.checked_mul(level_cap + 1).ok_or_else(|| {
709            XlogError::Compilation("GpuCircuitCache level offsets overflow".to_string())
710        })?;
711
712        let mut keys = memory.alloc::<u64>(table_len)?;
713        let mut slots = memory.alloc::<u32>(table_len)?;
714        let mut state = memory.alloc::<u32>(table_len)?;
715        let mut last_used = memory.alloc::<u64>(table_len)?;
716        let mut slot_states = memory.alloc::<u32>(slot_len)?;
717        let mut clock = memory.alloc::<u64>(1)?;
718
719        let mut node_type = memory.alloc::<u8>(node_slots)?;
720        let mut child_offsets = memory.alloc::<u32>(node_offsets)?;
721        let mut child_indices = memory.alloc::<u32>(edge_slots)?;
722        let mut lit = memory.alloc::<i32>(node_slots)?;
723        let mut decision_var = memory.alloc::<u32>(node_slots)?;
724        let mut decision_child_false = memory.alloc::<u32>(node_slots)?;
725        let mut decision_child_true = memory.alloc::<u32>(node_slots)?;
726        let mut level_nodes = memory.alloc::<u32>(node_slots)?;
727        let mut level_offsets = memory.alloc::<u32>(level_offsets)?;
728
729        let mut var_log_true = memory.alloc::<f64>(var_slots)?;
730        let mut var_log_false = memory.alloc::<f64>(var_slots)?;
731        let mut values = memory.alloc::<f64>(node_slots)?;
732        let mut adj = memory.alloc::<f64>(node_slots)?;
733        let mut grad_true = memory.alloc::<f64>(var_slots)?;
734        let mut grad_false = memory.alloc::<f64>(var_slots)?;
735        let mut free_var_mask = memory.alloc::<u8>(var_slots)?;
736        let mut meta_num_nodes = memory.alloc::<u32>(slot_len)?;
737        let mut meta_num_levels = memory.alloc::<u32>(slot_len)?;
738        let mut meta_root = memory.alloc::<u32>(slot_len)?;
739        let mut meta_max_var = memory.alloc::<u32>(slot_len)?;
740        let mut always_on = memory.alloc::<u32>(1)?;
741        let zero_len = node_cap.max(var_cap + 1);
742        let mut zero_f64 = memory.alloc::<f64>(zero_len)?;
743        let mut one_f64 = memory.alloc::<f64>(1)?;
744
745        device
746            .memset_zeros(&mut keys)
747            .map_err(|e| XlogError::Kernel(format!("GpuCircuitCache zero keys failed: {}", e)))?;
748        device
749            .memset_zeros(&mut slots)
750            .map_err(|e| XlogError::Kernel(format!("GpuCircuitCache zero slots failed: {}", e)))?;
751        device
752            .memset_zeros(&mut state)
753            .map_err(|e| XlogError::Kernel(format!("GpuCircuitCache zero state failed: {}", e)))?;
754        device.memset_zeros(&mut last_used).map_err(|e| {
755            XlogError::Kernel(format!("GpuCircuitCache zero last_used failed: {}", e))
756        })?;
757        device.memset_zeros(&mut slot_states).map_err(|e| {
758            XlogError::Kernel(format!("GpuCircuitCache zero slot_states failed: {}", e))
759        })?;
760        device
761            .memset_zeros(&mut clock)
762            .map_err(|e| XlogError::Kernel(format!("GpuCircuitCache zero clock failed: {}", e)))?;
763
764        device.memset_zeros(&mut node_type).map_err(|e| {
765            XlogError::Kernel(format!("GpuCircuitCache zero node_type failed: {}", e))
766        })?;
767        device.memset_zeros(&mut child_offsets).map_err(|e| {
768            XlogError::Kernel(format!("GpuCircuitCache zero child_offsets failed: {}", e))
769        })?;
770        device.memset_zeros(&mut child_indices).map_err(|e| {
771            XlogError::Kernel(format!("GpuCircuitCache zero child_indices failed: {}", e))
772        })?;
773        device
774            .memset_zeros(&mut lit)
775            .map_err(|e| XlogError::Kernel(format!("GpuCircuitCache zero lit failed: {}", e)))?;
776        device.memset_zeros(&mut decision_var).map_err(|e| {
777            XlogError::Kernel(format!("GpuCircuitCache zero decision_var failed: {}", e))
778        })?;
779        device
780            .memset_zeros(&mut decision_child_false)
781            .map_err(|e| {
782                XlogError::Kernel(format!(
783                    "GpuCircuitCache zero decision_child_false failed: {}",
784                    e
785                ))
786            })?;
787        device.memset_zeros(&mut decision_child_true).map_err(|e| {
788            XlogError::Kernel(format!(
789                "GpuCircuitCache zero decision_child_true failed: {}",
790                e
791            ))
792        })?;
793        device.memset_zeros(&mut level_nodes).map_err(|e| {
794            XlogError::Kernel(format!("GpuCircuitCache zero level_nodes failed: {}", e))
795        })?;
796        device.memset_zeros(&mut level_offsets).map_err(|e| {
797            XlogError::Kernel(format!("GpuCircuitCache zero level_offsets failed: {}", e))
798        })?;
799        device.memset_zeros(&mut var_log_true).map_err(|e| {
800            XlogError::Kernel(format!("GpuCircuitCache zero var_log_true failed: {}", e))
801        })?;
802        device.memset_zeros(&mut var_log_false).map_err(|e| {
803            XlogError::Kernel(format!("GpuCircuitCache zero var_log_false failed: {}", e))
804        })?;
805        device
806            .memset_zeros(&mut values)
807            .map_err(|e| XlogError::Kernel(format!("GpuCircuitCache zero values failed: {}", e)))?;
808        device
809            .memset_zeros(&mut adj)
810            .map_err(|e| XlogError::Kernel(format!("GpuCircuitCache zero adj failed: {}", e)))?;
811        device.memset_zeros(&mut grad_true).map_err(|e| {
812            XlogError::Kernel(format!("GpuCircuitCache zero grad_true failed: {}", e))
813        })?;
814        device.memset_zeros(&mut grad_false).map_err(|e| {
815            XlogError::Kernel(format!("GpuCircuitCache zero grad_false failed: {}", e))
816        })?;
817        device.memset_zeros(&mut free_var_mask).map_err(|e| {
818            XlogError::Kernel(format!("GpuCircuitCache zero free_var_mask failed: {}", e))
819        })?;
820        device.memset_zeros(&mut meta_num_nodes).map_err(|e| {
821            XlogError::Kernel(format!("GpuCircuitCache zero meta_num_nodes failed: {}", e))
822        })?;
823        device.memset_zeros(&mut meta_num_levels).map_err(|e| {
824            XlogError::Kernel(format!(
825                "GpuCircuitCache zero meta_num_levels failed: {}",
826                e
827            ))
828        })?;
829        device.memset_zeros(&mut meta_root).map_err(|e| {
830            XlogError::Kernel(format!("GpuCircuitCache zero meta_root failed: {}", e))
831        })?;
832        device.memset_zeros(&mut meta_max_var).map_err(|e| {
833            XlogError::Kernel(format!("GpuCircuitCache zero meta_max_var failed: {}", e))
834        })?;
835        device.memset_zeros(&mut zero_f64).map_err(|e| {
836            XlogError::Kernel(format!("GpuCircuitCache zero zero_f64 failed: {}", e))
837        })?;
838        provider
839            .htod_launch_metadata_sync_copy_into(&[1u32], &mut always_on)
840            .map_err(|e| {
841                XlogError::Kernel(format!("GpuCircuitCache init always_on failed: {}", e))
842            })?;
843        provider
844            .htod_launch_metadata_sync_copy_into(&[1.0f64], &mut one_f64)
845            .map_err(|e| {
846                XlogError::Kernel(format!("GpuCircuitCache init one_f64 failed: {}", e))
847            })?;
848
849        Ok(Self {
850            provider: provider.clone(),
851            table_size: config.table_size,
852            num_slots: config.num_slots,
853            node_cap: config.node_cap,
854            edge_cap: config.edge_cap,
855            level_cap: config.level_cap,
856            var_cap: config.var_cap,
857            keys,
858            slots,
859            state,
860            last_used,
861            slot_states,
862            clock,
863            node_type,
864            child_offsets,
865            child_indices,
866            lit,
867            decision_var,
868            decision_child_false,
869            decision_child_true,
870            level_nodes,
871            level_offsets,
872            var_log_true,
873            var_log_false,
874            values,
875            adj,
876            grad_true,
877            grad_false,
878            meta_num_nodes,
879            meta_num_levels,
880            meta_root,
881            meta_max_var,
882            always_on,
883            zero_f64,
884            one_f64,
885            free_var_mask,
886            has_free_var_mask: vec![false; config.num_slots as usize],
887        })
888    }
889
890    pub fn lookup_or_insert(&mut self, key: u64) -> Result<GpuCacheLookup> {
891        let memory = self.provider.memory();
892        let mut key_device = memory.alloc::<u64>(1)?;
893        self.provider
894            .htod_launch_metadata_sync_copy_into(&[key], &mut key_device)
895            .map_err(|e| XlogError::Kernel(format!("cache upload key failed: {}", e)))?;
896        self.lookup_or_insert_device(&key_device)
897    }
898
899    pub(crate) fn lookup_or_insert_device(
900        &mut self,
901        key_device: &TrackedCudaSlice<u64>,
902    ) -> Result<GpuCacheLookup> {
903        let memory = self.provider.memory();
904        let mut out_slot = memory.alloc::<u32>(1)?;
905        let mut out_compile_needed = memory.alloc::<u32>(1)?;
906
907        let func = self
908            .provider
909            .device()
910            .inner()
911            .get_func(CACHE_MODULE, cache_kernels::CACHE_LOOKUP_OR_INSERT)
912            .ok_or_else(|| {
913                XlogError::Kernel("cache_lookup_or_insert kernel not found".to_string())
914            })?;
915
916        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
917        unsafe {
918            func.clone().launch(
919                LaunchConfig {
920                    grid_dim: (1, 1, 1),
921                    block_dim: (1, 1, 1),
922                    shared_mem_bytes: 0,
923                },
924                (
925                    key_device,
926                    self.table_size,
927                    self.num_slots,
928                    &mut self.keys,
929                    &mut self.slots,
930                    &mut self.state,
931                    &mut self.last_used,
932                    &mut self.slot_states,
933                    &mut self.clock,
934                    &mut out_slot,
935                    &mut out_compile_needed,
936                ),
937            )
938        }
939        .map_err(|e| XlogError::Kernel(format!("cache_lookup_or_insert failed: {}", e)))?;
940        // No device synchronize: slot and compile_needed stay device-resident; same-stream ordering suffices.
941        Ok(GpuCacheLookup {
942            provider: self.provider.clone(),
943            slot: out_slot,
944            compile_needed: out_compile_needed,
945        })
946    }
947
948    pub fn claim_slot(&mut self, key: u64) -> Result<GpuCircuitCacheHandle> {
949        let lookup = self.lookup_or_insert(key)?;
950        lookup.into_handle()
951    }
952
953    pub fn store_from_xgcf(
954        &mut self,
955        handle: &mut GpuCircuitCacheHandle,
956        xgcf: &GpuXgcf,
957    ) -> Result<()> {
958        // Download the actual node/edge counts from device-resident metadata.
959        // xgcf.num_nodes() / num_edges() return the CAPACITY (node_cap / edge_cap),
960        // not the actual count produced by d4. Using the capacity would store garbage
961        // data beyond the actual circuit, corrupting disk cache artifacts.
962        let device = self.provider.device().inner();
963        let num_nodes_host: Vec<u32> = device
964            .dtoh_sync_copy(xgcf.num_nodes_device())
965            .map_err(|e| XlogError::Kernel(format!("dtoh meta_num_nodes: {}", e)))?;
966        let num_nodes = num_nodes_host[0];
967        if num_nodes == 0 {
968            return Err(XlogError::Compilation(
969                "GpuCircuitCache store: num_nodes must be > 0".to_string(),
970            ));
971        }
972        if num_nodes > self.node_cap {
973            return Err(XlogError::Compilation(format!(
974                "GpuCircuitCache store: num_nodes {} exceeds node_cap {}",
975                num_nodes, self.node_cap
976            )));
977        }
978
979        let num_edges_host: Vec<u32> = device
980            .dtoh_sync_copy(xgcf.num_edges_device())
981            .map_err(|e| XlogError::Kernel(format!("dtoh meta_num_edges: {}", e)))?;
982        let num_edges = num_edges_host[0];
983        if num_edges > self.edge_cap {
984            return Err(XlogError::Compilation(format!(
985                "GpuCircuitCache store: num_edges {} exceeds edge_cap {}",
986                num_edges, self.edge_cap
987            )));
988        }
989
990        let num_levels = xgcf.num_levels();
991        if num_levels == 0 {
992            return Err(XlogError::Compilation(
993                "GpuCircuitCache store: num_levels must be > 0".to_string(),
994            ));
995        }
996        if num_levels > self.level_cap {
997            return Err(XlogError::Compilation(format!(
998                "GpuCircuitCache store: num_levels {} exceeds level_cap {}",
999                num_levels, self.level_cap
1000            )));
1001        }
1002
1003        let root = xgcf.root();
1004        if root >= num_nodes {
1005            return Err(XlogError::Compilation(format!(
1006                "GpuCircuitCache store: root {} out of bounds (num_nodes={})",
1007                root, num_nodes
1008            )));
1009        }
1010
1011        let max_var = xgcf.max_var();
1012        if max_var > self.var_cap {
1013            return Err(XlogError::Compilation(format!(
1014                "GpuCircuitCache store: max_var {} exceeds var_cap {}",
1015                max_var, self.var_cap
1016            )));
1017        }
1018
1019        let expected_child_offsets = (num_nodes as usize) + 1;
1020        if xgcf.child_offsets().len() < expected_child_offsets {
1021            return Err(XlogError::Compilation(format!(
1022                "GpuCircuitCache store: child_offsets len {} < num_nodes+1 {}",
1023                xgcf.child_offsets().len(),
1024                expected_child_offsets
1025            )));
1026        }
1027        if xgcf.level_nodes().len() < num_nodes as usize {
1028            return Err(XlogError::Compilation(format!(
1029                "GpuCircuitCache store: level_nodes len {} < num_nodes {}",
1030                xgcf.level_nodes().len(),
1031                num_nodes
1032            )));
1033        }
1034        let expected_level_offsets = (num_levels as usize) + 1;
1035        if xgcf.level_offsets().len() != expected_level_offsets {
1036            return Err(XlogError::Compilation(format!(
1037                "GpuCircuitCache store: level_offsets len {} != num_levels+1 {}",
1038                xgcf.level_offsets().len(),
1039                expected_level_offsets
1040            )));
1041        }
1042
1043        handle.num_nodes = num_nodes;
1044        handle.num_levels = num_levels;
1045        handle.root = root;
1046        handle.max_var = max_var;
1047
1048        let store_u8 = device
1049            .get_func(CACHE_MODULE, cache_kernels::CACHE_STORE_U8)
1050            .ok_or_else(|| XlogError::Kernel("cache_store_u8 kernel not found".to_string()))?;
1051        let store_u32 = device
1052            .get_func(CACHE_MODULE, cache_kernels::CACHE_STORE_U32)
1053            .ok_or_else(|| XlogError::Kernel("cache_store_u32 kernel not found".to_string()))?;
1054        let store_i32 = device
1055            .get_func(CACHE_MODULE, cache_kernels::CACHE_STORE_I32)
1056            .ok_or_else(|| XlogError::Kernel("cache_store_i32 kernel not found".to_string()))?;
1057        let store_f64 = device
1058            .get_func(CACHE_MODULE, cache_kernels::CACHE_STORE_F64)
1059            .ok_or_else(|| XlogError::Kernel("cache_store_f64 kernel not found".to_string()))?;
1060        let store_meta = device
1061            .get_func(CACHE_MODULE, cache_kernels::CACHE_STORE_META)
1062            .ok_or_else(|| XlogError::Kernel("cache_store_meta kernel not found".to_string()))?;
1063
1064        let block_dim = 256u32;
1065
1066        let node_stride = self.node_cap;
1067        let offset_stride = self.node_cap.checked_add(1).ok_or_else(|| {
1068            XlogError::Compilation("GpuCircuitCache store: node_cap overflow".to_string())
1069        })?;
1070        let level_offset_stride = self.level_cap.checked_add(1).ok_or_else(|| {
1071            XlogError::Compilation("GpuCircuitCache store: level_cap overflow".to_string())
1072        })?;
1073        let var_stride = self.var_cap.checked_add(1).ok_or_else(|| {
1074            XlogError::Compilation("GpuCircuitCache store: var_cap overflow".to_string())
1075        })?;
1076
1077        let num_nodes_plus1 = num_nodes.checked_add(1).ok_or_else(|| {
1078            XlogError::Compilation("GpuCircuitCache store: num_nodes overflow".to_string())
1079        })?;
1080        let num_levels_plus1 = num_levels.checked_add(1).ok_or_else(|| {
1081            XlogError::Compilation("GpuCircuitCache store: num_levels overflow".to_string())
1082        })?;
1083        let weights_len = max_var.checked_add(1).ok_or_else(|| {
1084            XlogError::Compilation("GpuCircuitCache store: max_var overflow".to_string())
1085        })?;
1086
1087        let grid_nodes =
1088            cache_grid_dim_for_u32_count("GpuCircuitCache store node_type", num_nodes, block_dim)?;
1089        if grid_nodes != 0 {
1090            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1091            unsafe {
1092                store_u8.clone().launch(
1093                    LaunchConfig {
1094                        grid_dim: (grid_nodes, 1, 1),
1095                        block_dim: (block_dim, 1, 1),
1096                        shared_mem_bytes: 0,
1097                    },
1098                    (
1099                        handle.slot_device(),
1100                        handle.compile_needed_device(),
1101                        node_stride,
1102                        xgcf.node_type(),
1103                        &mut self.node_type,
1104                        num_nodes,
1105                    ),
1106                )
1107            }
1108            .map_err(|e| XlogError::Kernel(format!("cache_store_u8 failed: {}", e)))?;
1109        }
1110
1111        let grid_offsets = cache_grid_dim_for_u32_count(
1112            "GpuCircuitCache store child_offsets",
1113            num_nodes_plus1,
1114            block_dim,
1115        )?;
1116        if grid_offsets != 0 {
1117            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1118            unsafe {
1119                store_u32.clone().launch(
1120                    LaunchConfig {
1121                        grid_dim: (grid_offsets, 1, 1),
1122                        block_dim: (block_dim, 1, 1),
1123                        shared_mem_bytes: 0,
1124                    },
1125                    (
1126                        handle.slot_device(),
1127                        handle.compile_needed_device(),
1128                        offset_stride,
1129                        xgcf.child_offsets(),
1130                        &mut self.child_offsets,
1131                        num_nodes_plus1,
1132                    ),
1133                )
1134            }
1135            .map_err(|e| XlogError::Kernel(format!("cache_store_child_offsets failed: {}", e)))?;
1136        }
1137
1138        let grid_edges = cache_grid_dim_for_u32_count(
1139            "GpuCircuitCache store child_indices",
1140            num_edges,
1141            block_dim,
1142        )?;
1143        if grid_edges != 0 {
1144            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1145            unsafe {
1146                store_u32.clone().launch(
1147                    LaunchConfig {
1148                        grid_dim: (grid_edges, 1, 1),
1149                        block_dim: (block_dim, 1, 1),
1150                        shared_mem_bytes: 0,
1151                    },
1152                    (
1153                        handle.slot_device(),
1154                        handle.compile_needed_device(),
1155                        self.edge_cap,
1156                        xgcf.child_indices(),
1157                        &mut self.child_indices,
1158                        num_edges,
1159                    ),
1160                )
1161            }
1162            .map_err(|e| XlogError::Kernel(format!("cache_store_child_indices failed: {}", e)))?;
1163        }
1164
1165        if grid_nodes != 0 {
1166            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1167            unsafe {
1168                store_i32.clone().launch(
1169                    LaunchConfig {
1170                        grid_dim: (grid_nodes, 1, 1),
1171                        block_dim: (block_dim, 1, 1),
1172                        shared_mem_bytes: 0,
1173                    },
1174                    (
1175                        handle.slot_device(),
1176                        handle.compile_needed_device(),
1177                        node_stride,
1178                        xgcf.lit(),
1179                        &mut self.lit,
1180                        num_nodes,
1181                    ),
1182                )
1183            }
1184            .map_err(|e| XlogError::Kernel(format!("cache_store_lit failed: {}", e)))?;
1185
1186            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1187            unsafe {
1188                store_u32.clone().launch(
1189                    LaunchConfig {
1190                        grid_dim: (grid_nodes, 1, 1),
1191                        block_dim: (block_dim, 1, 1),
1192                        shared_mem_bytes: 0,
1193                    },
1194                    (
1195                        handle.slot_device(),
1196                        handle.compile_needed_device(),
1197                        node_stride,
1198                        xgcf.decision_var(),
1199                        &mut self.decision_var,
1200                        num_nodes,
1201                    ),
1202                )
1203            }
1204            .map_err(|e| XlogError::Kernel(format!("cache_store_decision_var failed: {}", e)))?;
1205
1206            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1207            unsafe {
1208                store_u32.clone().launch(
1209                    LaunchConfig {
1210                        grid_dim: (grid_nodes, 1, 1),
1211                        block_dim: (block_dim, 1, 1),
1212                        shared_mem_bytes: 0,
1213                    },
1214                    (
1215                        handle.slot_device(),
1216                        handle.compile_needed_device(),
1217                        node_stride,
1218                        xgcf.decision_child_false(),
1219                        &mut self.decision_child_false,
1220                        num_nodes,
1221                    ),
1222                )
1223            }
1224            .map_err(|e| {
1225                XlogError::Kernel(format!("cache_store_decision_child_false failed: {}", e))
1226            })?;
1227
1228            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1229            unsafe {
1230                store_u32.clone().launch(
1231                    LaunchConfig {
1232                        grid_dim: (grid_nodes, 1, 1),
1233                        block_dim: (block_dim, 1, 1),
1234                        shared_mem_bytes: 0,
1235                    },
1236                    (
1237                        handle.slot_device(),
1238                        handle.compile_needed_device(),
1239                        node_stride,
1240                        xgcf.decision_child_true(),
1241                        &mut self.decision_child_true,
1242                        num_nodes,
1243                    ),
1244                )
1245            }
1246            .map_err(|e| {
1247                XlogError::Kernel(format!("cache_store_decision_child_true failed: {}", e))
1248            })?;
1249
1250            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1251            unsafe {
1252                store_u32.clone().launch(
1253                    LaunchConfig {
1254                        grid_dim: (grid_nodes, 1, 1),
1255                        block_dim: (block_dim, 1, 1),
1256                        shared_mem_bytes: 0,
1257                    },
1258                    (
1259                        handle.slot_device(),
1260                        handle.compile_needed_device(),
1261                        node_stride,
1262                        xgcf.level_nodes(),
1263                        &mut self.level_nodes,
1264                        num_nodes,
1265                    ),
1266                )
1267            }
1268            .map_err(|e| XlogError::Kernel(format!("cache_store_level_nodes failed: {}", e)))?;
1269        }
1270
1271        let grid_levels = cache_grid_dim_for_u32_count(
1272            "GpuCircuitCache store level_offsets",
1273            num_levels_plus1,
1274            block_dim,
1275        )?;
1276        if grid_levels != 0 {
1277            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1278            unsafe {
1279                store_u32.clone().launch(
1280                    LaunchConfig {
1281                        grid_dim: (grid_levels, 1, 1),
1282                        block_dim: (block_dim, 1, 1),
1283                        shared_mem_bytes: 0,
1284                    },
1285                    (
1286                        handle.slot_device(),
1287                        handle.compile_needed_device(),
1288                        level_offset_stride,
1289                        xgcf.level_offsets(),
1290                        &mut self.level_offsets,
1291                        num_levels_plus1,
1292                    ),
1293                )
1294            }
1295            .map_err(|e| XlogError::Kernel(format!("cache_store_level_offsets failed: {}", e)))?;
1296        }
1297
1298        let grid_weights = cache_grid_dim_for_u32_count(
1299            "GpuCircuitCache store free_var_mask",
1300            weights_len,
1301            block_dim,
1302        )?;
1303        if grid_weights != 0 {
1304            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1305            unsafe {
1306                store_f64.clone().launch(
1307                    LaunchConfig {
1308                        grid_dim: (grid_weights, 1, 1),
1309                        block_dim: (block_dim, 1, 1),
1310                        shared_mem_bytes: 0,
1311                    },
1312                    (
1313                        handle.slot_device(),
1314                        handle.compile_needed_device(),
1315                        var_stride,
1316                        xgcf.var_log_true(),
1317                        &mut self.var_log_true,
1318                        weights_len,
1319                    ),
1320                )
1321            }
1322            .map_err(|e| XlogError::Kernel(format!("cache_store_var_log_true failed: {}", e)))?;
1323
1324            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1325            unsafe {
1326                store_f64.clone().launch(
1327                    LaunchConfig {
1328                        grid_dim: (grid_weights, 1, 1),
1329                        block_dim: (block_dim, 1, 1),
1330                        shared_mem_bytes: 0,
1331                    },
1332                    (
1333                        handle.slot_device(),
1334                        handle.compile_needed_device(),
1335                        var_stride,
1336                        xgcf.var_log_false(),
1337                        &mut self.var_log_false,
1338                        weights_len,
1339                    ),
1340                )
1341            }
1342            .map_err(|e| XlogError::Kernel(format!("cache_store_var_log_false failed: {}", e)))?;
1343        }
1344
1345        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1346        unsafe {
1347            store_meta.clone().launch(
1348                LaunchConfig {
1349                    grid_dim: (1, 1, 1),
1350                    block_dim: (1, 1, 1),
1351                    shared_mem_bytes: 0,
1352                },
1353                (
1354                    handle.slot_device(),
1355                    handle.compile_needed_device(),
1356                    self.num_slots,
1357                    num_nodes,
1358                    num_levels,
1359                    root,
1360                    max_var,
1361                    &mut self.meta_num_nodes,
1362                    &mut self.meta_num_levels,
1363                    &mut self.meta_root,
1364                    &mut self.meta_max_var,
1365                ),
1366            )
1367        }
1368        .map_err(|e| XlogError::Kernel(format!("cache_store_meta failed: {}", e)))?;
1369
1370        // No device synchronize needed: all stores are GPU-to-GPU on the same stream.
1371        // Same-stream ordering guarantees subsequent kernels see the stored data.
1372        Ok(())
1373    }
1374
1375    pub fn store_weights(
1376        &mut self,
1377        handle: &GpuCircuitCacheHandle,
1378        weights_true: &TrackedCudaSlice<f64>,
1379        weights_false: &TrackedCudaSlice<f64>,
1380    ) -> Result<()> {
1381        let weights_len = handle.max_var.checked_add(1).ok_or_else(|| {
1382            XlogError::Compilation("GpuCircuitCache store_weights max_var overflow".to_string())
1383        })?;
1384        let weights_len_usize = usize::try_from(weights_len).map_err(|_| {
1385            XlogError::Compilation("GpuCircuitCache store_weights len overflow".to_string())
1386        })?;
1387        if weights_true.len() < weights_len_usize || weights_false.len() < weights_len_usize {
1388            return Err(XlogError::Compilation(format!(
1389                "GpuCircuitCache store_weights requires weights len >= {}, got true={} false={}",
1390                weights_len,
1391                weights_true.len(),
1392                weights_false.len()
1393            )));
1394        }
1395
1396        let device = self.provider.device().inner();
1397        let store_f64 = device
1398            .get_func(CACHE_MODULE, cache_kernels::CACHE_STORE_F64)
1399            .ok_or_else(|| XlogError::Kernel("cache_store_f64 kernel not found".to_string()))?;
1400
1401        let block_dim = 256u32;
1402        let grid_dim = if weights_len == 0 {
1403            0
1404        } else {
1405            weights_len.div_ceil(block_dim)
1406        };
1407        if grid_dim != 0 {
1408            let var_stride = self.var_cap.checked_add(1).ok_or_else(|| {
1409                XlogError::Compilation("GpuCircuitCache store_weights var_cap overflow".to_string())
1410            })?;
1411            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1412            unsafe {
1413                store_f64.clone().launch(
1414                    LaunchConfig {
1415                        grid_dim: (grid_dim, 1, 1),
1416                        block_dim: (block_dim, 1, 1),
1417                        shared_mem_bytes: 0,
1418                    },
1419                    (
1420                        handle.slot_device(),
1421                        handle.compile_needed_device(),
1422                        var_stride,
1423                        weights_true,
1424                        &mut self.var_log_true,
1425                        weights_len,
1426                    ),
1427                )
1428            }
1429            .map_err(|e| XlogError::Kernel(format!("cache_store_weights_true failed: {}", e)))?;
1430
1431            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1432            unsafe {
1433                store_f64.clone().launch(
1434                    LaunchConfig {
1435                        grid_dim: (grid_dim, 1, 1),
1436                        block_dim: (block_dim, 1, 1),
1437                        shared_mem_bytes: 0,
1438                    },
1439                    (
1440                        handle.slot_device(),
1441                        handle.compile_needed_device(),
1442                        var_stride,
1443                        weights_false,
1444                        &mut self.var_log_false,
1445                        weights_len,
1446                    ),
1447                )
1448            }
1449            .map_err(|e| XlogError::Kernel(format!("cache_store_weights_false failed: {}", e)))?;
1450        }
1451
1452        // No device synchronize: same-stream ordering guarantees visibility.
1453        Ok(())
1454    }
1455
1456    pub fn overwrite_weights(
1457        &mut self,
1458        handle: &GpuCircuitCacheHandle,
1459        weights_true: &TrackedCudaSlice<f64>,
1460        weights_false: &TrackedCudaSlice<f64>,
1461    ) -> Result<()> {
1462        let weights_len = handle.max_var.checked_add(1).ok_or_else(|| {
1463            XlogError::Compilation("GpuCircuitCache overwrite_weights max_var overflow".to_string())
1464        })?;
1465        let weights_len_usize = usize::try_from(weights_len).map_err(|_| {
1466            XlogError::Compilation("GpuCircuitCache overwrite_weights len overflow".to_string())
1467        })?;
1468        if weights_true.len() < weights_len_usize || weights_false.len() < weights_len_usize {
1469            return Err(XlogError::Compilation(format!(
1470                "GpuCircuitCache overwrite_weights requires weights len >= {}, got true={} false={}",
1471                weights_len,
1472                weights_true.len(),
1473                weights_false.len()
1474            )));
1475        }
1476
1477        let device = self.provider.device().inner();
1478        let store_f64 = device
1479            .get_func(CACHE_MODULE, cache_kernels::CACHE_STORE_F64)
1480            .ok_or_else(|| XlogError::Kernel("cache_store_f64 kernel not found".to_string()))?;
1481
1482        let block_dim = 256u32;
1483        let grid_dim = if weights_len == 0 {
1484            0
1485        } else {
1486            weights_len.div_ceil(block_dim)
1487        };
1488        if grid_dim != 0 {
1489            let var_stride = self.var_cap.checked_add(1).ok_or_else(|| {
1490                XlogError::Compilation(
1491                    "GpuCircuitCache overwrite_weights var_cap overflow".to_string(),
1492                )
1493            })?;
1494            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1495            unsafe {
1496                store_f64.clone().launch(
1497                    LaunchConfig {
1498                        grid_dim: (grid_dim, 1, 1),
1499                        block_dim: (block_dim, 1, 1),
1500                        shared_mem_bytes: 0,
1501                    },
1502                    (
1503                        handle.slot_device(),
1504                        &self.always_on,
1505                        var_stride,
1506                        weights_true,
1507                        &mut self.var_log_true,
1508                        weights_len,
1509                    ),
1510                )
1511            }
1512            .map_err(|e| {
1513                XlogError::Kernel(format!("cache_overwrite_weights_true failed: {}", e))
1514            })?;
1515
1516            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1517            unsafe {
1518                store_f64.clone().launch(
1519                    LaunchConfig {
1520                        grid_dim: (grid_dim, 1, 1),
1521                        block_dim: (block_dim, 1, 1),
1522                        shared_mem_bytes: 0,
1523                    },
1524                    (
1525                        handle.slot_device(),
1526                        &self.always_on,
1527                        var_stride,
1528                        weights_false,
1529                        &mut self.var_log_false,
1530                        weights_len,
1531                    ),
1532                )
1533            }
1534            .map_err(|e| {
1535                XlogError::Kernel(format!("cache_overwrite_weights_false failed: {}", e))
1536            })?;
1537        }
1538
1539        // No device synchronize: same-stream ordering guarantees visibility.
1540        Ok(())
1541    }
1542
1543    pub fn store_free_var_mask(
1544        &mut self,
1545        handle: &GpuCircuitCacheHandle,
1546        mask: &TrackedCudaSlice<u8>,
1547    ) -> Result<()> {
1548        let mask_len = u32::try_from(mask.len()).map_err(|_| {
1549            XlogError::Compilation("GpuCircuitCache free_var_mask len overflow".to_string())
1550        })?;
1551        let expected_len = handle.max_var.checked_add(1).ok_or_else(|| {
1552            XlogError::Compilation("GpuCircuitCache free_var_mask max_var overflow".to_string())
1553        })?;
1554        if mask_len != expected_len {
1555            return Err(XlogError::Compilation(format!(
1556                "GpuCircuitCache free_var_mask len {} != expected {}",
1557                mask_len, expected_len
1558            )));
1559        }
1560        let var_stride = self.var_cap.checked_add(1).ok_or_else(|| {
1561            XlogError::Compilation("GpuCircuitCache free_var_mask var_cap overflow".to_string())
1562        })?;
1563        if expected_len > var_stride {
1564            return Err(XlogError::Compilation(format!(
1565                "GpuCircuitCache free_var_mask len {} exceeds var_cap+1 {}",
1566                expected_len, var_stride
1567            )));
1568        }
1569
1570        let device = self.provider.device().inner();
1571        let store_u8 = device
1572            .get_func(CACHE_MODULE, cache_kernels::CACHE_STORE_U8)
1573            .ok_or_else(|| XlogError::Kernel("cache_store_u8 kernel not found".to_string()))?;
1574
1575        let block_dim = 256u32;
1576        let grid_dim = mask_len.div_ceil(block_dim);
1577        if grid_dim == 0 {
1578            return Ok(());
1579        }
1580
1581        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1582        unsafe {
1583            store_u8.clone().launch(
1584                LaunchConfig {
1585                    grid_dim: (grid_dim, 1, 1),
1586                    block_dim: (block_dim, 1, 1),
1587                    shared_mem_bytes: 0,
1588                },
1589                (
1590                    handle.slot_device(),
1591                    handle.compile_needed_device(),
1592                    var_stride,
1593                    mask,
1594                    &mut self.free_var_mask,
1595                    mask_len,
1596                ),
1597            )
1598        }
1599        .map_err(|e| XlogError::Kernel(format!("cache_store_free_var_mask failed: {}", e)))?;
1600
1601        // No device synchronize: same-stream ordering guarantees visibility.
1602        let slot_idx = handle.slot_index() as usize;
1603        debug_assert!(
1604            slot_idx < self.has_free_var_mask.len(),
1605            "slot_index {} exceeds num_slots {}",
1606            slot_idx,
1607            self.has_free_var_mask.len()
1608        );
1609        if slot_idx < self.has_free_var_mask.len() {
1610            self.has_free_var_mask[slot_idx] = true;
1611        }
1612        Ok(())
1613    }
1614
1615    /// Populate a cache slot from host-resident arrays loaded from the disk cache.
1616    ///
1617    /// This mirrors [`store_from_xgcf`] but takes a [`disk_cache::CircuitArtifact`]
1618    /// (host `Vec`s) instead of a device-resident `GpuXgcf`. Each host array is
1619    /// uploaded to a temporary device buffer and then stored into the slot via the
1620    /// same `cache_store_*` kernels.
1621    pub(crate) fn restore_from_host_arrays(
1622        &mut self,
1623        handle: &mut GpuCircuitCacheHandle,
1624        artifact: &disk_cache::CircuitArtifact,
1625    ) -> Result<()> {
1626        // -- Validate sizes against cache caps --
1627        let num_nodes = artifact.num_nodes;
1628        if num_nodes == 0 {
1629            return Err(XlogError::Compilation(
1630                "GpuCircuitCache restore: num_nodes must be > 0".to_string(),
1631            ));
1632        }
1633        if num_nodes > self.node_cap {
1634            return Err(XlogError::Compilation(format!(
1635                "GpuCircuitCache restore: num_nodes {} exceeds node_cap {}",
1636                num_nodes, self.node_cap
1637            )));
1638        }
1639
1640        let num_edges = artifact.num_edges;
1641        if num_edges > self.edge_cap {
1642            return Err(XlogError::Compilation(format!(
1643                "GpuCircuitCache restore: num_edges {} exceeds edge_cap {}",
1644                num_edges, self.edge_cap
1645            )));
1646        }
1647
1648        let num_levels = artifact.num_levels;
1649        if num_levels == 0 {
1650            return Err(XlogError::Compilation(
1651                "GpuCircuitCache restore: num_levels must be > 0".to_string(),
1652            ));
1653        }
1654        if num_levels > self.level_cap {
1655            return Err(XlogError::Compilation(format!(
1656                "GpuCircuitCache restore: num_levels {} exceeds level_cap {}",
1657                num_levels, self.level_cap
1658            )));
1659        }
1660
1661        let root = artifact.root;
1662        if root >= num_nodes {
1663            return Err(XlogError::Compilation(format!(
1664                "GpuCircuitCache restore: root {} out of bounds (num_nodes={})",
1665                root, num_nodes
1666            )));
1667        }
1668
1669        let max_var = artifact.max_var;
1670        if max_var > self.var_cap {
1671            return Err(XlogError::Compilation(format!(
1672                "GpuCircuitCache restore: max_var {} exceeds var_cap {}",
1673                max_var, self.var_cap
1674            )));
1675        }
1676
1677        let expected_child_offsets = (num_nodes as usize) + 1;
1678        if artifact.child_offsets.len() < expected_child_offsets {
1679            return Err(XlogError::Compilation(format!(
1680                "GpuCircuitCache restore: child_offsets len {} < num_nodes+1 {}",
1681                artifact.child_offsets.len(),
1682                expected_child_offsets
1683            )));
1684        }
1685        if artifact.level_nodes.len() < num_nodes as usize {
1686            return Err(XlogError::Compilation(format!(
1687                "GpuCircuitCache restore: level_nodes len {} < num_nodes {}",
1688                artifact.level_nodes.len(),
1689                num_nodes
1690            )));
1691        }
1692        let expected_level_offsets = (num_levels as usize) + 1;
1693        if artifact.level_offsets.len() != expected_level_offsets {
1694            return Err(XlogError::Compilation(format!(
1695                "GpuCircuitCache restore: level_offsets len {} != num_levels+1 {}",
1696                artifact.level_offsets.len(),
1697                expected_level_offsets
1698            )));
1699        }
1700
1701        // -- Set handle metadata --
1702        handle.num_nodes = num_nodes;
1703        handle.num_levels = num_levels;
1704        handle.root = root;
1705        handle.max_var = max_var;
1706
1707        // -- Load kernels --
1708        let device = self.provider.device().inner();
1709        let memory = self.provider.memory();
1710
1711        let store_u8 = device
1712            .get_func(CACHE_MODULE, cache_kernels::CACHE_STORE_U8)
1713            .ok_or_else(|| XlogError::Kernel("cache_store_u8 kernel not found".to_string()))?;
1714        let store_u32 = device
1715            .get_func(CACHE_MODULE, cache_kernels::CACHE_STORE_U32)
1716            .ok_or_else(|| XlogError::Kernel("cache_store_u32 kernel not found".to_string()))?;
1717        let store_i32 = device
1718            .get_func(CACHE_MODULE, cache_kernels::CACHE_STORE_I32)
1719            .ok_or_else(|| XlogError::Kernel("cache_store_i32 kernel not found".to_string()))?;
1720        let store_meta = device
1721            .get_func(CACHE_MODULE, cache_kernels::CACHE_STORE_META)
1722            .ok_or_else(|| XlogError::Kernel("cache_store_meta kernel not found".to_string()))?;
1723
1724        let block_dim = 256u32;
1725
1726        let node_stride = self.node_cap;
1727        let offset_stride = self.node_cap.checked_add(1).ok_or_else(|| {
1728            XlogError::Compilation("GpuCircuitCache restore: node_cap overflow".to_string())
1729        })?;
1730        let level_offset_stride = self.level_cap.checked_add(1).ok_or_else(|| {
1731            XlogError::Compilation("GpuCircuitCache restore: level_cap overflow".to_string())
1732        })?;
1733        let var_stride = self.var_cap.checked_add(1).ok_or_else(|| {
1734            XlogError::Compilation("GpuCircuitCache restore: var_cap overflow".to_string())
1735        })?;
1736
1737        let num_nodes_plus1 = num_nodes.checked_add(1).ok_or_else(|| {
1738            XlogError::Compilation("GpuCircuitCache restore: num_nodes overflow".to_string())
1739        })?;
1740        let num_levels_plus1 = num_levels.checked_add(1).ok_or_else(|| {
1741            XlogError::Compilation("GpuCircuitCache restore: num_levels overflow".to_string())
1742        })?;
1743
1744        // -- Upload node_type (u8, num_nodes elements) --
1745        let grid_nodes = cache_grid_dim_for_u32_count(
1746            "GpuCircuitCache restore node_type",
1747            num_nodes,
1748            block_dim,
1749        )?;
1750        if grid_nodes != 0 {
1751            let mut d_node_type = memory.alloc::<u8>(num_nodes as usize)?;
1752            self.provider
1753                .htod_sync_copy_into_tracked(
1754                    &artifact.node_type[..num_nodes as usize],
1755                    &mut d_node_type,
1756                )
1757                .map_err(|e| XlogError::Kernel(format!("restore htod node_type failed: {}", e)))?;
1758            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1759            unsafe {
1760                store_u8.clone().launch(
1761                    LaunchConfig {
1762                        grid_dim: (grid_nodes, 1, 1),
1763                        block_dim: (block_dim, 1, 1),
1764                        shared_mem_bytes: 0,
1765                    },
1766                    (
1767                        handle.slot_device(),
1768                        handle.compile_needed_device(),
1769                        node_stride,
1770                        &d_node_type,
1771                        &mut self.node_type,
1772                        num_nodes,
1773                    ),
1774                )
1775            }
1776            .map_err(|e| {
1777                XlogError::Kernel(format!("restore cache_store node_type failed: {}", e))
1778            })?;
1779        }
1780
1781        // -- Upload child_offsets (u32, num_nodes+1 elements) --
1782        let grid_offsets = cache_grid_dim_for_u32_count(
1783            "GpuCircuitCache restore child_offsets",
1784            num_nodes_plus1,
1785            block_dim,
1786        )?;
1787        if grid_offsets != 0 {
1788            let mut d_child_offsets = memory.alloc::<u32>(num_nodes_plus1 as usize)?;
1789            self.provider
1790                .htod_sync_copy_into_tracked(
1791                    &artifact.child_offsets[..num_nodes_plus1 as usize],
1792                    &mut d_child_offsets,
1793                )
1794                .map_err(|e| {
1795                    XlogError::Kernel(format!("restore htod child_offsets failed: {}", e))
1796                })?;
1797            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1798            unsafe {
1799                store_u32.clone().launch(
1800                    LaunchConfig {
1801                        grid_dim: (grid_offsets, 1, 1),
1802                        block_dim: (block_dim, 1, 1),
1803                        shared_mem_bytes: 0,
1804                    },
1805                    (
1806                        handle.slot_device(),
1807                        handle.compile_needed_device(),
1808                        offset_stride,
1809                        &d_child_offsets,
1810                        &mut self.child_offsets,
1811                        num_nodes_plus1,
1812                    ),
1813                )
1814            }
1815            .map_err(|e| {
1816                XlogError::Kernel(format!("restore cache_store child_offsets failed: {}", e))
1817            })?;
1818        }
1819
1820        // -- Upload child_indices (u32, num_edges elements) --
1821        let grid_edges = cache_grid_dim_for_u32_count(
1822            "GpuCircuitCache restore child_indices",
1823            num_edges,
1824            block_dim,
1825        )?;
1826        if grid_edges != 0 {
1827            let mut d_child_indices = memory.alloc::<u32>(num_edges as usize)?;
1828            self.provider
1829                .htod_sync_copy_into_tracked(
1830                    &artifact.child_indices[..num_edges as usize],
1831                    &mut d_child_indices,
1832                )
1833                .map_err(|e| {
1834                    XlogError::Kernel(format!("restore htod child_indices failed: {}", e))
1835                })?;
1836            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1837            unsafe {
1838                store_u32.clone().launch(
1839                    LaunchConfig {
1840                        grid_dim: (grid_edges, 1, 1),
1841                        block_dim: (block_dim, 1, 1),
1842                        shared_mem_bytes: 0,
1843                    },
1844                    (
1845                        handle.slot_device(),
1846                        handle.compile_needed_device(),
1847                        self.edge_cap,
1848                        &d_child_indices,
1849                        &mut self.child_indices,
1850                        num_edges,
1851                    ),
1852                )
1853            }
1854            .map_err(|e| {
1855                XlogError::Kernel(format!("restore cache_store child_indices failed: {}", e))
1856            })?;
1857        }
1858
1859        // -- Upload lit (i32, num_nodes elements) --
1860        if grid_nodes != 0 {
1861            let mut d_lit = memory.alloc::<i32>(num_nodes as usize)?;
1862            self.provider
1863                .htod_sync_copy_into_tracked(&artifact.lit[..num_nodes as usize], &mut d_lit)
1864                .map_err(|e| XlogError::Kernel(format!("restore htod lit failed: {}", e)))?;
1865            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1866            unsafe {
1867                store_i32.clone().launch(
1868                    LaunchConfig {
1869                        grid_dim: (grid_nodes, 1, 1),
1870                        block_dim: (block_dim, 1, 1),
1871                        shared_mem_bytes: 0,
1872                    },
1873                    (
1874                        handle.slot_device(),
1875                        handle.compile_needed_device(),
1876                        node_stride,
1877                        &d_lit,
1878                        &mut self.lit,
1879                        num_nodes,
1880                    ),
1881                )
1882            }
1883            .map_err(|e| XlogError::Kernel(format!("restore cache_store lit failed: {}", e)))?;
1884
1885            // -- Upload decision_var (u32, num_nodes elements) --
1886            let mut d_decision_var = memory.alloc::<u32>(num_nodes as usize)?;
1887            self.provider
1888                .htod_sync_copy_into_tracked(
1889                    &artifact.decision_var[..num_nodes as usize],
1890                    &mut d_decision_var,
1891                )
1892                .map_err(|e| {
1893                    XlogError::Kernel(format!("restore htod decision_var failed: {}", e))
1894                })?;
1895            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1896            unsafe {
1897                store_u32.clone().launch(
1898                    LaunchConfig {
1899                        grid_dim: (grid_nodes, 1, 1),
1900                        block_dim: (block_dim, 1, 1),
1901                        shared_mem_bytes: 0,
1902                    },
1903                    (
1904                        handle.slot_device(),
1905                        handle.compile_needed_device(),
1906                        node_stride,
1907                        &d_decision_var,
1908                        &mut self.decision_var,
1909                        num_nodes,
1910                    ),
1911                )
1912            }
1913            .map_err(|e| {
1914                XlogError::Kernel(format!("restore cache_store decision_var failed: {}", e))
1915            })?;
1916
1917            // -- Upload decision_child_false (u32, num_nodes elements) --
1918            let mut d_decision_child_false = memory.alloc::<u32>(num_nodes as usize)?;
1919            self.provider
1920                .htod_sync_copy_into_tracked(
1921                    &artifact.decision_child_false[..num_nodes as usize],
1922                    &mut d_decision_child_false,
1923                )
1924                .map_err(|e| {
1925                    XlogError::Kernel(format!("restore htod decision_child_false failed: {}", e))
1926                })?;
1927            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1928            unsafe {
1929                store_u32.clone().launch(
1930                    LaunchConfig {
1931                        grid_dim: (grid_nodes, 1, 1),
1932                        block_dim: (block_dim, 1, 1),
1933                        shared_mem_bytes: 0,
1934                    },
1935                    (
1936                        handle.slot_device(),
1937                        handle.compile_needed_device(),
1938                        node_stride,
1939                        &d_decision_child_false,
1940                        &mut self.decision_child_false,
1941                        num_nodes,
1942                    ),
1943                )
1944            }
1945            .map_err(|e| {
1946                XlogError::Kernel(format!(
1947                    "restore cache_store decision_child_false failed: {}",
1948                    e
1949                ))
1950            })?;
1951
1952            // -- Upload decision_child_true (u32, num_nodes elements) --
1953            let mut d_decision_child_true = memory.alloc::<u32>(num_nodes as usize)?;
1954            self.provider
1955                .htod_sync_copy_into_tracked(
1956                    &artifact.decision_child_true[..num_nodes as usize],
1957                    &mut d_decision_child_true,
1958                )
1959                .map_err(|e| {
1960                    XlogError::Kernel(format!("restore htod decision_child_true failed: {}", e))
1961                })?;
1962            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1963            unsafe {
1964                store_u32.clone().launch(
1965                    LaunchConfig {
1966                        grid_dim: (grid_nodes, 1, 1),
1967                        block_dim: (block_dim, 1, 1),
1968                        shared_mem_bytes: 0,
1969                    },
1970                    (
1971                        handle.slot_device(),
1972                        handle.compile_needed_device(),
1973                        node_stride,
1974                        &d_decision_child_true,
1975                        &mut self.decision_child_true,
1976                        num_nodes,
1977                    ),
1978                )
1979            }
1980            .map_err(|e| {
1981                XlogError::Kernel(format!(
1982                    "restore cache_store decision_child_true failed: {}",
1983                    e
1984                ))
1985            })?;
1986
1987            // -- Upload level_nodes (u32, num_nodes elements) --
1988            let mut d_level_nodes = memory.alloc::<u32>(num_nodes as usize)?;
1989            self.provider
1990                .htod_sync_copy_into_tracked(
1991                    &artifact.level_nodes[..num_nodes as usize],
1992                    &mut d_level_nodes,
1993                )
1994                .map_err(|e| {
1995                    XlogError::Kernel(format!("restore htod level_nodes failed: {}", e))
1996                })?;
1997            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
1998            unsafe {
1999                store_u32.clone().launch(
2000                    LaunchConfig {
2001                        grid_dim: (grid_nodes, 1, 1),
2002                        block_dim: (block_dim, 1, 1),
2003                        shared_mem_bytes: 0,
2004                    },
2005                    (
2006                        handle.slot_device(),
2007                        handle.compile_needed_device(),
2008                        node_stride,
2009                        &d_level_nodes,
2010                        &mut self.level_nodes,
2011                        num_nodes,
2012                    ),
2013                )
2014            }
2015            .map_err(|e| {
2016                XlogError::Kernel(format!("restore cache_store level_nodes failed: {}", e))
2017            })?;
2018        }
2019
2020        // -- Upload level_offsets (u32, num_levels+1 elements) --
2021        let grid_levels = cache_grid_dim_for_u32_count(
2022            "GpuCircuitCache restore level_offsets",
2023            num_levels_plus1,
2024            block_dim,
2025        )?;
2026        if grid_levels != 0 {
2027            let mut d_level_offsets = memory.alloc::<u32>(num_levels_plus1 as usize)?;
2028            self.provider
2029                .htod_sync_copy_into_tracked(
2030                    &artifact.level_offsets[..num_levels_plus1 as usize],
2031                    &mut d_level_offsets,
2032                )
2033                .map_err(|e| {
2034                    XlogError::Kernel(format!("restore htod level_offsets failed: {}", e))
2035                })?;
2036            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2037            unsafe {
2038                store_u32.clone().launch(
2039                    LaunchConfig {
2040                        grid_dim: (grid_levels, 1, 1),
2041                        block_dim: (block_dim, 1, 1),
2042                        shared_mem_bytes: 0,
2043                    },
2044                    (
2045                        handle.slot_device(),
2046                        handle.compile_needed_device(),
2047                        level_offset_stride,
2048                        &d_level_offsets,
2049                        &mut self.level_offsets,
2050                        num_levels_plus1,
2051                    ),
2052                )
2053            }
2054            .map_err(|e| {
2055                XlogError::Kernel(format!("restore cache_store level_offsets failed: {}", e))
2056            })?;
2057        }
2058
2059        // -- Store metadata via cache_store_meta kernel --
2060        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2061        unsafe {
2062            store_meta.clone().launch(
2063                LaunchConfig {
2064                    grid_dim: (1, 1, 1),
2065                    block_dim: (1, 1, 1),
2066                    shared_mem_bytes: 0,
2067                },
2068                (
2069                    handle.slot_device(),
2070                    handle.compile_needed_device(),
2071                    self.num_slots,
2072                    num_nodes,
2073                    num_levels,
2074                    root,
2075                    max_var,
2076                    &mut self.meta_num_nodes,
2077                    &mut self.meta_num_levels,
2078                    &mut self.meta_root,
2079                    &mut self.meta_max_var,
2080                ),
2081            )
2082        }
2083        .map_err(|e| XlogError::Kernel(format!("restore cache_store_meta failed: {}", e)))?;
2084
2085        // -- Zero the free_var_mask region for this slot, then conditionally write --
2086        let slot_idx = handle.slot_index() as usize;
2087
2088        // Zero the slot's mask region by uploading a zero buffer and storing it.
2089        // We always zero to ensure stale mask data from a previous occupant is cleared.
2090        let mask_cap = var_stride; // max_var+1 capacity per slot
2091        let grid_mask_zero = cache_grid_dim_for_u32_count(
2092            "GpuCircuitCache restore zero free_var_mask",
2093            mask_cap,
2094            block_dim,
2095        )?;
2096        if grid_mask_zero != 0 {
2097            let mut d_zeros = memory.alloc::<u8>(mask_cap as usize)?;
2098            device.memset_zeros(&mut d_zeros).map_err(|e| {
2099                XlogError::Kernel(format!("restore memset_zeros free_var_mask failed: {}", e))
2100            })?;
2101            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2102            unsafe {
2103                store_u8.clone().launch(
2104                    LaunchConfig {
2105                        grid_dim: (grid_mask_zero, 1, 1),
2106                        block_dim: (block_dim, 1, 1),
2107                        shared_mem_bytes: 0,
2108                    },
2109                    (
2110                        handle.slot_device(),
2111                        handle.compile_needed_device(),
2112                        var_stride,
2113                        &d_zeros,
2114                        &mut self.free_var_mask,
2115                        mask_cap,
2116                    ),
2117                )
2118            }
2119            .map_err(|e| {
2120                XlogError::Kernel(format!(
2121                    "restore cache_store zero free_var_mask failed: {}",
2122                    e
2123                ))
2124            })?;
2125        }
2126
2127        // Write the actual free_var_mask if the artifact has one.
2128        let has_mask = artifact.has_free_var_mask && !artifact.free_var_mask.is_empty();
2129        if has_mask {
2130            let mask_len = max_var.checked_add(1).ok_or_else(|| {
2131                XlogError::Compilation(
2132                    "GpuCircuitCache restore: free_var_mask max_var overflow".to_string(),
2133                )
2134            })?;
2135            let actual_len = std::cmp::min(mask_len as usize, artifact.free_var_mask.len());
2136            if actual_len > 0 {
2137                let actual_len_u32 = u32::try_from(actual_len).map_err(|_| {
2138                    XlogError::Compilation(
2139                        "GpuCircuitCache restore free_var_mask len exceeds u32".to_string(),
2140                    )
2141                })?;
2142                let grid_mask = cache_grid_dim_for_u32_count(
2143                    "GpuCircuitCache restore free_var_mask",
2144                    actual_len_u32,
2145                    block_dim,
2146                )?;
2147                if grid_mask != 0 {
2148                    let mut d_mask = memory.alloc::<u8>(actual_len)?;
2149                    self.provider
2150                        .htod_sync_copy_into_tracked(
2151                            &artifact.free_var_mask[..actual_len],
2152                            &mut d_mask,
2153                        )
2154                        .map_err(|e| {
2155                            XlogError::Kernel(format!("restore htod free_var_mask failed: {}", e))
2156                        })?;
2157                    // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2158                    unsafe {
2159                        store_u8.clone().launch(
2160                            LaunchConfig {
2161                                grid_dim: (grid_mask, 1, 1),
2162                                block_dim: (block_dim, 1, 1),
2163                                shared_mem_bytes: 0,
2164                            },
2165                            (
2166                                handle.slot_device(),
2167                                handle.compile_needed_device(),
2168                                var_stride,
2169                                &d_mask,
2170                                &mut self.free_var_mask,
2171                                actual_len_u32,
2172                            ),
2173                        )
2174                    }
2175                    .map_err(|e| {
2176                        XlogError::Kernel(format!(
2177                            "restore cache_store free_var_mask failed: {}",
2178                            e
2179                        ))
2180                    })?;
2181                }
2182            }
2183        }
2184
2185        // Set per-slot has_free_var_mask flag.
2186        debug_assert!(
2187            slot_idx < self.has_free_var_mask.len(),
2188            "slot_index {} exceeds num_slots {}",
2189            slot_idx,
2190            self.has_free_var_mask.len()
2191        );
2192        if slot_idx < self.has_free_var_mask.len() {
2193            self.has_free_var_mask[slot_idx] = has_mask;
2194        }
2195
2196        // No device synchronize needed: all stores are H→D copies followed by
2197        // same-stream kernel launches, so ordering is guaranteed.
2198        Ok(())
2199    }
2200
2201    /// Extract a [`disk_cache::CircuitArtifact`] from a populated GPU cache slot.
2202    ///
2203    /// This is the inverse of [`restore_from_host_arrays`]: it reads device-resident
2204    /// topology arrays from the cache slot and builds host vectors suitable for disk
2205    /// serialization. The caller must ensure the slot has been populated (i.e. after
2206    /// `store_from_xgcf` + `store_free_var_mask`).
2207    pub(crate) fn build_artifact_from_device(
2208        &self,
2209        handle: &GpuCircuitCacheHandle,
2210        provider: &Arc<CudaKernelProvider>,
2211    ) -> Result<disk_cache::CircuitArtifact> {
2212        let device = provider.device().inner();
2213        let slot = handle.slot_index() as usize;
2214        let num_nodes = handle.num_nodes();
2215        let num_levels = handle.num_levels();
2216        let root = handle.root();
2217        let max_var = handle.max_var();
2218
2219        if num_nodes == 0 {
2220            return Err(XlogError::Compilation(
2221                "build_artifact_from_device: num_nodes is 0".to_string(),
2222            ));
2223        }
2224
2225        let node_stride = self.node_cap as usize;
2226        let offset_stride = (self.node_cap as usize) + 1;
2227        let edge_stride = self.edge_cap as usize;
2228        let level_offset_stride = (self.level_cap as usize) + 1;
2229        let var_stride = (self.var_cap as usize) + 1;
2230
2231        let slot_node_start = slot * node_stride;
2232        let slot_offset_start = slot * offset_stride;
2233        let slot_level_offset_start = slot * level_offset_stride;
2234        let slot_var_start = slot * var_stride;
2235
2236        let nn = num_nodes as usize;
2237        let nn1 = nn + 1;
2238        let nl1 = (num_levels as usize) + 1;
2239
2240        // Determine num_edges from child_offsets[num_nodes] - child_offsets[0].
2241        // We read child_offsets first, then derive num_edges from it.
2242        let child_offsets_view = self
2243            .child_offsets
2244            .slice(slot_offset_start..(slot_offset_start + nn1));
2245        let child_offsets: Vec<u32> = device
2246            .dtoh_sync_copy(&child_offsets_view)
2247            .map_err(|e| XlogError::Kernel(format!("build_artifact dtoh child_offsets: {}", e)))?;
2248        let num_edges = if nn1 > 0 {
2249            child_offsets[nn]
2250                .checked_sub(child_offsets[0])
2251                .ok_or_else(|| {
2252                    XlogError::Compilation(
2253                        "build_artifact_from_device: child_offsets[num_nodes] < child_offsets[0]"
2254                            .to_string(),
2255                    )
2256                })?
2257        } else {
2258            0
2259        };
2260
2261        // Read child_indices from the edge region.
2262        let slot_edge_start = slot * edge_stride;
2263        let ne = num_edges as usize;
2264        let child_indices: Vec<u32> = if ne > 0 {
2265            let view = self
2266                .child_indices
2267                .slice(slot_edge_start..(slot_edge_start + ne));
2268            device.dtoh_sync_copy(&view).map_err(|e| {
2269                XlogError::Kernel(format!("build_artifact dtoh child_indices: {}", e))
2270            })?
2271        } else {
2272            Vec::new()
2273        };
2274
2275        // node_type (u8)
2276        let node_type_view = self
2277            .node_type
2278            .slice(slot_node_start..(slot_node_start + nn));
2279        let node_type: Vec<u8> = device
2280            .dtoh_sync_copy(&node_type_view)
2281            .map_err(|e| XlogError::Kernel(format!("build_artifact dtoh node_type: {}", e)))?;
2282
2283        // lit (i32)
2284        let lit_view = self.lit.slice(slot_node_start..(slot_node_start + nn));
2285        let lit: Vec<i32> = device
2286            .dtoh_sync_copy(&lit_view)
2287            .map_err(|e| XlogError::Kernel(format!("build_artifact dtoh lit: {}", e)))?;
2288
2289        // decision_var (u32)
2290        let dv_view = self
2291            .decision_var
2292            .slice(slot_node_start..(slot_node_start + nn));
2293        let decision_var: Vec<u32> = device
2294            .dtoh_sync_copy(&dv_view)
2295            .map_err(|e| XlogError::Kernel(format!("build_artifact dtoh decision_var: {}", e)))?;
2296
2297        // decision_child_false (u32)
2298        let dcf_view = self
2299            .decision_child_false
2300            .slice(slot_node_start..(slot_node_start + nn));
2301        let decision_child_false: Vec<u32> = device.dtoh_sync_copy(&dcf_view).map_err(|e| {
2302            XlogError::Kernel(format!("build_artifact dtoh decision_child_false: {}", e))
2303        })?;
2304
2305        // decision_child_true (u32)
2306        let dct_view = self
2307            .decision_child_true
2308            .slice(slot_node_start..(slot_node_start + nn));
2309        let decision_child_true: Vec<u32> = device.dtoh_sync_copy(&dct_view).map_err(|e| {
2310            XlogError::Kernel(format!("build_artifact dtoh decision_child_true: {}", e))
2311        })?;
2312
2313        // level_nodes (u32)
2314        let ln_view = self
2315            .level_nodes
2316            .slice(slot_node_start..(slot_node_start + nn));
2317        let level_nodes: Vec<u32> = device
2318            .dtoh_sync_copy(&ln_view)
2319            .map_err(|e| XlogError::Kernel(format!("build_artifact dtoh level_nodes: {}", e)))?;
2320
2321        // level_offsets (u32)
2322        let lo_view = self
2323            .level_offsets
2324            .slice(slot_level_offset_start..(slot_level_offset_start + nl1));
2325        let level_offsets: Vec<u32> = device
2326            .dtoh_sync_copy(&lo_view)
2327            .map_err(|e| XlogError::Kernel(format!("build_artifact dtoh level_offsets: {}", e)))?;
2328
2329        // free_var_mask (u8)
2330        let has_free_var_mask = self.has_free_var_mask_for_slot(slot as u32);
2331        let mask_len = (max_var as usize) + 1;
2332        let free_var_mask: Vec<u8> = if mask_len > 0 {
2333            let fvm_view = self
2334                .free_var_mask
2335                .slice(slot_var_start..(slot_var_start + mask_len));
2336            device.dtoh_sync_copy(&fvm_view).map_err(|e| {
2337                XlogError::Kernel(format!("build_artifact dtoh free_var_mask: {}", e))
2338            })?
2339        } else {
2340            Vec::new()
2341        };
2342
2343        Ok(disk_cache::CircuitArtifact {
2344            num_nodes,
2345            num_edges,
2346            num_levels,
2347            root,
2348            max_var,
2349            has_free_var_mask,
2350            node_type,
2351            child_offsets,
2352            child_indices,
2353            lit,
2354            decision_var,
2355            decision_child_false,
2356            decision_child_true,
2357            level_nodes,
2358            level_offsets,
2359            free_var_mask,
2360        })
2361    }
2362
2363    /// Evaluates cached logZ without a host transfer.
2364    ///
2365    /// Callers must ensure evaluated resident weights are NaN-free. NaN or undefined arithmetic
2366    /// that reaches evaluation is represented by a NaN sentinel in `out_log_z`; callers must reject
2367    /// that sentinel at an existing readback. Positive infinity remains valid for value-only
2368    /// evaluation.
2369    pub fn eval_log_wmc_device_inplace(
2370        &mut self,
2371        handle: &GpuCircuitCacheHandle,
2372        out_log_z: &mut TrackedCudaSlice<f64>,
2373    ) -> Result<()> {
2374        self.eval_log_wmc_device_only(handle, out_log_z)
2375    }
2376
2377    /// Device-only implementation of cached value evaluation; see
2378    /// [`Self::eval_log_wmc_device_inplace`] for its numeric preconditions and sentinel contract.
2379    pub fn eval_log_wmc_device_only(
2380        &mut self,
2381        handle: &GpuCircuitCacheHandle,
2382        out_log_z: &mut TrackedCudaSlice<f64>,
2383    ) -> Result<()> {
2384        if out_log_z.len() != 1 {
2385            return Err(XlogError::Compilation(format!(
2386                "GPU cache logZ output len {} != 1",
2387                out_log_z.len()
2388            )));
2389        }
2390
2391        {
2392            let device = self.provider.device().inner();
2393            let eval_all = device
2394                .get_func(
2395                    xlog_cuda::CIRCUIT_MODULE,
2396                    xlog_cuda::circuit_kernels::XGCF_EVAL_ALL_LEVELS_CACHED,
2397                )
2398                .ok_or_else(|| {
2399                    XlogError::Kernel("xgcf_eval_all_levels_cached kernel not found".to_string())
2400                })?;
2401
2402            let block_size: u32 = 256;
2403            let mut params: Vec<*mut std::ffi::c_void> = vec![
2404                handle.slot_device().as_kernel_param(),
2405                self.node_cap.as_kernel_param(),
2406                self.edge_cap.as_kernel_param(),
2407                self.level_cap.as_kernel_param(),
2408                self.var_cap.as_kernel_param(),
2409                (&self.node_type).as_kernel_param(),
2410                (&self.child_offsets).as_kernel_param(),
2411                (&self.child_indices).as_kernel_param(),
2412                (&self.lit).as_kernel_param(),
2413                (&self.decision_var).as_kernel_param(),
2414                (&self.decision_child_false).as_kernel_param(),
2415                (&self.decision_child_true).as_kernel_param(),
2416                (&self.level_nodes).as_kernel_param(),
2417                (&self.level_offsets).as_kernel_param(),
2418                (&self.var_log_true).as_kernel_param(),
2419                (&self.var_log_false).as_kernel_param(),
2420                (&self.values).as_kernel_param(),
2421                (&self.meta_num_levels).as_kernel_param(),
2422            ];
2423            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2424            unsafe {
2425                eval_all.clone().launch(
2426                    LaunchConfig {
2427                        grid_dim: (1, 1, 1),
2428                        block_dim: (block_size, 1, 1),
2429                        shared_mem_bytes: 0,
2430                    },
2431                    &mut params,
2432                )
2433            }
2434            .map_err(|e| XlogError::Kernel(format!("xgcf_eval_all_levels_cached failed: {}", e)))?;
2435        }
2436
2437        self.apply_free_var_correction_cached(handle, true, false)?;
2438
2439        let device = self.provider.device().inner();
2440        let copy_root = device
2441            .get_func(
2442                xlog_cuda::CIRCUIT_MODULE,
2443                xlog_cuda::circuit_kernels::XGCF_COPY_ROOT_CACHED_META,
2444            )
2445            .ok_or_else(|| {
2446                XlogError::Kernel("xgcf_copy_root_cached_meta kernel not found".to_string())
2447            })?;
2448        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2449        unsafe {
2450            copy_root.clone().launch(
2451                LaunchConfig {
2452                    grid_dim: (1, 1, 1),
2453                    block_dim: (1, 1, 1),
2454                    shared_mem_bytes: 0,
2455                },
2456                (
2457                    handle.slot_device(),
2458                    self.node_cap,
2459                    &self.values,
2460                    &self.meta_root,
2461                    out_log_z,
2462                ),
2463            )
2464        }
2465        .map_err(|e| XlogError::Kernel(format!("xgcf_copy_root_cached_meta failed: {}", e)))?;
2466
2467        // No device synchronize: callers read back with a synchronous host copy
2468        // or pass the result to subsequent GPU operations (same-stream ordering).
2469        Ok(())
2470    }
2471
2472    /// Evaluates cached gradients without a host transfer.
2473    ///
2474    /// Callers must validate numeric inputs and eventual outputs. Positive infinity is valid when
2475    /// the actual backward path remains finite. This asynchronous API cannot return a host-side
2476    /// numeric error; NaN or undefined normalization that reaches evaluation is observable as a
2477    /// non-finite device result.
2478    pub fn eval_grads_inplace(&mut self, handle: &GpuCircuitCacheHandle) -> Result<()> {
2479        let device = self.provider.device().inner();
2480        let eval_all = device
2481            .get_func(
2482                xlog_cuda::CIRCUIT_MODULE,
2483                xlog_cuda::circuit_kernels::XGCF_EVAL_ALL_LEVELS_CACHED,
2484            )
2485            .ok_or_else(|| {
2486                XlogError::Kernel("xgcf_eval_all_levels_cached kernel not found".to_string())
2487            })?;
2488        let block_size: u32 = 256;
2489        let mut params: Vec<*mut std::ffi::c_void> = vec![
2490            handle.slot_device().as_kernel_param(),
2491            self.node_cap.as_kernel_param(),
2492            self.edge_cap.as_kernel_param(),
2493            self.level_cap.as_kernel_param(),
2494            self.var_cap.as_kernel_param(),
2495            (&self.node_type).as_kernel_param(),
2496            (&self.child_offsets).as_kernel_param(),
2497            (&self.child_indices).as_kernel_param(),
2498            (&self.lit).as_kernel_param(),
2499            (&self.decision_var).as_kernel_param(),
2500            (&self.decision_child_false).as_kernel_param(),
2501            (&self.decision_child_true).as_kernel_param(),
2502            (&self.level_nodes).as_kernel_param(),
2503            (&self.level_offsets).as_kernel_param(),
2504            (&self.var_log_true).as_kernel_param(),
2505            (&self.var_log_false).as_kernel_param(),
2506            (&self.values).as_kernel_param(),
2507            (&self.meta_num_levels).as_kernel_param(),
2508        ];
2509        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2510        unsafe {
2511            eval_all.clone().launch(
2512                LaunchConfig {
2513                    grid_dim: (1, 1, 1),
2514                    block_dim: (block_size, 1, 1),
2515                    shared_mem_bytes: 0,
2516                },
2517                &mut params,
2518            )
2519        }
2520        .map_err(|e| XlogError::Kernel(format!("xgcf_eval_all_levels_cached failed: {}", e)))?;
2521
2522        let device = self.provider.device().inner();
2523        let store_f64 = device
2524            .get_func(CACHE_MODULE, cache_kernels::CACHE_STORE_F64)
2525            .ok_or_else(|| XlogError::Kernel("cache_store_f64 kernel not found".to_string()))?;
2526
2527        let node_stride = self.node_cap;
2528        let var_stride = self.var_cap.checked_add(1).ok_or_else(|| {
2529            XlogError::Compilation("GPU cache eval_grads var_cap overflow".to_string())
2530        })?;
2531        let weights_len = self.var_cap.checked_add(1).ok_or_else(|| {
2532            XlogError::Compilation("GPU cache eval_grads var_cap overflow".to_string())
2533        })?;
2534
2535        let grid_nodes = cache_grid_dim_for_u32_count(
2536            "GpuCircuitCache eval_grads zero adj",
2537            self.node_cap,
2538            block_size,
2539        )?;
2540        if grid_nodes != 0 {
2541            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2542            unsafe {
2543                store_f64.clone().launch(
2544                    LaunchConfig {
2545                        grid_dim: (grid_nodes, 1, 1),
2546                        block_dim: (block_size, 1, 1),
2547                        shared_mem_bytes: 0,
2548                    },
2549                    (
2550                        handle.slot_device(),
2551                        &self.always_on,
2552                        node_stride,
2553                        &self.zero_f64,
2554                        &mut self.adj,
2555                        self.node_cap,
2556                    ),
2557                )
2558            }
2559            .map_err(|e| XlogError::Kernel(format!("cache zero adj failed: {}", e)))?;
2560        }
2561
2562        let grid_weights = cache_grid_dim_for_u32_count(
2563            "GpuCircuitCache eval_grads zero weights",
2564            weights_len,
2565            block_size,
2566        )?;
2567        if grid_weights != 0 {
2568            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2569            unsafe {
2570                store_f64.clone().launch(
2571                    LaunchConfig {
2572                        grid_dim: (grid_weights, 1, 1),
2573                        block_dim: (block_size, 1, 1),
2574                        shared_mem_bytes: 0,
2575                    },
2576                    (
2577                        handle.slot_device(),
2578                        &self.always_on,
2579                        var_stride,
2580                        &self.zero_f64,
2581                        &mut self.grad_true,
2582                        weights_len,
2583                    ),
2584                )
2585            }
2586            .map_err(|e| XlogError::Kernel(format!("cache zero grad_true failed: {}", e)))?;
2587
2588            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2589            unsafe {
2590                store_f64.clone().launch(
2591                    LaunchConfig {
2592                        grid_dim: (grid_weights, 1, 1),
2593                        block_dim: (block_size, 1, 1),
2594                        shared_mem_bytes: 0,
2595                    },
2596                    (
2597                        handle.slot_device(),
2598                        &self.always_on,
2599                        var_stride,
2600                        &self.zero_f64,
2601                        &mut self.grad_false,
2602                        weights_len,
2603                    ),
2604                )
2605            }
2606            .map_err(|e| XlogError::Kernel(format!("cache zero grad_false failed: {}", e)))?;
2607        }
2608
2609        let add_scalar = device
2610            .get_func(
2611                xlog_cuda::CIRCUIT_MODULE,
2612                xlog_cuda::circuit_kernels::XGCF_ADD_SCALAR_CACHED,
2613            )
2614            .ok_or_else(|| {
2615                XlogError::Kernel("xgcf_add_scalar_cached kernel not found".to_string())
2616            })?;
2617        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2618        unsafe {
2619            add_scalar.clone().launch(
2620                LaunchConfig {
2621                    grid_dim: (1, 1, 1),
2622                    block_dim: (1, 1, 1),
2623                    shared_mem_bytes: 0,
2624                },
2625                (
2626                    handle.slot_device(),
2627                    self.node_cap,
2628                    &mut self.adj,
2629                    &self.meta_root,
2630                    &self.one_f64,
2631                ),
2632            )
2633        }
2634        .map_err(|e| XlogError::Kernel(format!("xgcf_add_scalar_cached (adj) failed: {}", e)))?;
2635
2636        let propagate = device
2637            .get_func(
2638                xlog_cuda::CIRCUIT_MODULE,
2639                xlog_cuda::circuit_kernels::XGCF_BACKWARD_LEVEL_PROPAGATE_CACHED,
2640            )
2641            .ok_or_else(|| {
2642                XlogError::Kernel(
2643                    "xgcf_backward_level_propagate_cached kernel not found".to_string(),
2644                )
2645            })?;
2646        let decision_grad = device
2647            .get_func(
2648                xlog_cuda::CIRCUIT_MODULE,
2649                xlog_cuda::circuit_kernels::XGCF_BACKWARD_LEVEL_DECISION_GRAD_CACHED,
2650            )
2651            .ok_or_else(|| {
2652                XlogError::Kernel(
2653                    "xgcf_backward_level_decision_grad_cached kernel not found".to_string(),
2654                )
2655            })?;
2656        let lit_grad = device
2657            .get_func(
2658                xlog_cuda::CIRCUIT_MODULE,
2659                xlog_cuda::circuit_kernels::XGCF_BACKWARD_LEVEL_LIT_GRAD_CACHED,
2660            )
2661            .ok_or_else(|| {
2662                XlogError::Kernel(
2663                    "xgcf_backward_level_lit_grad_cached kernel not found".to_string(),
2664                )
2665            })?;
2666
2667        let num_blocks = self.node_cap.div_ceil(block_size);
2668        let num_levels = self.level_cap;
2669        for level in (0..num_levels).rev() {
2670            if num_blocks == 0 {
2671                continue;
2672            }
2673            let level_u32: u32 = level;
2674            let mut params: Vec<*mut std::ffi::c_void> = vec![
2675                handle.slot_device().as_kernel_param(),
2676                self.node_cap.as_kernel_param(),
2677                self.edge_cap.as_kernel_param(),
2678                self.level_cap.as_kernel_param(),
2679                self.var_cap.as_kernel_param(),
2680                (&self.node_type).as_kernel_param(),
2681                (&self.child_offsets).as_kernel_param(),
2682                (&self.child_indices).as_kernel_param(),
2683                (&self.decision_var).as_kernel_param(),
2684                (&self.decision_child_false).as_kernel_param(),
2685                (&self.decision_child_true).as_kernel_param(),
2686                (&self.level_nodes).as_kernel_param(),
2687                (&self.level_offsets).as_kernel_param(),
2688                level_u32.as_kernel_param(),
2689                (&self.var_log_true).as_kernel_param(),
2690                (&self.var_log_false).as_kernel_param(),
2691                (&self.values).as_kernel_param(),
2692                (&self.adj).as_kernel_param(),
2693                (&self.meta_num_levels).as_kernel_param(),
2694            ];
2695
2696            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2697            unsafe {
2698                propagate.clone().launch(
2699                    LaunchConfig {
2700                        grid_dim: (num_blocks, 1, 1),
2701                        block_dim: (block_size, 1, 1),
2702                        shared_mem_bytes: 0,
2703                    },
2704                    &mut params,
2705                )
2706            }
2707            .map_err(|e| {
2708                XlogError::Kernel(format!(
2709                    "xgcf_backward_level_propagate_cached failed: {}",
2710                    e
2711                ))
2712            })?;
2713
2714            let mut params: Vec<*mut std::ffi::c_void> = vec![
2715                handle.slot_device().as_kernel_param(),
2716                self.node_cap.as_kernel_param(),
2717                self.edge_cap.as_kernel_param(),
2718                self.level_cap.as_kernel_param(),
2719                self.var_cap.as_kernel_param(),
2720                (&self.node_type).as_kernel_param(),
2721                (&self.decision_var).as_kernel_param(),
2722                (&self.decision_child_false).as_kernel_param(),
2723                (&self.decision_child_true).as_kernel_param(),
2724                (&self.level_nodes).as_kernel_param(),
2725                (&self.level_offsets).as_kernel_param(),
2726                level_u32.as_kernel_param(),
2727                (&self.var_log_true).as_kernel_param(),
2728                (&self.var_log_false).as_kernel_param(),
2729                (&self.values).as_kernel_param(),
2730                (&self.adj).as_kernel_param(),
2731                (&self.grad_true).as_kernel_param(),
2732                (&self.grad_false).as_kernel_param(),
2733                (&self.meta_num_levels).as_kernel_param(),
2734            ];
2735
2736            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2737            unsafe {
2738                decision_grad.clone().launch(
2739                    LaunchConfig {
2740                        grid_dim: (num_blocks, 1, 1),
2741                        block_dim: (block_size, 1, 1),
2742                        shared_mem_bytes: 0,
2743                    },
2744                    &mut params,
2745                )
2746            }
2747            .map_err(|e| {
2748                XlogError::Kernel(format!(
2749                    "xgcf_backward_level_decision_grad_cached failed: {}",
2750                    e
2751                ))
2752            })?;
2753
2754            let mut params: Vec<*mut std::ffi::c_void> = vec![
2755                handle.slot_device().as_kernel_param(),
2756                self.node_cap.as_kernel_param(),
2757                self.edge_cap.as_kernel_param(),
2758                self.level_cap.as_kernel_param(),
2759                self.var_cap.as_kernel_param(),
2760                (&self.node_type).as_kernel_param(),
2761                (&self.lit).as_kernel_param(),
2762                (&self.level_nodes).as_kernel_param(),
2763                (&self.level_offsets).as_kernel_param(),
2764                level_u32.as_kernel_param(),
2765                (&self.adj).as_kernel_param(),
2766                (&self.grad_true).as_kernel_param(),
2767                (&self.grad_false).as_kernel_param(),
2768                (&self.meta_num_levels).as_kernel_param(),
2769            ];
2770
2771            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2772            unsafe {
2773                lit_grad.clone().launch(
2774                    LaunchConfig {
2775                        grid_dim: (num_blocks, 1, 1),
2776                        block_dim: (block_size, 1, 1),
2777                        shared_mem_bytes: 0,
2778                    },
2779                    &mut params,
2780                )
2781            }
2782            .map_err(|e| {
2783                XlogError::Kernel(format!("xgcf_backward_level_lit_grad_cached failed: {}", e))
2784            })?;
2785        }
2786
2787        self.apply_free_var_correction_cached(handle, true, true)?;
2788        // No device synchronize: callers batch multiple eval/backward calls
2789        // before syncing at the query boundary.
2790        Ok(())
2791    }
2792
2793    /// Like [`eval_grads_inplace`] but replaces the per-level backward loop
2794    /// with a single launch of `xgcf_backward_all_levels_cached`, and omits the
2795    /// trailing `device().synchronize()` so that the caller can batch multiple
2796    /// queries before syncing.
2797    /// The same caller-side numeric validation contract applies.
2798    pub fn eval_grads_inplace_fused(&mut self, handle: &GpuCircuitCacheHandle) -> Result<()> {
2799        let device = self.provider.device().inner();
2800        let eval_all = device
2801            .get_func(
2802                xlog_cuda::CIRCUIT_MODULE,
2803                xlog_cuda::circuit_kernels::XGCF_EVAL_ALL_LEVELS_CACHED,
2804            )
2805            .ok_or_else(|| {
2806                XlogError::Kernel("xgcf_eval_all_levels_cached kernel not found".to_string())
2807            })?;
2808        let block_size: u32 = 256;
2809        let mut params: Vec<*mut std::ffi::c_void> = vec![
2810            handle.slot_device().as_kernel_param(),
2811            self.node_cap.as_kernel_param(),
2812            self.edge_cap.as_kernel_param(),
2813            self.level_cap.as_kernel_param(),
2814            self.var_cap.as_kernel_param(),
2815            (&self.node_type).as_kernel_param(),
2816            (&self.child_offsets).as_kernel_param(),
2817            (&self.child_indices).as_kernel_param(),
2818            (&self.lit).as_kernel_param(),
2819            (&self.decision_var).as_kernel_param(),
2820            (&self.decision_child_false).as_kernel_param(),
2821            (&self.decision_child_true).as_kernel_param(),
2822            (&self.level_nodes).as_kernel_param(),
2823            (&self.level_offsets).as_kernel_param(),
2824            (&self.var_log_true).as_kernel_param(),
2825            (&self.var_log_false).as_kernel_param(),
2826            (&self.values).as_kernel_param(),
2827            (&self.meta_num_levels).as_kernel_param(),
2828        ];
2829        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2830        unsafe {
2831            eval_all.clone().launch(
2832                LaunchConfig {
2833                    grid_dim: (1, 1, 1),
2834                    block_dim: (block_size, 1, 1),
2835                    shared_mem_bytes: 0,
2836                },
2837                &mut params,
2838            )
2839        }
2840        .map_err(|e| XlogError::Kernel(format!("xgcf_eval_all_levels_cached failed: {}", e)))?;
2841
2842        let device = self.provider.device().inner();
2843        let store_f64 = device
2844            .get_func(CACHE_MODULE, cache_kernels::CACHE_STORE_F64)
2845            .ok_or_else(|| XlogError::Kernel("cache_store_f64 kernel not found".to_string()))?;
2846
2847        let node_stride = self.node_cap;
2848        let var_stride = self.var_cap.checked_add(1).ok_or_else(|| {
2849            XlogError::Compilation("GPU cache eval_grads var_cap overflow".to_string())
2850        })?;
2851        let weights_len = self.var_cap.checked_add(1).ok_or_else(|| {
2852            XlogError::Compilation("GPU cache eval_grads var_cap overflow".to_string())
2853        })?;
2854
2855        let grid_nodes = cache_grid_dim_for_u32_count(
2856            "GpuCircuitCache batched eval_grads zero adj",
2857            self.node_cap,
2858            block_size,
2859        )?;
2860        if grid_nodes != 0 {
2861            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2862            unsafe {
2863                store_f64.clone().launch(
2864                    LaunchConfig {
2865                        grid_dim: (grid_nodes, 1, 1),
2866                        block_dim: (block_size, 1, 1),
2867                        shared_mem_bytes: 0,
2868                    },
2869                    (
2870                        handle.slot_device(),
2871                        &self.always_on,
2872                        node_stride,
2873                        &self.zero_f64,
2874                        &mut self.adj,
2875                        self.node_cap,
2876                    ),
2877                )
2878            }
2879            .map_err(|e| XlogError::Kernel(format!("cache zero adj failed: {}", e)))?;
2880        }
2881
2882        let grid_weights = cache_grid_dim_for_u32_count(
2883            "GpuCircuitCache batched eval_grads zero weights",
2884            weights_len,
2885            block_size,
2886        )?;
2887        if grid_weights != 0 {
2888            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2889            unsafe {
2890                store_f64.clone().launch(
2891                    LaunchConfig {
2892                        grid_dim: (grid_weights, 1, 1),
2893                        block_dim: (block_size, 1, 1),
2894                        shared_mem_bytes: 0,
2895                    },
2896                    (
2897                        handle.slot_device(),
2898                        &self.always_on,
2899                        var_stride,
2900                        &self.zero_f64,
2901                        &mut self.grad_true,
2902                        weights_len,
2903                    ),
2904                )
2905            }
2906            .map_err(|e| XlogError::Kernel(format!("cache zero grad_true failed: {}", e)))?;
2907
2908            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2909            unsafe {
2910                store_f64.clone().launch(
2911                    LaunchConfig {
2912                        grid_dim: (grid_weights, 1, 1),
2913                        block_dim: (block_size, 1, 1),
2914                        shared_mem_bytes: 0,
2915                    },
2916                    (
2917                        handle.slot_device(),
2918                        &self.always_on,
2919                        var_stride,
2920                        &self.zero_f64,
2921                        &mut self.grad_false,
2922                        weights_len,
2923                    ),
2924                )
2925            }
2926            .map_err(|e| XlogError::Kernel(format!("cache zero grad_false failed: {}", e)))?;
2927        }
2928
2929        let add_scalar = device
2930            .get_func(
2931                xlog_cuda::CIRCUIT_MODULE,
2932                xlog_cuda::circuit_kernels::XGCF_ADD_SCALAR_CACHED,
2933            )
2934            .ok_or_else(|| {
2935                XlogError::Kernel("xgcf_add_scalar_cached kernel not found".to_string())
2936            })?;
2937        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2938        unsafe {
2939            add_scalar.clone().launch(
2940                LaunchConfig {
2941                    grid_dim: (1, 1, 1),
2942                    block_dim: (1, 1, 1),
2943                    shared_mem_bytes: 0,
2944                },
2945                (
2946                    handle.slot_device(),
2947                    self.node_cap,
2948                    &mut self.adj,
2949                    &self.meta_root,
2950                    &self.one_f64,
2951                ),
2952            )
2953        }
2954        .map_err(|e| XlogError::Kernel(format!("xgcf_add_scalar_cached (adj) failed: {}", e)))?;
2955
2956        // Fused backward: single kernel replaces the per-level loop.
2957        let backward_all = device
2958            .get_func(
2959                xlog_cuda::CIRCUIT_MODULE,
2960                xlog_cuda::circuit_kernels::XGCF_BACKWARD_ALL_LEVELS_CACHED,
2961            )
2962            .ok_or_else(|| XlogError::Kernel("xgcf_backward_all_levels_cached not found".into()))?;
2963
2964        let mut params: Vec<*mut std::ffi::c_void> = vec![
2965            handle.slot_device().as_kernel_param(),
2966            self.node_cap.as_kernel_param(),
2967            self.edge_cap.as_kernel_param(),
2968            self.level_cap.as_kernel_param(),
2969            self.var_cap.as_kernel_param(),
2970            (&self.node_type).as_kernel_param(),
2971            (&self.child_offsets).as_kernel_param(),
2972            (&self.child_indices).as_kernel_param(),
2973            (&self.decision_var).as_kernel_param(),
2974            (&self.decision_child_false).as_kernel_param(),
2975            (&self.decision_child_true).as_kernel_param(),
2976            (&self.lit).as_kernel_param(),
2977            (&self.level_nodes).as_kernel_param(),
2978            (&self.level_offsets).as_kernel_param(),
2979            (&self.var_log_true).as_kernel_param(),
2980            (&self.var_log_false).as_kernel_param(),
2981            (&self.values).as_kernel_param(),
2982            (&self.adj).as_kernel_param(),
2983            (&self.grad_true).as_kernel_param(),
2984            (&self.grad_false).as_kernel_param(),
2985            (&self.meta_num_levels).as_kernel_param(),
2986        ];
2987
2988        // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
2989        unsafe {
2990            backward_all.clone().launch(
2991                LaunchConfig {
2992                    grid_dim: (1, 1, 1),
2993                    block_dim: (block_size, 1, 1),
2994                    shared_mem_bytes: 0,
2995                },
2996                &mut params,
2997            )
2998        }
2999        .map_err(|e| XlogError::Kernel(format!("xgcf_backward_all_levels_cached failed: {}", e)))?;
3000
3001        self.apply_free_var_correction_cached(handle, true, true)?;
3002        Ok(())
3003    }
3004
3005    fn apply_free_var_correction_cached(
3006        &mut self,
3007        handle: &GpuCircuitCacheHandle,
3008        apply_log_z: bool,
3009        apply_grads: bool,
3010    ) -> Result<()> {
3011        if !self.has_free_var_mask_for_slot(handle.slot_index()) {
3012            return Ok(());
3013        }
3014        let n = self
3015            .var_cap
3016            .checked_add(1)
3017            .ok_or_else(|| XlogError::Compilation("GPU cache free-var overflow".to_string()))?;
3018        if n == 0 {
3019            return Ok(());
3020        }
3021
3022        let device = self.provider.device().inner();
3023        let block_dim = 256u32;
3024        let grid_dim = n.div_ceil(block_dim);
3025
3026        if apply_grads {
3027            let apply_grad = device
3028                .get_func(
3029                    xlog_cuda::CIRCUIT_MODULE,
3030                    xlog_cuda::circuit_kernels::XGCF_FREE_VAR_APPLY_GRAD_CACHED,
3031                )
3032                .ok_or_else(|| {
3033                    XlogError::Kernel(
3034                        "xgcf_free_var_apply_grad_cached kernel not found".to_string(),
3035                    )
3036                })?;
3037            // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
3038            unsafe {
3039                apply_grad.clone().launch(
3040                    LaunchConfig {
3041                        grid_dim: (grid_dim, 1, 1),
3042                        block_dim: (block_dim, 1, 1),
3043                        shared_mem_bytes: 0,
3044                    },
3045                    (
3046                        handle.slot_device(),
3047                        self.var_cap,
3048                        &self.free_var_mask,
3049                        &self.var_log_true,
3050                        &self.var_log_false,
3051                        n,
3052                        &mut self.grad_true,
3053                        &mut self.grad_false,
3054                    ),
3055                )
3056            }
3057            .map_err(|e| {
3058                XlogError::Kernel(format!("xgcf_free_var_apply_grad_cached failed: {}", e))
3059            })?;
3060        }
3061
3062        if apply_log_z {
3063            let reduce_stage = device
3064                .get_func(
3065                    xlog_cuda::CIRCUIT_MODULE,
3066                    xlog_cuda::circuit_kernels::XGCF_FREE_VAR_REDUCE_STAGE_CACHED,
3067                )
3068                .ok_or_else(|| {
3069                    XlogError::Kernel(
3070                        "xgcf_free_var_reduce_stage_cached kernel not found".to_string(),
3071                    )
3072                })?;
3073            let add_scalar = device
3074                .get_func(
3075                    xlog_cuda::CIRCUIT_MODULE,
3076                    xlog_cuda::circuit_kernels::XGCF_ADD_SCALAR_CACHED,
3077                )
3078                .ok_or_else(|| {
3079                    XlogError::Kernel("xgcf_add_scalar_cached kernel not found".to_string())
3080                })?;
3081
3082            let memory = self.provider.memory();
3083            let mut buf_a = memory.alloc::<f64>(n as usize)?;
3084            let mut buf_b = memory.alloc::<f64>(n as usize)?;
3085
3086            let mut stage_n = n;
3087            let mut stage0 = true;
3088            let mut output_is_a = true;
3089            loop {
3090                let out_len = stage_n.div_ceil(2);
3091                let stage_grid = out_len.div_ceil(block_dim);
3092
3093                let (in_buf, out_buf): (&TrackedCudaSlice<f64>, &mut TrackedCudaSlice<f64>) =
3094                    if output_is_a {
3095                        (&buf_b, &mut buf_a)
3096                    } else {
3097                        (&buf_a, &mut buf_b)
3098                    };
3099                let mode = if stage0 { 0u32 } else { 1u32 };
3100
3101                // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
3102                unsafe {
3103                    reduce_stage.clone().launch(
3104                        LaunchConfig {
3105                            grid_dim: (stage_grid, 1, 1),
3106                            block_dim: (block_dim, 1, 1),
3107                            shared_mem_bytes: 0,
3108                        },
3109                        (
3110                            handle.slot_device(),
3111                            self.var_cap,
3112                            &self.free_var_mask,
3113                            &self.var_log_true,
3114                            &self.var_log_false,
3115                            in_buf,
3116                            stage_n,
3117                            mode,
3118                            out_buf,
3119                        ),
3120                    )
3121                }
3122                .map_err(|e| {
3123                    XlogError::Kernel(format!("xgcf_free_var_reduce_stage_cached failed: {}", e))
3124                })?;
3125
3126                if out_len == 1 {
3127                    let result_buf = if output_is_a { &buf_a } else { &buf_b };
3128                    // SAFETY: kernel arguments match the PTX signature; device buffers were allocated with sufficient size
3129                    unsafe {
3130                        add_scalar.clone().launch(
3131                            LaunchConfig {
3132                                grid_dim: (1, 1, 1),
3133                                block_dim: (1, 1, 1),
3134                                shared_mem_bytes: 0,
3135                            },
3136                            (
3137                                handle.slot_device(),
3138                                self.node_cap,
3139                                &mut self.values,
3140                                &self.meta_root,
3141                                result_buf,
3142                            ),
3143                        )
3144                    }
3145                    .map_err(|e| {
3146                        XlogError::Kernel(format!("xgcf_add_scalar_cached failed: {}", e))
3147                    })?;
3148                    break;
3149                }
3150
3151                stage_n = out_len;
3152                stage0 = false;
3153                output_is_a = !output_is_a;
3154            }
3155        }
3156
3157        Ok(())
3158    }
3159}