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