1#![allow(clippy::missing_safety_doc)]
2
3use derive_new::new;
4use openvm_cuda_backend::prelude::F;
5use openvm_cuda_common::{d_buffer::DeviceBuffer, error::CudaError, stream::cudaStream_t};
6
7#[repr(C)]
9#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, new)]
10pub struct UInt2 {
11 pub x: u32,
12 pub y: u32,
13}
14
15pub mod bitwise_op_lookup {
16 #[cfg(test)]
17 use openvm_cuda_common::d_buffer::DeviceBufferView;
18
19 use super::*;
20
21 extern "C" {
22 fn _bitwise_op_lookup_tracegen(
23 d_count: *const u32,
24 d_cpu_count: *const u32,
25 d_trace: *mut F,
26 num_bits: u32,
27 stream: cudaStream_t,
28 ) -> i32;
29
30 #[cfg(test)]
31 fn _bitwise_dummy_tracegen(
32 d_trace: *mut F,
33 records: DeviceBufferView,
34 bitwise_count: *mut u32,
35 bitwise_num_bits: u32,
36 stream: cudaStream_t,
37 ) -> i32;
38 }
39
40 pub unsafe fn tracegen(
41 d_count: &DeviceBuffer<F>,
42 d_cpu_count: &Option<DeviceBuffer<u32>>,
43 d_trace: &DeviceBuffer<F>,
44 num_bits: u32,
45 stream: cudaStream_t,
46 ) -> Result<(), CudaError> {
47 CudaError::from_result(_bitwise_op_lookup_tracegen(
48 d_count.as_ptr() as *const u32,
49 d_cpu_count
50 .as_ref()
51 .map(|b| b.as_ptr())
52 .unwrap_or(std::ptr::null()),
53 d_trace.as_mut_ptr(),
54 num_bits,
55 stream,
56 ))
57 }
58
59 #[cfg(test)]
60 pub unsafe fn dummy_tracegen(
61 d_trace: &DeviceBuffer<F>,
62 records: &DeviceBuffer<u32>,
63 bitwise_count: &DeviceBuffer<F>,
64 bitwise_num_bits: u32,
65 stream: cudaStream_t,
66 ) -> Result<(), CudaError> {
67 CudaError::from_result(_bitwise_dummy_tracegen(
68 d_trace.as_mut_ptr(),
69 records.view(),
70 bitwise_count.as_mut_ptr() as *mut u32,
71 bitwise_num_bits,
72 stream,
73 ))
74 }
75}
76
77pub mod range_tuple {
78 use super::*;
79
80 extern "C" {
81 fn _range_tuple_checker_tracegen(
82 d_count: *const u32,
83 d_cpu_count: *const u32,
84 d_trace: *mut F,
85 d_sizes: *const u32,
86 num_dims: u32,
87 num_bins: usize,
88 stream: cudaStream_t,
89 ) -> i32;
90
91 #[cfg(test)]
92 fn _range_tuple_dummy_tracegen(
93 d_data: *const u32,
94 d_trace: *mut F,
95 d_rc_count: *mut u32,
96 data_height: usize,
97 sizes: *const u32,
98 num_sizes: usize,
99 stream: cudaStream_t,
100 ) -> i32;
101 }
102
103 pub unsafe fn tracegen(
104 d_count: &DeviceBuffer<F>,
105 d_cpu_count: &Option<DeviceBuffer<u32>>,
106 d_trace: &DeviceBuffer<F>,
107 d_sizes: &DeviceBuffer<u32>,
108 stream: cudaStream_t,
109 ) -> Result<(), CudaError> {
110 CudaError::from_result(_range_tuple_checker_tracegen(
111 d_count.as_ptr() as *const u32,
112 d_cpu_count
113 .as_ref()
114 .map(|b| b.as_ptr())
115 .unwrap_or(std::ptr::null()),
116 d_trace.as_mut_ptr(),
117 d_sizes.as_ptr(),
118 d_sizes.len() as u32,
119 d_count.len(),
120 stream,
121 ))
122 }
123
124 #[cfg(test)]
125 pub unsafe fn dummy_tracegen(
126 d_data: &DeviceBuffer<u32>,
127 d_trace: &DeviceBuffer<F>,
128 d_rc_count: &DeviceBuffer<F>,
129 sizes: &DeviceBuffer<u32>,
130 stream: cudaStream_t,
131 ) -> Result<(), CudaError> {
132 CudaError::from_result(_range_tuple_dummy_tracegen(
133 d_data.as_ptr(),
134 d_trace.as_mut_ptr(),
135 d_rc_count.as_mut_ptr() as *mut u32,
136 d_data.len() / sizes.len(),
137 sizes.as_ptr(),
138 sizes.len(),
139 stream,
140 ))
141 }
142}
143
144pub mod var_range {
145 use super::*;
146
147 extern "C" {
148 fn _range_checker_tracegen(
149 d_count: *const u32,
150 d_cpu_count: *const u32,
151 d_trace: *mut F,
152 num_bins: usize,
153 stream: cudaStream_t,
154 ) -> i32;
155
156 #[cfg(test)]
157 fn _var_range_dummy_tracegen(
158 d_data: *const u32,
159 d_trace: *mut F,
160 d_rc_count: *mut u32,
161 data_len: usize,
162 range_max_bits: usize,
163 stream: cudaStream_t,
164 ) -> i32;
165 }
166
167 pub unsafe fn tracegen(
168 d_count: &DeviceBuffer<F>,
169 d_cpu_count: &Option<DeviceBuffer<u32>>,
170 d_trace: &DeviceBuffer<F>,
171 stream: cudaStream_t,
172 ) -> Result<(), CudaError> {
173 CudaError::from_result(_range_checker_tracegen(
174 d_count.as_ptr() as *const u32,
175 d_cpu_count
176 .as_ref()
177 .map(|b| b.as_ptr())
178 .unwrap_or(std::ptr::null()),
179 d_trace.as_mut_ptr(),
180 d_count.len(),
181 stream,
182 ))
183 }
184
185 #[cfg(test)]
186 pub unsafe fn dummy_tracegen(
187 d_data: &DeviceBuffer<u32>,
188 d_trace: &DeviceBuffer<F>,
189 d_rc_count: &DeviceBuffer<F>,
190 stream: cudaStream_t,
191 ) -> Result<(), CudaError> {
192 CudaError::from_result(_var_range_dummy_tracegen(
193 d_data.as_ptr(),
194 d_trace.as_mut_ptr(),
195 d_rc_count.as_mut_ptr() as *mut u32,
196 d_data.len(),
197 d_rc_count.len(),
198 stream,
199 ))
200 }
201}
202
203#[cfg(test)]
204pub mod encoder {
205 use super::*;
206
207 extern "C" {
208 fn _encoder_tracegen(
209 trace: *mut F,
210 num_flags: u32,
211 max_degree: u32,
212 reserve_invalid: bool,
213 expected_k: u32,
214 stream: cudaStream_t,
215 ) -> i32;
216 }
217
218 pub unsafe fn dummy_tracegen(
219 d_trace: &DeviceBuffer<F>,
220 num_flags: u32,
221 max_degree: u32,
222 reserve_invalid: bool,
223 expected_k: u32,
224 stream: cudaStream_t,
225 ) -> Result<(), CudaError> {
226 CudaError::from_result(_encoder_tracegen(
227 d_trace.as_mut_ptr(),
228 num_flags,
229 max_degree,
230 reserve_invalid,
231 expected_k,
232 stream,
233 ))
234 }
235}
236
237#[cfg(test)]
238pub mod is_equal {
239 use super::*;
240
241 extern "C" {
242 fn _isequal_tracegen(
243 output: *mut F,
244 inputs_x: *mut F,
245 inputs_y: *mut F,
246 n: u32,
247 stream: cudaStream_t,
248 ) -> i32;
249
250 fn _isequal_array_tracegen(
251 output: *mut F,
252 inputs_x: *mut F,
253 inputs_y: *mut F,
254 array_len: u32,
255 n: u32,
256 stream: cudaStream_t,
257 ) -> i32;
258 }
259
260 pub unsafe fn dummy_tracegen(
261 d_output: &DeviceBuffer<F>,
262 d_inputs_x: &DeviceBuffer<F>,
263 d_inputs_y: &DeviceBuffer<F>,
264 stream: cudaStream_t,
265 ) -> Result<(), CudaError> {
266 CudaError::from_result(_isequal_tracegen(
267 d_output.as_mut_ptr(),
268 d_inputs_x.as_mut_ptr(),
269 d_inputs_y.as_mut_ptr(),
270 d_inputs_x.len() as u32,
271 stream,
272 ))
273 }
274
275 pub unsafe fn dummy_tracegen_array(
276 d_output: &DeviceBuffer<F>,
277 d_inputs_x: &DeviceBuffer<F>,
278 d_inputs_y: &DeviceBuffer<F>,
279 array_len: usize,
280 stream: cudaStream_t,
281 ) -> Result<(), CudaError> {
282 CudaError::from_result(_isequal_array_tracegen(
283 d_output.as_mut_ptr(),
284 d_inputs_x.as_mut_ptr(),
285 d_inputs_y.as_mut_ptr(),
286 array_len as u32,
287 (d_inputs_x.len() / array_len) as u32,
288 stream,
289 ))
290 }
291}
292
293#[cfg(test)]
294pub mod is_zero {
295 use super::*;
296
297 extern "C" {
298 fn _iszero_tracegen(output: *mut F, inputs: *mut F, n: u32, stream: cudaStream_t) -> i32;
299 }
300
301 pub unsafe fn dummy_tracegen(
302 d_output: &DeviceBuffer<F>,
303 d_inputs: &DeviceBuffer<F>,
304 stream: cudaStream_t,
305 ) -> Result<(), CudaError> {
306 CudaError::from_result(_iszero_tracegen(
307 d_output.as_mut_ptr(),
308 d_inputs.as_mut_ptr(),
309 d_inputs.len() as u32,
310 stream,
311 ))
312 }
313}
314
315#[cfg(test)]
316pub mod less_than {
317 use super::*;
318
319 extern "C" {
320 fn _assert_less_than_tracegen(
321 trace: *mut F,
322 trace_height: usize,
323 pairs: *const u32,
324 max_bits: u32,
325 aux_len: u32,
326 rc_count: *mut u32,
327 rc_num_bins: u32,
328 stream: cudaStream_t,
329 ) -> i32;
330
331 fn _less_than_tracegen(
332 trace: *mut F,
333 trace_height: usize,
334 pairs: *const u32,
335 max_bits: u32,
336 aux_len: u32,
337 rc_count: *mut u32,
338 rc_num_bins: u32,
339 stream: cudaStream_t,
340 ) -> i32;
341
342 fn _less_than_array_tracegen(
343 trace: *mut F,
344 trace_height: usize,
345 pairs: *const u32,
346 max_bits: u32,
347 array_len: u32,
348 aux_len: u32,
349 rc_count: *mut u32,
350 rc_num_bins: u32,
351 stream: cudaStream_t,
352 ) -> i32;
353 }
354
355 pub unsafe fn assert_less_than_dummy_tracegen(
356 trace: &DeviceBuffer<F>,
357 trace_height: usize,
358 pairs: &DeviceBuffer<u32>,
359 max_bits: usize,
360 aux_len: usize,
361 rc_count: &DeviceBuffer<u32>,
362 stream: cudaStream_t,
363 ) -> Result<(), CudaError> {
364 CudaError::from_result(_assert_less_than_tracegen(
365 trace.as_mut_ptr(),
366 trace_height,
367 pairs.as_ptr(),
368 max_bits as u32,
369 aux_len as u32,
370 rc_count.as_mut_ptr(),
371 rc_count.len() as u32,
372 stream,
373 ))
374 }
375
376 pub unsafe fn less_than_dummy_tracegen(
377 trace: &DeviceBuffer<F>,
378 trace_height: usize,
379 pairs: &DeviceBuffer<u32>,
380 max_bits: usize,
381 aux_len: usize,
382 rc_count: &DeviceBuffer<u32>,
383 stream: cudaStream_t,
384 ) -> Result<(), CudaError> {
385 CudaError::from_result(_less_than_tracegen(
386 trace.as_mut_ptr(),
387 trace_height,
388 pairs.as_ptr(),
389 max_bits as u32,
390 aux_len as u32,
391 rc_count.as_mut_ptr(),
392 rc_count.len() as u32,
393 stream,
394 ))
395 }
396
397 #[allow(clippy::too_many_arguments)]
398 pub unsafe fn less_than_array_dummy_tracegen(
399 trace: &DeviceBuffer<F>,
400 trace_height: usize,
401 pairs: &DeviceBuffer<u32>,
402 max_bits: usize,
403 array_len: usize,
404 aux_len: usize,
405 rc_count: &DeviceBuffer<u32>,
406 stream: cudaStream_t,
407 ) -> Result<(), CudaError> {
408 CudaError::from_result(_less_than_array_tracegen(
409 trace.as_mut_ptr(),
410 trace_height,
411 pairs.as_ptr(),
412 max_bits as u32,
413 array_len as u32,
414 aux_len as u32,
415 rc_count.as_mut_ptr(),
416 rc_count.len() as u32,
417 stream,
418 ))
419 }
420}