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
24pub trait BatchQueryMerkle: Copy + Sized + 'static {
31 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 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 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 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
124impl<D: Copy + Send + Sync + 'static> MerkleTreeGpu<F, D> {
126 #[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 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 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 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 Vec<
210 Vec<F>, >,
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 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 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
300impl MerkleTreeGpu<F, Digest> {
302 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
320impl<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 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 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 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 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 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 pub fn batch_query_merkle_proofs(
421 trees: &[&Self],
422 query_indices: &[usize],
423 device_ctx: &GpuDeviceCtx,
424 ) -> Result<
425 Vec<
426 Vec<
428 Vec<D>, >,
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
462impl<D: Copy + Send + Sync + 'static> MerkleTreeGpu<EF, D> {
464 #[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 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 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 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
524impl MerkleTreeGpu<EF, Digest> {
526 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}