openvm_deferral_circuit/poseidon2/
cuda.rs

1use std::sync::Arc;
2
3use openvm_circuit::{arch::DenseRecordArena, utils::next_power_of_two_or_zero};
4use openvm_circuit_primitives::Chip;
5use openvm_cuda_backend::{base::DeviceMatrix, prelude::F, GpuBackend};
6use openvm_cuda_common::{
7    copy::{MemCopyD2H, MemCopyH2D},
8    d_buffer::DeviceBuffer,
9    stream::GpuDeviceCtx,
10};
11use openvm_stark_backend::prover::{AirProvingContext, MatrixDimensions};
12use openvm_stark_sdk::config::baby_bear_poseidon2::DIGEST_SIZE;
13
14use crate::{
15    cuda_abi::poseidon2::{self, DeferralPoseidon2Count},
16    poseidon2::DeferralPoseidon2Cols,
17};
18
19#[derive(Clone)]
20pub struct DeferralPoseidon2SharedBuffer {
21    pub records: Arc<DeviceBuffer<F>>,
22    pub counts: Arc<DeviceBuffer<DeferralPoseidon2Count>>,
23    pub idx: Arc<DeviceBuffer<u32>>,
24}
25
26pub struct DeferralPoseidon2ChipGpu {
27    pub device_ctx: GpuDeviceCtx,
28    pub records: Arc<DeviceBuffer<F>>,
29    pub counts: Arc<DeviceBuffer<DeferralPoseidon2Count>>,
30    pub idx: Arc<DeviceBuffer<u32>>,
31    pub sbox_registers: usize,
32}
33
34impl DeferralPoseidon2ChipGpu {
35    /// Creates a new deferral Poseidon2 chip configured for `max_trace_height` records. Each
36    /// Poseidon2 record occupies `POSEIDON2_WIDTH` (16) field elements, and a buffer of that
37    /// size is allocated.
38    pub fn new(max_trace_height: usize, sbox_registers: usize, device_ctx: GpuDeviceCtx) -> Self {
39        let max_num_records = max_trace_height.next_power_of_two();
40        let max_record_buf_size = max_num_records * (DIGEST_SIZE * 2);
41
42        let idx = Arc::new(DeviceBuffer::<u32>::with_capacity_on(1, &device_ctx));
43        idx.fill_zero_on(&device_ctx).unwrap();
44
45        Self {
46            device_ctx: device_ctx.clone(),
47            records: Arc::new(DeviceBuffer::<F>::with_capacity_on(
48                max_record_buf_size,
49                &device_ctx,
50            )),
51            counts: Arc::new(DeviceBuffer::<DeferralPoseidon2Count>::with_capacity_on(
52                max_num_records,
53                &device_ctx,
54            )),
55            idx,
56            sbox_registers,
57        }
58    }
59
60    pub fn shared_buffer(&self) -> DeferralPoseidon2SharedBuffer {
61        DeferralPoseidon2SharedBuffer {
62            records: self.records.clone(),
63            counts: self.counts.clone(),
64            idx: self.idx.clone(),
65        }
66    }
67
68    pub fn trace_width() -> usize {
69        DeferralPoseidon2Cols::<F>::width()
70    }
71}
72
73impl Chip<DenseRecordArena, GpuBackend> for DeferralPoseidon2ChipGpu {
74    fn generate_proving_ctx(&self, _: DenseRecordArena) -> AirProvingContext<GpuBackend> {
75        let mut num_records = self.idx.to_host_on(&self.device_ctx).unwrap()[0] as usize;
76        if num_records == 0 {
77            return AirProvingContext::simple_no_pis(DeviceMatrix::dummy());
78        }
79
80        let dedup_records =
81            DeviceBuffer::<F>::with_capacity_on(num_records * DIGEST_SIZE * 2, &self.device_ctx);
82        let dedup_counts =
83            DeviceBuffer::<DeferralPoseidon2Count>::with_capacity_on(num_records, &self.device_ctx);
84        unsafe {
85            let d_num_records = [num_records].to_device_on(&self.device_ctx).unwrap();
86            let mut temp_bytes = 0;
87            poseidon2::deduplicate_records_get_temp_bytes(
88                &self.records,
89                &self.counts,
90                num_records,
91                &d_num_records,
92                &mut temp_bytes,
93                self.device_ctx.stream.as_raw(),
94            )
95            .expect("Failed to get deferral poseidon2 temp bytes");
96
97            let d_temp_storage = if temp_bytes == 0 {
98                DeviceBuffer::<u8>::new()
99            } else {
100                DeviceBuffer::<u8>::with_capacity_on(temp_bytes, &self.device_ctx)
101            };
102
103            poseidon2::deduplicate_records(
104                &self.records,
105                &self.counts,
106                &dedup_records,
107                &dedup_counts,
108                num_records,
109                &d_num_records,
110                &d_temp_storage,
111                temp_bytes,
112                self.device_ctx.stream.as_raw(),
113            )
114            .expect("Failed to deduplicate deferral poseidon2 records");
115
116            num_records = *d_num_records
117                .to_host_on(&self.device_ctx)
118                .unwrap()
119                .first()
120                .unwrap();
121        }
122
123        let trace_height = next_power_of_two_or_zero(num_records);
124        let trace = DeviceMatrix::<F>::with_capacity_on(
125            trace_height,
126            Self::trace_width(),
127            &self.device_ctx,
128        );
129
130        unsafe {
131            poseidon2::tracegen(
132                trace.buffer(),
133                trace.height(),
134                trace.width(),
135                &dedup_records,
136                &dedup_counts,
137                num_records,
138                self.sbox_registers,
139                self.device_ctx.stream.as_raw(),
140            )
141            .expect("Failed to generate deferral poseidon2 trace");
142        }
143
144        self.idx
145            .fill_zero_on(&self.device_ctx)
146            .expect("Failed to reset deferral poseidon2 record index");
147
148        AirProvingContext::simple_no_pis(trace)
149    }
150}