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