1use 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#[derive(Debug, Clone, Copy)]
22#[non_exhaustive]
23pub struct GpuCircuitCacheConfig {
24 pub num_slots: u32,
26 pub table_size: u32,
28 pub node_cap: u32,
30 pub edge_cap: u32,
32 pub level_cap: u32,
34 pub var_cap: u32,
36}
37
38impl Default for GpuCircuitCacheConfig {
39 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
211pub 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 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 Ok(out_hash)
249}
250
251impl GpuCircuitCache {
252 pub fn provider(&self) -> &Arc<CudaKernelProvider> {
253 &self.provider
254 }
255
256 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 pub(crate) fn restore_from_host_arrays(
1622 &mut self,
1623 handle: &mut GpuCircuitCacheHandle,
1624 artifact: &disk_cache::CircuitArtifact,
1625 ) -> Result<()> {
1626 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 handle.num_nodes = num_nodes;
1703 handle.num_levels = num_levels;
1704 handle.root = root;
1705 handle.max_var = max_var;
1706
1707 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 let slot_idx = handle.slot_index() as usize;
2087
2088 let mask_cap = var_stride; 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 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 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 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 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 Ok(())
2199 }
2200
2201 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 Ok(())
2470 }
2471
2472 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 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 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 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 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 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 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 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 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 Ok(())
2791 }
2792
2793 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 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 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 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 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 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 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 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 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 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 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}