Skip to main content

openvm_cuda_backend/
merkle_tree.rs

1use std::array::from_fn;
2
3use itertools::Itertools;
4use openvm_cuda_common::{
5    copy::{MemCopyD2H, MemCopyH2D},
6    d_buffer::DeviceBuffer,
7    memory_manager::MemTracker,
8    stream::GpuDeviceCtx,
9};
10use openvm_stark_backend::prover::MatrixDimensions;
11use p3_util::log2_strict_usize;
12use tracing::{info_span, instrument};
13
14#[cfg(feature = "baby-bear-bn254-poseidon2")]
15use crate::cuda::bn254_merkle_tree::Bn254Digest;
16use crate::{
17    base::DeviceMatrix,
18    cuda::{matrix::matrix_get_rows_fp_kernel, merkle_tree::query_digest_layers},
19    hash_scheme::{GpuMerkleHash, Poseidon2MerkleHash},
20    prelude::{Digest, DIGEST_SIZE, EF, F},
21    MerkleTreeError,
22};
23
24/// Trait for reconstructing a digest from a flat slice of F elements as
25/// produced by the `query_digest_layers` CUDA kernel.
26///
27/// Both `Digest = [F; 8]` (BabyBear Poseidon2) and `Bn254Digest = [Bn254Scalar; 1]`
28/// occupy exactly `DIGEST_SIZE * size_of::<F>()` = 32 bytes, so the same kernel
29/// can be reused for both.
30pub trait BatchQueryMerkle: Copy + Sized + 'static {
31    /// Reconstruct one digest from `DIGEST_SIZE` consecutive F-valued words in `out`
32    /// starting at index `base`.
33    fn reconstruct_from_f(out: &[F], base: usize) -> Self;
34}
35
36const MAX_MERKLE_ROWS_PER_QUERY: usize = 512;
37
38fn validate_merkle_rows_per_query(
39    rows_per_query: usize,
40    height: usize,
41) -> Result<usize, MerkleTreeError> {
42    // The generic Merkle constructor still treats "power of two" and "fits within the matrix
43    // height" as caller invariants. The CUDA-specific maximum rows-per-query is returned as a
44    // recoverable error because it depends on backend support, not generic Merkle correctness.
45    let k = log2_strict_usize(rows_per_query);
46    assert!(
47        rows_per_query <= height,
48        "rows_per_query ({rows_per_query}) must not exceed height ({height})"
49    );
50    if rows_per_query > MAX_MERKLE_ROWS_PER_QUERY {
51        return Err(MerkleTreeError::UnsupportedRowsPerQuery {
52            rows_per_query,
53            max_rows_per_query: MAX_MERKLE_ROWS_PER_QUERY,
54        });
55    }
56    Ok(k)
57}
58
59impl BatchQueryMerkle for Digest {
60    fn reconstruct_from_f(out: &[F], base: usize) -> Self {
61        from_fn(|i| out[base + i])
62    }
63}
64
65#[cfg(feature = "baby-bear-bn254-poseidon2")]
66impl BatchQueryMerkle for Bn254Digest {
67    fn reconstruct_from_f(out: &[F], base: usize) -> Self {
68        // Safety: [F; DIGEST_SIZE] and Bn254Digest have the same size (32 bytes).
69        const _: () =
70            assert!(std::mem::size_of::<Bn254Digest>() == DIGEST_SIZE * std::mem::size_of::<F>());
71        let f_arr: [F; DIGEST_SIZE] = from_fn(|i| out[base + i]);
72        unsafe { std::ptr::read_unaligned(f_arr.as_ptr() as *const Bn254Digest) }
73    }
74}
75
76pub struct MerkleTreeGpu<F, Digest> {
77    /// The matrix that is used to form the leaves of the Merkle tree, which are
78    /// in turn hashed into the bottom digest layer.
79    ///
80    /// This matrix is optionally cached depending on the prover configuration:
81    /// - Caching increases the peak GPU memory but avoids a recomputation of MLE eval-to-coeffs,
82    ///   batch_expand, and forward NTT.
83    /// - Not caching pays a performance penalty due to the above recomputation.
84    pub(crate) backing_matrix: Option<DeviceMatrix<F>>,
85    pub(crate) digest_layers: Vec<DeviceBuffer<Digest>>,
86    pub(crate) rows_per_query: usize,
87    pub(crate) root: Digest,
88}
89
90pub trait MerkleTreeConstructor: GpuMerkleHash {
91    fn new_merkle_tree(
92        matrix: DeviceMatrix<F>,
93        rows_per_query: usize,
94        cache_backing_matrix: bool,
95        device_ctx: &GpuDeviceCtx,
96    ) -> Result<MerkleTreeGpu<F, Self::Digest>, MerkleTreeError>;
97}
98
99pub trait MerkleProofQueryDigest: BatchQueryMerkle + Copy + Send + Sync + 'static {
100    fn batch_query_merkle_proofs(
101        trees: &[&MerkleTreeGpu<F, Self>],
102        query_indices: &[usize],
103        device_ctx: &GpuDeviceCtx,
104    ) -> Result<Vec<Vec<Vec<Self>>>, MerkleTreeError>;
105}
106
107impl<F, Digest> MerkleTreeGpu<F, Digest> {
108    pub fn root(&self) -> Digest
109    where
110        Digest: Clone,
111    {
112        self.root.clone()
113    }
114
115    pub fn query_stride(&self) -> usize {
116        self.digest_layers[0].len()
117    }
118
119    pub fn proof_depth(&self) -> usize {
120        self.digest_layers.len() - 1
121    }
122}
123
124// Base field merkle tree — generic constructor
125impl<D: Copy + Send + Sync + 'static> MerkleTreeGpu<F, D> {
126    /// Build a Merkle tree using the given hash scheme `MH`.
127    ///
128    /// This is the primary constructor; `new()` is a convenience wrapper that
129    /// fixes `MH = Poseidon2MerkleHash`.
130    #[instrument(name = "merkle_tree", skip_all)]
131    pub fn new_with_hash<MH: MerkleTreeConstructor<Digest = D>>(
132        matrix: DeviceMatrix<F>,
133        rows_per_query: usize,
134        cache_backing_matrix: bool,
135        device_ctx: &GpuDeviceCtx,
136    ) -> Result<Self, MerkleTreeError> {
137        MH::new_merkle_tree(matrix, rows_per_query, cache_backing_matrix, device_ctx)
138    }
139
140    fn new_generic_with_hash<MH: GpuMerkleHash<Digest = D>>(
141        matrix: DeviceMatrix<F>,
142        rows_per_query: usize,
143        cache_backing_matrix: bool,
144        device_ctx: &GpuDeviceCtx,
145    ) -> Result<Self, MerkleTreeError> {
146        let mem = MemTracker::start("prover.merkle_tree");
147        let height = matrix.height();
148        assert!(height.is_power_of_two());
149        let k = validate_merkle_rows_per_query(rows_per_query, height)?;
150        let query_stride = height / rows_per_query;
151        let mut query_digest_layer = DeviceBuffer::<D>::with_capacity_on(query_stride, device_ctx);
152        // SAFETY: query_digest_layer properly allocated
153        unsafe {
154            MH::compress_rows(
155                &mut query_digest_layer,
156                matrix.buffer(),
157                matrix.width(),
158                query_stride,
159                k,
160                device_ctx,
161            )
162            .map_err(MerkleTreeError::CompressingRowHashes)?;
163        }
164        // If not caching, drop the backing matrix at this point to save memory.
165        let backing_matrix = cache_backing_matrix.then_some(matrix);
166
167        let mut digest_layers = vec![query_digest_layer];
168        while digest_layers.last().unwrap().len() > 1 {
169            let prev_layer = digest_layers.last().unwrap();
170            let mut layer = DeviceBuffer::<D>::with_capacity_on(prev_layer.len() / 2, device_ctx);
171            let layer_len = layer.len();
172            let layer_idx = digest_layers.len();
173            // SAFETY:
174            // - `layer` is properly allocated with half the size of `prev_layer` and does not
175            //   overlap with it.
176            unsafe {
177                MH::compress_layer(&mut layer, prev_layer, layer_len, device_ctx).map_err(
178                    |error| MerkleTreeError::AdjacentCompressLayer {
179                        error,
180                        layer: layer_idx,
181                    },
182                )?;
183            }
184            digest_layers.push(layer);
185        }
186        let d_root = digest_layers.last().unwrap();
187        assert_eq!(d_root.len(), 1, "Only one root is supported");
188        let root = d_root.to_host_on(device_ctx)?.pop().unwrap();
189
190        mem.emit_metrics();
191        Ok(Self {
192            backing_matrix,
193            digest_layers,
194            rows_per_query,
195            root,
196        })
197    }
198
199    #[instrument(name = "batch_open_rows", skip_all)]
200    pub fn batch_open_rows(
201        backing_matrices: &[&DeviceMatrix<F>],
202        query_indices: &[usize],
203        query_stride: usize,
204        rows_per_query: usize,
205        device_ctx: &GpuDeviceCtx,
206    ) -> Result<
207        Vec<
208            // per tree
209            Vec<
210                // per query index
211                Vec<F>, // opened rows, concatenated for rows_per_query strided rows
212            >,
213        >,
214        MerkleTreeError,
215    > {
216        if query_indices.is_empty() {
217            return Ok(vec![Vec::new(); backing_matrices.len()]);
218        }
219        let row_idxs = query_indices
220            .iter()
221            .flat_map(|&query_idx| {
222                debug_assert!(query_idx < query_stride);
223                (0..rows_per_query)
224                    .map(move |row_offset| (row_offset * query_stride + query_idx) as u32)
225            })
226            .collect_vec();
227        let d_row_idxs = row_idxs.to_device_on(device_ctx)?;
228        // PERF[jpw]: I did not batch across trees into a single kernel call because widths are
229        // different so it was inconvenient. Make a new kernel if slow.
230        // NOTE: par_iter requires cross-stream waits, not worth the effort
231        backing_matrices
232            .iter()
233            .enumerate()
234            .map(|(matrix_idx, matrix)| {
235                let d_out = DeviceBuffer::<F>::with_capacity_on(
236                    row_idxs.len() * matrix.width(),
237                    device_ctx,
238                );
239                // SAFETY:
240                // - `output_rows` is allocated with row_idxs.len() * width
241                // - row indices are within bounds
242                unsafe {
243                    matrix_get_rows_fp_kernel(
244                        &d_out,
245                        matrix.buffer(),
246                        &d_row_idxs,
247                        matrix.width() as u64,
248                        matrix.height() as u64,
249                        d_row_idxs.len(),
250                        device_ctx.stream.as_raw(),
251                    )
252                    .map_err(|error| MerkleTreeError::MatrixGetRows { error, matrix_idx })?;
253                }
254                let width = matrix.width();
255                let out =
256                    info_span!("opened_rows_d2h").in_scope(|| d_out.to_host_on(device_ctx))?;
257                let opened_rows_per_query = out
258                    .chunks_exact(rows_per_query * width)
259                    .map(|rows| rows.to_vec())
260                    .collect_vec();
261                Ok(opened_rows_per_query)
262            })
263            .collect::<Result<Vec<_>, MerkleTreeError>>()
264    }
265}
266
267impl MerkleTreeConstructor for Poseidon2MerkleHash {
268    fn new_merkle_tree(
269        matrix: DeviceMatrix<F>,
270        rows_per_query: usize,
271        cache_backing_matrix: bool,
272        device_ctx: &GpuDeviceCtx,
273    ) -> Result<MerkleTreeGpu<F, Self::Digest>, MerkleTreeError> {
274        MerkleTreeGpu::<F, Self::Digest>::new_generic_with_hash::<Self>(
275            matrix,
276            rows_per_query,
277            cache_backing_matrix,
278            device_ctx,
279        )
280    }
281}
282
283#[cfg(feature = "baby-bear-bn254-poseidon2")]
284impl MerkleTreeConstructor for crate::hash_scheme::Bn254Poseidon2MerkleHash {
285    fn new_merkle_tree(
286        matrix: DeviceMatrix<F>,
287        rows_per_query: usize,
288        cache_backing_matrix: bool,
289        device_ctx: &GpuDeviceCtx,
290    ) -> Result<MerkleTreeGpu<F, Self::Digest>, MerkleTreeError> {
291        MerkleTreeGpu::<F, Self::Digest>::new_generic_with_hash::<Self>(
292            matrix,
293            rows_per_query,
294            cache_backing_matrix,
295            device_ctx,
296        )
297    }
298}
299
300// Base field merkle tree — Poseidon2 default constructor
301impl MerkleTreeGpu<F, Digest> {
302    /// Build a Merkle tree using the default Poseidon2 hash.
303    ///
304    /// Equivalent to `new_with_hash::<Poseidon2MerkleHash>(...)`.
305    pub fn new(
306        matrix: DeviceMatrix<F>,
307        rows_per_query: usize,
308        cache_backing_matrix: bool,
309        device_ctx: &GpuDeviceCtx,
310    ) -> Result<Self, MerkleTreeError> {
311        Self::new_with_hash::<Poseidon2MerkleHash>(
312            matrix,
313            rows_per_query,
314            cache_backing_matrix,
315            device_ctx,
316        )
317    }
318}
319
320// Base field merkle tree — generic batch query (works for any BatchQueryMerkle digest)
321impl<D: BatchQueryMerkle + Send + Sync + 'static> MerkleTreeGpu<F, D> {
322    fn batch_query_proofs(
323        trees: &[&Self],
324        query_indices: &[usize],
325        device_ctx: &GpuDeviceCtx,
326    ) -> Result<Vec<Vec<Vec<D>>>, MerkleTreeError> {
327        if trees.is_empty() {
328            return Ok(Vec::new());
329        }
330        // The kernel treats each layer as a separate array and does parallel accesses;
331        // we lay out all the layer pointers flattened into a vec.
332        let num_trees = trees.len();
333        let num_queries = query_indices.len();
334        let depth = trees[0].proof_depth();
335        debug_assert!(
336            trees.iter().all(|tree| tree.proof_depth() == depth),
337            "Merkle trees don't have same depth"
338        );
339        if num_queries == 0 {
340            return Ok(vec![Vec::new(); num_trees]);
341        }
342        if depth == 0 {
343            return Ok(vec![vec![Vec::new(); num_queries]; num_trees]);
344        }
345        let layers_ptr = trees
346            .iter()
347            .flat_map(|tree| {
348                // skip root layer [depth]
349                tree.digest_layers
350                    .iter()
351                    .take(depth)
352                    .map(|layer| layer.as_ptr() as u64)
353            })
354            .collect_vec();
355        let d_layers_ptr = layers_ptr.to_device_on(device_ctx)?;
356        debug_assert_eq!(d_layers_ptr.len(), num_trees * depth);
357
358        // [query_idx] is the top level grouping
359        let indices = query_indices
360            .iter()
361            .flat_map(|&index| {
362                (0..num_trees).flat_map(move |tree_idx| {
363                    (0..depth).map(move |layer_idx| {
364                        debug_assert!(index < trees[tree_idx].query_stride());
365                        ((index >> layer_idx) ^ 1) as u64
366                    })
367                })
368            })
369            .collect_vec();
370        let d_indices = indices.to_device_on(device_ctx)?;
371        debug_assert_eq!(d_indices.len(), d_layers_ptr.len() * num_queries);
372
373        let mut d_out = DeviceBuffer::<F>::with_capacity_on(
374            d_layers_ptr.len() * num_queries * DIGEST_SIZE,
375            device_ctx,
376        );
377        // SAFETY:
378        // - d_out has size num_trees * depth * num_queries * DIGEST_SIZE in `F` elements
379        // - d_layers_ptr is size `num_trees * depth` device pointers
380        // - d_indices is size `num_trees * depth * num_queries` indices for merkle proof sibling
381        //   indices
382        // - Both Digest=[F;8] and Bn254Digest=[Bn254Scalar;1] are 32 bytes == DIGEST_SIZE * 4, so
383        //   the same kernel correctly copies the raw bytes for either digest type.
384        unsafe {
385            query_digest_layers(
386                &mut d_out,
387                &d_layers_ptr,
388                &d_indices,
389                num_queries,
390                d_layers_ptr.len(),
391                device_ctx.stream.as_raw(),
392            )
393            .map_err(MerkleTreeError::QueryDigestLayers)?;
394        }
395        let out = d_out.to_host_on(device_ctx)?;
396        // Chunk up the array using D::reconstruct_from_f
397        let res = (0..num_trees)
398            .map(|tree_idx| {
399                (0..num_queries)
400                    .map(|query_idx| {
401                        (0..depth)
402                            .map(|layer_idx| {
403                                let base =
404                                    (query_idx * num_trees * depth + tree_idx * depth + layer_idx)
405                                        * DIGEST_SIZE;
406                                D::reconstruct_from_f(&out, base)
407                            })
408                            .collect_vec()
409                    })
410                    .collect_vec()
411            })
412            .collect_vec();
413        Ok(res)
414    }
415
416    /// Batch queries multiple `trees` at _the same_ `query_indices` for merkle proofs.
417    ///
418    /// # Assumptions
419    /// - All `trees` have the same depth.
420    pub fn batch_query_merkle_proofs(
421        trees: &[&Self],
422        query_indices: &[usize],
423        device_ctx: &GpuDeviceCtx,
424    ) -> Result<
425        Vec<
426            // per tree
427            Vec<
428                // per query index
429                Vec<D>, // merkle proof
430            >,
431        >,
432        MerkleTreeError,
433    >
434    where
435        D: MerkleProofQueryDigest,
436    {
437        D::batch_query_merkle_proofs(trees, query_indices, device_ctx)
438    }
439}
440
441impl MerkleProofQueryDigest for Digest {
442    fn batch_query_merkle_proofs(
443        trees: &[&MerkleTreeGpu<F, Self>],
444        query_indices: &[usize],
445        device_ctx: &GpuDeviceCtx,
446    ) -> Result<Vec<Vec<Vec<Self>>>, MerkleTreeError> {
447        MerkleTreeGpu::<F, Self>::batch_query_proofs(trees, query_indices, device_ctx)
448    }
449}
450
451#[cfg(feature = "baby-bear-bn254-poseidon2")]
452impl MerkleProofQueryDigest for Bn254Digest {
453    fn batch_query_merkle_proofs(
454        trees: &[&MerkleTreeGpu<F, Self>],
455        query_indices: &[usize],
456        device_ctx: &GpuDeviceCtx,
457    ) -> Result<Vec<Vec<Vec<Self>>>, MerkleTreeError> {
458        MerkleTreeGpu::<F, Self>::batch_query_proofs(trees, query_indices, device_ctx)
459    }
460}
461
462// Extension field merkle tree — generic constructor
463impl<D: Copy + Send + Sync + 'static> MerkleTreeGpu<EF, D> {
464    /// Build a Merkle tree from an extension-field matrix using hash scheme `MH`.
465    #[instrument(name = "merkle_tree_ext", skip_all)]
466    pub fn new_with_hash<MH: GpuMerkleHash<Digest = D>>(
467        matrix: DeviceMatrix<EF>,
468        rows_per_query: usize,
469        cache_backing_matrix: bool,
470        device_ctx: &GpuDeviceCtx,
471    ) -> Result<Self, MerkleTreeError> {
472        let height = matrix.height();
473        assert!(height.is_power_of_two());
474        let k = validate_merkle_rows_per_query(rows_per_query, height)?;
475        let query_stride = height / rows_per_query;
476        let mut query_digest_layer = DeviceBuffer::<D>::with_capacity_on(query_stride, device_ctx);
477        // SAFETY: query_digest_layer properly allocated
478        unsafe {
479            MH::compress_rows_ext(
480                &mut query_digest_layer,
481                matrix.buffer(),
482                matrix.width(),
483                query_stride,
484                k,
485                device_ctx,
486            )
487            .map_err(MerkleTreeError::CompressingRowHashesExt)?;
488        }
489        // If not caching, drop the backing matrix at this point to save memory.
490        let backing_matrix = cache_backing_matrix.then_some(matrix);
491
492        let mut digest_layers = vec![query_digest_layer];
493        while digest_layers.last().unwrap().len() > 1 {
494            let prev_layer = digest_layers.last().unwrap();
495            let mut layer = DeviceBuffer::<D>::with_capacity_on(prev_layer.len() / 2, device_ctx);
496            let layer_len = layer.len();
497            let layer_idx = digest_layers.len();
498            // SAFETY:
499            // - `layer` is properly allocated with half the size of `prev_layer` and does not
500            //   overlap with it.
501            unsafe {
502                MH::compress_layer(&mut layer, prev_layer, layer_len, device_ctx).map_err(
503                    |error| MerkleTreeError::AdjacentCompressLayer {
504                        error,
505                        layer: layer_idx,
506                    },
507                )?;
508            }
509            digest_layers.push(layer);
510        }
511        let d_root = digest_layers.last().unwrap();
512        assert_eq!(d_root.len(), 1, "Only one root is supported");
513        let root = d_root.to_host_on(device_ctx)?.pop().unwrap();
514
515        Ok(Self {
516            backing_matrix,
517            digest_layers,
518            rows_per_query,
519            root,
520        })
521    }
522}
523
524// Extension field merkle tree — Poseidon2 default constructor
525impl MerkleTreeGpu<EF, Digest> {
526    /// Build a Merkle tree from an extension-field matrix using Poseidon2.
527    ///
528    /// NOTE: currently unused because we transpose `DeviceMatrix<EF>` to
529    /// `DeviceMatrix<F>` beforehand in our use cases.
530    pub fn new(
531        matrix: DeviceMatrix<EF>,
532        rows_per_query: usize,
533        cache_backing_matrix: bool,
534        device_ctx: &GpuDeviceCtx,
535    ) -> Result<Self, MerkleTreeError> {
536        Self::new_with_hash::<Poseidon2MerkleHash>(
537            matrix,
538            rows_per_query,
539            cache_backing_matrix,
540            device_ctx,
541        )
542    }
543}