openvm_deferral_circuit/poseidon2/
cuda.rs1use 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 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}