Skip to main content

openvm_cuda_backend/
hash_scheme.rs

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
30/// Dispatch trait for GPU Merkle hash kernels.
31///
32/// Each implementation routes the three kernel entry points
33/// (`compress_rows`, `compress_rows_ext`, `compress_layer`) to the
34/// appropriate CUDA FFI wrappers, and declares the concrete `Digest` type
35/// those kernels produce.
36pub 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    /// Compress rows of a base-field matrix into digest leaves.
48    ///
49    /// # Safety
50    ///
51    /// `out` must be allocated with capacity `query_stride` and `matrix` must
52    /// contain `width * query_stride * (1 << log_rows_per_query)` valid elements.
53    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    /// Compress rows of an extension-field matrix into digest leaves.
63    ///
64    /// # Safety
65    ///
66    /// `out` must be allocated with capacity `query_stride` and `matrix` must
67    /// contain `width * query_stride * (1 << log_rows_per_query)` valid elements.
68    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    /// Compress adjacent pairs of digests to build an inner Merkle layer.
78    ///
79    /// # Safety
80    ///
81    /// `output` must be allocated with capacity `output_size`, `prev_layer` must
82    /// contain at least `output_size * 2` valid elements, and the two buffers
83    /// must not overlap.
84    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
92/// Binding trait that couples a `StarkProtocolConfig`, a Merkle hash scheme,
93/// and a transcript type into a single coherent GPU proving configuration.
94pub 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// ---------------------------------------------------------------------------
114// Poseidon2 / BabyBear concrete implementations
115// ---------------------------------------------------------------------------
116
117/// Poseidon2 Merkle hash over BabyBear — delegates to the existing CUDA FFI.
118#[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/// BabyBear Poseidon2 hash scheme — the only scheme implemented in this crate.
176#[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// ---------------------------------------------------------------------------
197// BN254 Poseidon2 concrete implementations
198// ---------------------------------------------------------------------------
199
200#[cfg(feature = "baby-bear-bn254-poseidon2")]
201/// BN254 Poseidon2 Merkle hash — delegates to the BN254 CUDA FFI.
202#[derive(Clone, Copy, Debug, Default)]
203pub struct Bn254Poseidon2MerkleHash;
204
205#[cfg(feature = "baby-bear-bn254-poseidon2")]
206impl GpuMerkleHash for Bn254Poseidon2MerkleHash {
207    // `Bn254Digest` from stark-sdk = `[Bn254Scalar; 1]`, which is the same concrete type as
208    // `Bn254Digest` in `cuda::bn254_merkle_tree` — both are type aliases for `[p3_bn254::Bn254;
209    // 1]`.
210    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/// BabyBear + BN254 Poseidon2 hash scheme (Groth16-friendly transcript).
265#[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}