1use openvm_cuda_common::{d_buffer::DeviceBuffer, error::CudaError, stream::GpuDeviceCtx};
2use openvm_stark_backend::{StarkProtocolConfig, SystemParams};
3#[cfg(feature = "baby-bear-bn254-poseidon2")]
4use openvm_stark_sdk::config::baby_bear_bn254_poseidon2::{
5 BabyBearBn254Poseidon2Config, Digest as Bn254Digest,
6};
7use openvm_stark_sdk::config::baby_bear_poseidon2::{
8 BabyBearPoseidon2Config, Digest as BabyBearPoseidon2Digest,
9};
10use serde::{de::DeserializeOwned, Serialize};
11
12#[cfg(feature = "baby-bear-bn254-poseidon2")]
13use crate::{
14 bn254_sponge::MultiFieldTranscriptGpu,
15 cuda::bn254_merkle_tree::{
16 bn254_poseidon2_adjacent_compress_layer, bn254_poseidon2_compressing_row_hashes,
17 bn254_poseidon2_compressing_row_hashes_ext,
18 },
19};
20use crate::{
21 cuda::merkle_tree::{
22 poseidon2_adjacent_compress_layer, poseidon2_compressing_row_hashes,
23 poseidon2_compressing_row_hashes_ext,
24 },
25 merkle_tree::BatchQueryMerkle,
26 sponge::{DuplexSpongeGpu, GpuFiatShamirTranscript},
27 types::{EF, F},
28};
29
30pub trait GpuMerkleHash: Copy + Clone + Send + Sync + 'static {
37 type Digest: Copy
38 + Clone
39 + PartialEq
40 + Send
41 + Sync
42 + Serialize
43 + DeserializeOwned
44 + BatchQueryMerkle
45 + 'static;
46
47 unsafe fn compress_rows(
54 out: &mut DeviceBuffer<Self::Digest>,
55 matrix: &DeviceBuffer<F>,
56 width: usize,
57 query_stride: usize,
58 log_rows_per_query: usize,
59 device_ctx: &GpuDeviceCtx,
60 ) -> Result<(), CudaError>;
61
62 unsafe fn compress_rows_ext(
69 out: &mut DeviceBuffer<Self::Digest>,
70 matrix: &DeviceBuffer<EF>,
71 width: usize,
72 query_stride: usize,
73 log_rows_per_query: usize,
74 device_ctx: &GpuDeviceCtx,
75 ) -> Result<(), CudaError>;
76
77 unsafe fn compress_layer(
85 output: &mut DeviceBuffer<Self::Digest>,
86 prev_layer: &DeviceBuffer<Self::Digest>,
87 output_size: usize,
88 device_ctx: &GpuDeviceCtx,
89 ) -> Result<(), CudaError>;
90}
91
92pub trait GpuHashScheme: Copy + Clone + Send + Sync + 'static {
95 type SC: StarkProtocolConfig<F = F, EF = EF, Digest = Self::Digest>;
96 type Digest: Copy
97 + Clone
98 + PartialEq
99 + Send
100 + Sync
101 + Serialize
102 + DeserializeOwned
103 + BatchQueryMerkle
104 + 'static;
105 type Transcript: GpuFiatShamirTranscript<Self::SC> + Default + Clone + Send + Sync + 'static;
106 type MerkleHash: GpuMerkleHash<Digest = Self::Digest>;
107
108 fn default_config(params: SystemParams) -> Self::SC;
109
110 fn default_transcript() -> Self::Transcript;
111}
112
113#[derive(Clone, Copy, Debug, Default)]
119pub struct Poseidon2MerkleHash;
120
121impl GpuMerkleHash for Poseidon2MerkleHash {
122 type Digest = BabyBearPoseidon2Digest;
123
124 unsafe fn compress_rows(
125 out: &mut DeviceBuffer<Self::Digest>,
126 matrix: &DeviceBuffer<F>,
127 width: usize,
128 query_stride: usize,
129 log_rows_per_query: usize,
130 device_ctx: &GpuDeviceCtx,
131 ) -> Result<(), CudaError> {
132 poseidon2_compressing_row_hashes(
133 out,
134 matrix,
135 width,
136 query_stride,
137 log_rows_per_query,
138 device_ctx.stream.as_raw(),
139 )
140 }
141
142 unsafe fn compress_rows_ext(
143 out: &mut DeviceBuffer<Self::Digest>,
144 matrix: &DeviceBuffer<EF>,
145 width: usize,
146 query_stride: usize,
147 log_rows_per_query: usize,
148 device_ctx: &GpuDeviceCtx,
149 ) -> Result<(), CudaError> {
150 poseidon2_compressing_row_hashes_ext(
151 out,
152 matrix,
153 width,
154 query_stride,
155 log_rows_per_query,
156 device_ctx.stream.as_raw(),
157 )
158 }
159
160 unsafe fn compress_layer(
161 output: &mut DeviceBuffer<Self::Digest>,
162 prev_layer: &DeviceBuffer<Self::Digest>,
163 output_size: usize,
164 device_ctx: &GpuDeviceCtx,
165 ) -> Result<(), CudaError> {
166 poseidon2_adjacent_compress_layer(
167 output,
168 prev_layer,
169 output_size,
170 device_ctx.stream.as_raw(),
171 )
172 }
173}
174
175#[derive(Clone, Copy, Debug, Default)]
177pub struct BabyBearPoseidon2HashScheme;
178
179impl GpuHashScheme for BabyBearPoseidon2HashScheme {
180 type SC = BabyBearPoseidon2Config;
181 type Digest = BabyBearPoseidon2Digest;
182 type Transcript = DuplexSpongeGpu;
183 type MerkleHash = Poseidon2MerkleHash;
184
185 fn default_config(params: SystemParams) -> Self::SC {
186 Self::SC::default_from_params(params)
187 }
188
189 fn default_transcript() -> Self::Transcript {
190 Self::Transcript::default()
191 }
192}
193
194pub type DefaultHashScheme = BabyBearPoseidon2HashScheme;
195
196#[cfg(feature = "baby-bear-bn254-poseidon2")]
201#[derive(Clone, Copy, Debug, Default)]
203pub struct Bn254Poseidon2MerkleHash;
204
205#[cfg(feature = "baby-bear-bn254-poseidon2")]
206impl GpuMerkleHash for Bn254Poseidon2MerkleHash {
207 type Digest = Bn254Digest;
211
212 unsafe fn compress_rows(
213 out: &mut DeviceBuffer<Self::Digest>,
214 matrix: &DeviceBuffer<F>,
215 width: usize,
216 query_stride: usize,
217 log_rows_per_query: usize,
218 device_ctx: &GpuDeviceCtx,
219 ) -> Result<(), CudaError> {
220 bn254_poseidon2_compressing_row_hashes(
221 out,
222 matrix,
223 width,
224 query_stride,
225 log_rows_per_query,
226 device_ctx.stream.as_raw(),
227 )
228 }
229
230 unsafe fn compress_rows_ext(
231 out: &mut DeviceBuffer<Self::Digest>,
232 matrix: &DeviceBuffer<EF>,
233 width: usize,
234 query_stride: usize,
235 log_rows_per_query: usize,
236 device_ctx: &GpuDeviceCtx,
237 ) -> Result<(), CudaError> {
238 bn254_poseidon2_compressing_row_hashes_ext(
239 out,
240 matrix,
241 width,
242 query_stride,
243 log_rows_per_query,
244 device_ctx.stream.as_raw(),
245 )
246 }
247
248 unsafe fn compress_layer(
249 output: &mut DeviceBuffer<Self::Digest>,
250 prev_layer: &DeviceBuffer<Self::Digest>,
251 output_size: usize,
252 device_ctx: &GpuDeviceCtx,
253 ) -> Result<(), CudaError> {
254 bn254_poseidon2_adjacent_compress_layer(
255 output,
256 prev_layer,
257 output_size,
258 device_ctx.stream.as_raw(),
259 )
260 }
261}
262
263#[cfg(feature = "baby-bear-bn254-poseidon2")]
264#[derive(Clone, Copy, Debug, Default)]
266pub struct BabyBearBn254Poseidon2HashScheme;
267
268#[cfg(feature = "baby-bear-bn254-poseidon2")]
269impl GpuHashScheme for BabyBearBn254Poseidon2HashScheme {
270 type SC = BabyBearBn254Poseidon2Config;
271 type Digest = Bn254Digest;
272 type Transcript = MultiFieldTranscriptGpu;
273 type MerkleHash = Bn254Poseidon2MerkleHash;
274
275 fn default_config(params: SystemParams) -> Self::SC {
276 Self::SC::default_from_params(params)
277 }
278
279 fn default_transcript() -> Self::Transcript {
280 Self::Transcript::default()
281 }
282}