openvm_deferral_circuit/call/
cuda.rs1use std::{mem::size_of, sync::Arc};
2
3use derive_new::new;
4use openvm_circuit::{arch::DenseRecordArena, utils::next_power_of_two_or_zero};
5use openvm_circuit_primitives::{
6 bitwise_op_lookup::BitwiseOperationLookupChipGPU, var_range::VariableRangeCheckerChipGPU, Chip,
7};
8use openvm_cuda_backend::{base::DeviceMatrix, prelude::F, GpuBackend};
9use openvm_cuda_common::{copy::MemCopyH2D, d_buffer::DeviceBuffer};
10use openvm_instructions::riscv::RV32_CELL_BITS;
11use openvm_stark_backend::prover::AirProvingContext;
12
13use super::{
14 DeferralCallAdapterCols, DeferralCallAdapterRecord, DeferralCallCoreCols,
15 DeferralCallCoreRecord,
16};
17use crate::{cuda_abi::call, poseidon2::DeferralPoseidon2SharedBuffer};
18
19#[derive(new)]
20pub struct DeferralCallChipGpu {
21 pub range_checker: Arc<VariableRangeCheckerChipGPU>,
22 pub bitwise_lookup: Arc<BitwiseOperationLookupChipGPU<RV32_CELL_BITS>>,
23 pub address_bits: usize,
24 pub timestamp_max_bits: usize,
25 pub count: Arc<DeviceBuffer<u32>>,
26 pub num_deferral_circuits: usize,
27 pub poseidon2: DeferralPoseidon2SharedBuffer,
28}
29
30impl Chip<DenseRecordArena, GpuBackend> for DeferralCallChipGpu {
31 fn generate_proving_ctx(&self, arena: DenseRecordArena) -> AirProvingContext<GpuBackend> {
32 type Record = (DeferralCallAdapterRecord<F>, DeferralCallCoreRecord<F>);
33 const RECORD_SIZE: usize = size_of::<Record>();
34
35 let records = arena.allocated();
36 if records.is_empty() {
37 return AirProvingContext::simple_no_pis(DeviceMatrix::dummy());
38 }
39 debug_assert_eq!(records.len() % RECORD_SIZE, 0);
40
41 let num_records = records.len() / RECORD_SIZE;
42 let trace_height = next_power_of_two_or_zero(num_records);
43 let trace_width =
44 DeferralCallAdapterCols::<F>::width() + DeferralCallCoreCols::<F>::width();
45 let device_ctx = &self.range_checker.device_ctx;
46
47 let d_records = records.to_device_on(device_ctx).unwrap();
48 let trace = DeviceMatrix::<F>::with_capacity_on(trace_height, trace_width, device_ctx);
49
50 unsafe {
51 call::tracegen(
52 trace.buffer(),
53 trace_height,
54 trace_width,
55 &d_records,
56 num_records,
57 &self.count,
58 self.num_deferral_circuits,
59 &self.range_checker.count,
60 self.timestamp_max_bits as u32,
61 &self.bitwise_lookup.count,
62 RV32_CELL_BITS as u32,
63 &self.poseidon2.records,
64 &self.poseidon2.counts,
65 &self.poseidon2.idx,
66 self.poseidon2.records.len(),
68 self.address_bits,
69 device_ctx.stream.as_raw(),
70 )
71 .expect("Failed to generate deferral call trace");
72 }
73
74 AirProvingContext::simple_no_pis(trace)
75 }
76}