openvm_circuit/system/cuda/
poseidon2.rs

1#[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    /// Prepare an exact one-segment scratch buffer for Poseidon2 records.
57    ///
58    /// Each Poseidon2 record occupies `POSEIDON2_WIDTH` field elements.
59    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        // Reset state of this chip.
172        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}