openvm_circuit/system/cuda/
poseidon2.rs1#[cfg(feature = "metrics")]
2use std::sync::atomic::AtomicUsize;
3use std::sync::{Arc, Mutex};
4
5use openvm_circuit::{
6 primitives::Chip, system::poseidon2::columns::Poseidon2PeripheryCols,
7 utils::next_power_of_two_or_zero,
8};
9use openvm_cuda_backend::{base::DeviceMatrix, prelude::F, GpuBackend};
10use openvm_cuda_common::{
11 copy::{MemCopyD2H, MemCopyH2D},
12 d_buffer::DeviceBuffer,
13 stream::GpuDeviceCtx,
14};
15use openvm_poseidon2_air::POSEIDON2_WIDTH;
16use openvm_stark_backend::prover::{AirProvingContext, MatrixDimensions};
17
18use crate::cuda_abi::poseidon2;
19
20#[derive(Clone)]
21pub struct SharedBuffer<T> {
22 buffer: Arc<Mutex<Option<Arc<DeviceBuffer<T>>>>>,
23 pub idx: Arc<DeviceBuffer<u32>>,
24}
25
26impl<T> SharedBuffer<T> {
27 pub fn records(&self) -> Arc<DeviceBuffer<T>> {
28 let records = self.buffer.lock().unwrap();
29 records
30 .clone()
31 .expect("Poseidon2 records buffer must be prepared before tracegen")
32 }
33}
34
35pub struct Poseidon2ChipGPU<const SBOX_REGISTERS: usize> {
36 pub device_ctx: GpuDeviceCtx,
37 pub records: Arc<Mutex<Option<Arc<DeviceBuffer<F>>>>>,
38 pub idx: Arc<DeviceBuffer<u32>>,
39 #[cfg(feature = "metrics")]
40 pub(crate) current_trace_height: Arc<AtomicUsize>,
41}
42
43impl<const SBOX_REGISTERS: usize> Poseidon2ChipGPU<SBOX_REGISTERS> {
44 pub fn new(device_ctx: GpuDeviceCtx) -> Self {
45 let idx = Arc::new(DeviceBuffer::<u32>::with_capacity_on(1, &device_ctx));
46 idx.fill_zero_on(&device_ctx).unwrap();
47 Self {
48 device_ctx: device_ctx.clone(),
49 records: Arc::new(Mutex::new(None)),
50 idx,
51 #[cfg(feature = "metrics")]
52 current_trace_height: Arc::new(AtomicUsize::new(0)),
53 }
54 }
55
56 pub fn prepare_records(&self, num_records: usize) {
60 self.idx.fill_zero_on(&self.device_ctx).unwrap();
61 let mut records = self.records.lock().unwrap();
62 assert!(
63 records.is_none(),
64 "Poseidon2 records buffer already prepared"
65 );
66 if num_records == 0 {
67 return;
68 }
69 let num_elements = num_records
70 .checked_mul(POSEIDON2_WIDTH)
71 .expect("Poseidon2 records buffer size overflow");
72 records.replace(Arc::new(DeviceBuffer::<F>::with_capacity_on(
73 num_elements,
74 &self.device_ctx,
75 )));
76 }
77
78 pub fn shared_buffer(&self) -> SharedBuffer<F> {
79 SharedBuffer {
80 buffer: self.records.clone(),
81 idx: self.idx.clone(),
82 }
83 }
84
85 pub fn trace_width() -> usize {
86 Poseidon2PeripheryCols::<F, SBOX_REGISTERS>::width()
87 }
88}
89
90impl<RA, const SBOX_REGISTERS: usize> Chip<RA, GpuBackend> for Poseidon2ChipGPU<SBOX_REGISTERS> {
91 fn generate_proving_ctx(&self, _: RA) -> AirProvingContext<GpuBackend> {
92 let Some(records) = self.records.lock().unwrap().take() else {
93 self.idx.fill_zero_on(&self.device_ctx).unwrap();
94 return AirProvingContext::simple_no_pis(DeviceMatrix::dummy());
95 };
96 debug_assert_eq!(records.len() % POSEIDON2_WIDTH, 0);
97 let capacity_records = records.len() / POSEIDON2_WIDTH;
98 let mut num_records = self.idx.to_host_on(&self.device_ctx).unwrap()[0] as usize;
99 assert!(
100 num_records <= capacity_records,
101 "Poseidon2 records buffer overflow: pushed {num_records} records into capacity {capacity_records}"
102 );
103 if num_records == 0 {
104 self.idx.fill_zero_on(&self.device_ctx).unwrap();
105 return AirProvingContext::simple_no_pis(DeviceMatrix::dummy());
106 }
107 let counts = DeviceBuffer::<u32>::with_capacity_on(num_records, &self.device_ctx);
108 let dedup_records =
109 DeviceBuffer::<F>::with_capacity_on(num_records * POSEIDON2_WIDTH, &self.device_ctx);
110 let dedup_counts = DeviceBuffer::<u32>::with_capacity_on(num_records, &self.device_ctx);
111 unsafe {
112 let d_num_records = [num_records].to_device_on(&self.device_ctx).unwrap();
113 let mut temp_bytes = 0;
114 poseidon2::deduplicate_records_get_temp_bytes(
115 &records,
116 &counts,
117 num_records,
118 &d_num_records,
119 &mut temp_bytes,
120 self.device_ctx.stream.as_raw(),
121 )
122 .expect("Failed to get temp bytes");
123 let d_temp_storage = if temp_bytes == 0 {
124 DeviceBuffer::<u8>::new()
125 } else {
126 DeviceBuffer::<u8>::with_capacity_on(temp_bytes, &self.device_ctx)
127 };
128 poseidon2::deduplicate_records(
129 &records,
130 &counts,
131 &dedup_records,
132 &dedup_counts,
133 num_records,
134 &d_num_records,
135 &d_temp_storage,
136 temp_bytes,
137 self.device_ctx.stream.as_raw(),
138 )
139 .expect("Failed to deduplicate records");
140 num_records = *d_num_records
141 .to_host_on(&self.device_ctx)
142 .unwrap()
143 .first()
144 .unwrap();
145 }
146 drop(records);
147 drop(counts);
148 #[cfg(feature = "metrics")]
149 self.current_trace_height
150 .store(num_records, std::sync::atomic::Ordering::Relaxed);
151 let trace_height = next_power_of_two_or_zero(num_records);
152 let trace = DeviceMatrix::<F>::with_capacity_on(
153 trace_height,
154 Self::trace_width(),
155 &self.device_ctx,
156 );
157 trace.buffer().fill_zero_on(&self.device_ctx).unwrap();
158 unsafe {
159 poseidon2::tracegen(
160 trace.buffer(),
161 trace.height(),
162 trace.width(),
163 &dedup_records,
164 &dedup_counts,
165 num_records,
166 SBOX_REGISTERS,
167 self.device_ctx.stream.as_raw(),
168 )
169 .expect("Failed to generate trace");
170 }
171 self.idx.fill_zero_on(&self.device_ctx).unwrap();
173 AirProvingContext::simple_no_pis(trace)
174 }
175}
176
177pub enum Poseidon2PeripheryChipGPU {
178 Register0(Poseidon2ChipGPU<0>),
179 Register1(Poseidon2ChipGPU<1>),
180}
181
182impl Poseidon2PeripheryChipGPU {
183 pub fn new(sbox_registers: usize, device_ctx: GpuDeviceCtx) -> Self {
184 match sbox_registers {
185 0 => Self::Register0(Poseidon2ChipGPU::new(device_ctx)),
186 1 => Self::Register1(Poseidon2ChipGPU::new(device_ctx)),
187 _ => panic!("Invalid number of sbox registers: {sbox_registers}"),
188 }
189 }
190
191 pub fn prepare_records(&self, num_records: usize) {
192 match self {
193 Self::Register0(chip) => chip.prepare_records(num_records),
194 Self::Register1(chip) => chip.prepare_records(num_records),
195 }
196 }
197
198 pub fn shared_buffer(&self) -> SharedBuffer<F> {
199 match self {
200 Self::Register0(chip) => chip.shared_buffer(),
201 Self::Register1(chip) => chip.shared_buffer(),
202 }
203 }
204
205 pub fn device_ctx(&self) -> &GpuDeviceCtx {
206 match self {
207 Self::Register0(chip) => &chip.device_ctx,
208 Self::Register1(chip) => &chip.device_ctx,
209 }
210 }
211}
212
213impl<RA> Chip<RA, GpuBackend> for Poseidon2PeripheryChipGPU {
214 fn generate_proving_ctx(&self, _: RA) -> AirProvingContext<GpuBackend> {
215 match self {
216 Self::Register0(chip) => chip.generate_proving_ctx(()),
217 Self::Register1(chip) => chip.generate_proving_ctx(()),
218 }
219 }
220}