1use std::{ffi::c_void, sync::Arc};
2
3use getset::Getters;
4use itertools::Itertools;
5use openvm_cuda_common::{
6 copy::cuda_memcpy_on, d_buffer::DeviceBuffer, memory_manager::MemTracker, stream::GpuDeviceCtx,
7};
8use openvm_stark_backend::{
9 p3_util::log2_strict_usize,
10 prover::{stacked_pcs::StackedLayout, MatrixDimensions},
11};
12use tracing::instrument;
13
14use crate::{
15 base::{DeviceMatrix, DeviceMatrixView},
16 cuda::{
17 batch_ntt_small::batch_ntt_small,
18 matrix::{batch_expand_pad, batch_expand_pad_wide},
19 ntt::bit_rev,
20 },
21 hash_scheme::GpuMerkleHash,
22 merkle_tree::{MerkleTreeConstructor, MerkleTreeGpu},
23 ntt::batch_ntt,
24 poly::{mle_interpolate_stages, PleMatrix},
25 prelude::F,
26 GpuProverConfig, ProverError, RsCodeMatrixError, StackTracesError,
27};
28
29#[derive(Getters)]
30pub struct StackedPcsDataGpu<F, Digest> {
31 #[getset(get = "pub")]
33 pub(crate) layout: StackedLayout,
34 #[getset(get = "pub")]
41 pub(crate) matrix: Option<PleMatrix<F>>,
42 #[getset(get = "pub")]
45 pub(crate) tree: MerkleTreeGpu<F, Digest>,
46}
47
48#[allow(clippy::type_complexity)]
49#[instrument(level = "info", skip_all)]
50pub fn stacked_commit<MH: GpuMerkleHash + MerkleTreeConstructor>(
51 l_skip: usize,
52 n_stack: usize,
53 log_blowup: usize,
54 k_whir: usize,
55 traces: &[&DeviceMatrix<F>],
56 prover_config: GpuProverConfig,
57 device_ctx: &GpuDeviceCtx,
58) -> Result<(MH::Digest, StackedPcsDataGpu<F, MH::Digest>), ProverError> {
59 let mut mem = MemTracker::start("prover.stacked_commit");
60 mem.tracing_info("before stacked_commit");
61 mem.reset_peak();
62 let layout = get_stacked_layout(l_skip, n_stack, traces);
63 tracing::info!(
64 height = layout.height(),
65 width = layout.width(),
66 "stacked_matrix_dimensions"
67 );
68 let opt_stacked_matrix = if prover_config.cache_stacked_matrix {
69 Some(stack_traces(&layout, traces, device_ctx)?)
70 } else {
71 None
72 };
73 let rs_matrix = rs_code_matrix(log_blowup, &layout, traces, &opt_stacked_matrix, device_ctx)?;
74 let tree = MerkleTreeGpu::<F, MH::Digest>::new_with_hash::<MH>(
75 rs_matrix,
76 1 << k_whir,
77 prover_config.cache_rs_code_matrix,
78 device_ctx,
79 )?;
80 let root = tree.root();
81 let data = StackedPcsDataGpu {
82 layout,
83 matrix: opt_stacked_matrix,
84 tree,
85 };
86 mem.emit_metrics();
87 Ok((root, data))
88}
89
90#[instrument(skip_all)]
94pub fn stacked_matrix(
95 l_skip: usize,
96 n_stack: usize,
97 traces: &[&DeviceMatrix<F>],
98 device_ctx: &GpuDeviceCtx,
99) -> Result<(PleMatrix<F>, StackedLayout), ProverError> {
100 let layout = get_stacked_layout(l_skip, n_stack, traces);
101 let matrix = stack_traces(&layout, traces, device_ctx)?;
102 Ok((matrix, layout))
103}
104
105pub(crate) fn get_stacked_layout(
106 l_skip: usize,
107 n_stack: usize,
108 traces: &[&DeviceMatrix<F>],
109) -> StackedLayout {
110 let sorted_meta = traces
111 .iter()
112 .map(|trace| {
113 let log_height = log2_strict_usize(trace.height());
115 (trace.width(), log_height)
116 })
117 .collect_vec();
118 debug_assert!(sorted_meta.is_sorted_by(|a, b| a.1 >= b.1));
119 StackedLayout::new(l_skip, l_skip + n_stack, sorted_meta).unwrap()
120}
121
122pub(crate) fn stack_traces(
123 layout: &StackedLayout,
124 traces: &[&DeviceMatrix<F>],
125 device_ctx: &GpuDeviceCtx,
126) -> Result<PleMatrix<F>, StackTracesError> {
127 let mem = MemTracker::start("prover.stack_traces");
128 let l_skip = layout.l_skip();
129 let height = layout.height();
130 let width = layout.width();
131 let mut q_evals =
132 DeviceBuffer::<F>::with_capacity_on(width.checked_mul(height).unwrap(), device_ctx);
133 stack_traces_into_expanded(layout, traces, &mut q_evals, height, device_ctx)?;
134 mem.emit_metrics();
135 Ok(PleMatrix::from_evals(
136 l_skip, q_evals, height, width, device_ctx,
137 ))
138}
139
140pub(crate) fn stack_traces_into_expanded(
144 layout: &StackedLayout,
145 traces: &[&DeviceMatrix<F>],
146 buffer: &mut DeviceBuffer<F>,
147 padded_height: usize,
148 device_ctx: &GpuDeviceCtx,
149) -> Result<(), StackTracesError> {
150 let l_skip = layout.l_skip();
151 debug_assert_eq!(padded_height % layout.height(), 0);
152 debug_assert_eq!(buffer.len() % padded_height, 0);
153 debug_assert_eq!(buffer.len() / padded_height, layout.width());
154 buffer
155 .fill_zero_on(device_ctx)
156 .map_err(StackTracesError::FillZero)?;
157 let mut idx = 0;
158 while idx < layout.sorted_cols.len() {
159 let (mat_idx, j, s) = &layout.sorted_cols[idx];
160 let start = s.col_idx * padded_height + s.row_idx;
161 let trace = traces[*mat_idx];
162 let s_len = s.len(l_skip);
163 debug_assert_eq!(trace.height(), 1 << s.log_height());
164 if s.log_height() >= l_skip {
165 debug_assert_eq!(trace.height(), s_len);
166 let mut copy_len = s_len;
167 let mut end = idx + 1;
168 while end < layout.sorted_cols.len() {
169 let (next_mat_idx, next_j, next_s) = &layout.sorted_cols[end];
170 if *next_mat_idx != *mat_idx || next_s.log_height() != s.log_height() {
171 break;
172 }
173 let expected_j = *j + (end - idx);
174 let next_len = next_s.len(l_skip);
175 let next_start = next_s.col_idx * padded_height + next_s.row_idx;
176 if *next_j != expected_j || next_len != s_len || next_start != start + copy_len {
177 break;
178 }
179 copy_len += next_len;
180 end += 1;
181 }
182
183 unsafe {
186 let src = trace.buffer().as_ptr().add(*j * s_len);
187 let dst = buffer.as_mut_ptr().add(start);
188 cuda_memcpy_on::<true, true>(
189 dst as *mut c_void,
190 src as *const c_void,
191 copy_len * size_of::<F>(),
192 device_ctx,
193 )?;
194 }
195 idx = end;
196 } else {
197 let stride = s.stride(l_skip);
198 debug_assert_eq!(stride * trace.height(), s_len);
199 unsafe {
204 let src = trace.buffer().as_ptr().add(*j * trace.height());
205 let dst = buffer.as_mut_ptr().add(start);
206 batch_expand_pad_wide(
207 dst,
208 src,
209 trace.height() as u32,
210 stride as u32,
211 1,
212 device_ctx.stream.as_raw(),
213 )
214 .map_err(StackTracesError::BatchExpandPadWide)?;
215 }
216 idx += 1;
217 }
218 }
219 Ok(())
220}
221
222#[instrument(skip_all)]
229pub fn rs_code_matrix(
230 log_blowup: usize,
231 layout: &StackedLayout,
232 traces: &[&DeviceMatrix<F>],
233 stacked_matrix: &Option<PleMatrix<F>>,
234 device_ctx: &GpuDeviceCtx,
235) -> Result<DeviceMatrix<F>, RsCodeMatrixError> {
236 let mem = MemTracker::start_and_reset_peak("prover.rs_code_matrix");
237 let l_skip = layout.l_skip();
238 let height = layout.height();
239 let width = layout.width();
240 debug_assert!(height >= (1 << l_skip));
241 let codeword_height = height.checked_shl(log_blowup as u32).unwrap();
242 let mut codewords = DeviceBuffer::<F>::with_capacity_on(codeword_height * width, device_ctx);
243 if let Some(stacked_matrix) = stacked_matrix.as_ref() {
246 unsafe {
249 batch_expand_pad(
250 codewords.as_mut_ptr(),
251 stacked_matrix.mixed.as_ptr(),
252 width as u32,
253 codeword_height as u32,
254 height as u32,
255 device_ctx.stream.as_raw(),
256 )
257 .map_err(RsCodeMatrixError::BatchExpandPad)?;
258 }
259 } else {
260 stack_traces_into_expanded(layout, traces, &mut codewords, codeword_height, device_ctx)
261 .map_err(RsCodeMatrixError::StackTraces)?;
262 if l_skip > 0 {
267 let num_uni_poly = width * (codeword_height >> l_skip);
270 unsafe {
271 batch_ntt_small(
272 &mut codewords,
273 l_skip,
274 num_uni_poly,
275 true,
276 device_ctx.stream.as_raw(),
277 )
278 .map_err(RsCodeMatrixError::CustomBatchIntt)?;
279 }
280 }
281 }
282 let log_codeword_height = log2_strict_usize(codeword_height);
287
288 if l_skip > 0 {
292 unsafe {
295 mle_interpolate_stages(
296 codewords.as_mut_ptr(),
297 width,
298 codeword_height as u32,
299 log_blowup as u32,
300 0, l_skip as u32 - 1, false, false, device_ctx.stream.as_raw(),
305 )
306 .map_err(|error| RsCodeMatrixError::MleInterpolateStage2d { error, step: 1 })?;
307 }
308 }
309
310 unsafe {
312 bit_rev(
313 &codewords,
314 &codewords,
315 log_codeword_height as u32,
316 codeword_height as u32,
317 width as u32,
318 device_ctx.stream.as_raw(),
319 )
320 .map_err(RsCodeMatrixError::BitRev)?;
321 }
322
323 batch_ntt(
325 &codewords,
326 log_codeword_height as u32,
327 0u32,
328 width as u32,
329 false, false,
331 device_ctx,
332 );
333 let code_matrix = DeviceMatrix::new(Arc::new(codewords), codeword_height, width);
334 mem.emit_metrics();
335
336 Ok(code_matrix)
337}
338
339impl<F, Digest> StackedPcsDataGpu<F, Digest> {
340 pub fn mixed_view<'a>(
346 &'a self,
347 mat_idx: usize,
348 width: usize,
349 ) -> Option<DeviceMatrixView<'a, F>> {
350 if let Some(matrix) = self.matrix.as_ref() {
351 debug_assert_eq!(self.layout.width_of(mat_idx), width);
352 let s = self
353 .layout
354 .get(mat_idx, 0)
355 .unwrap_or_else(|| panic!("Invalid matrix index: {mat_idx}"));
356 let l_skip = self.layout.l_skip();
357 let lifted_height = s.len(l_skip);
358 let offset = s.col_idx * matrix.height() + s.row_idx;
359 unsafe {
363 let ptr = matrix.mixed.as_ptr().add(offset);
364 Some(DeviceMatrixView::from_raw_parts(ptr, lifted_height, width))
365 }
366 } else {
367 None
368 }
369 }
370}
371
372#[cfg(test)]
373mod tests {
374 use itertools::Itertools;
375 use openvm_cuda_common::{
376 common::get_device,
377 stream::{CudaStream, GpuDeviceCtx, StreamGuard},
378 };
379 use openvm_stark_backend::{
380 prover::ColMajorMatrix,
381 test_utils::{InteractionsFixture11, TestFixture},
382 };
383 use p3_field::PrimeCharacteristicRing;
384
385 use super::*;
386 use crate::{
387 data_transporter::{transport_matrix_d2h_col_major, transport_matrix_h2d_col_major},
388 prelude::{F, SC},
389 };
390
391 fn test_ctx() -> GpuDeviceCtx {
392 GpuDeviceCtx {
393 device_id: get_device().unwrap() as u32,
394 stream: StreamGuard::new(CudaStream::new_non_blocking().unwrap()),
395 }
396 }
397
398 #[test]
399 fn test_stacked_matrix_manual_0() {
400 let device_ctx = test_ctx();
401 let columns = [vec![1, 2, 3, 4], vec![5, 6], vec![7]]
402 .map(|v| v.into_iter().map(F::from_u32).collect_vec());
403 let mats = columns
404 .into_iter()
405 .map(|c| {
406 transport_matrix_h2d_col_major(&ColMajorMatrix::new(c, 1), &device_ctx).unwrap()
407 })
408 .collect_vec();
409 let mat_refs = mats.iter().collect_vec();
410 let l_skip = 0;
411 let (stacked_mat, _layout) = stacked_matrix(0, 2, &mat_refs, &device_ctx).unwrap();
412 assert_eq!(stacked_mat.height(), 4);
413 assert_eq!(stacked_mat.width(), 2);
414 let stacked_h_mat = transport_matrix_d2h_col_major(
415 &stacked_mat.to_evals(l_skip, &device_ctx).unwrap(),
416 &device_ctx,
417 )
418 .unwrap();
419 assert_eq!(
420 stacked_h_mat.values,
421 [1, 2, 3, 4, 5, 6, 7, 0].map(F::from_u32).to_vec()
422 );
423 }
424
425 #[test]
426 fn test_stacked_matrix_manual_1() {
427 let gpu_ctx = test_ctx();
428 let proving_ctx = TestFixture::<SC>::generate_proving_ctx(&InteractionsFixture11);
429 let [send_trace, rcv_trace] = [0, 1].map(|i| {
430 transport_matrix_h2d_col_major(&proving_ctx.per_trace[i].1.common_main, &gpu_ctx)
431 .unwrap()
432 });
433 let l_skip = 2;
434 let n_stack = 8;
435 let (stacked_mat, _layout) =
436 stacked_matrix(l_skip, n_stack, &[&rcv_trace, &send_trace], &gpu_ctx).unwrap();
437 assert_eq!(stacked_mat.height(), 1 << (l_skip + n_stack));
438 assert_eq!(stacked_mat.width(), 1);
439 let stacked_h_mat = transport_matrix_d2h_col_major(
440 &stacked_mat.to_evals(l_skip, &gpu_ctx).unwrap(),
441 &gpu_ctx,
442 )
443 .unwrap();
444 let mut expected = vec![F::ZERO; 1 << (l_skip + n_stack)];
445 expected[..24].copy_from_slice(
446 &[
447 1, 3, 4, 2, 0, 545, 1, 0, 5, 4, 4, 5, 123, 889, 889, 456, 0, 3, 7, 546, 1, 5, 4,
448 889,
449 ]
450 .map(F::from_u32),
451 );
452 assert_eq!(stacked_h_mat.values, expected);
453 }
454
455 #[test]
456 fn test_stacked_matrix_manual_strided_0() {
457 let device_ctx = test_ctx();
458 let columns = [vec![1, 2, 3, 4], vec![5, 6], vec![7]]
459 .map(|v| v.into_iter().map(F::from_u32).collect_vec());
460 let mats = columns
461 .into_iter()
462 .map(|c| {
463 transport_matrix_h2d_col_major(&ColMajorMatrix::new(c, 1), &device_ctx).unwrap()
464 })
465 .collect_vec();
466 let mat_refs = mats.iter().collect_vec();
467 let l_skip = 2;
468 let (stacked_mat, _layout) = stacked_matrix(l_skip, 0, &mat_refs, &device_ctx).unwrap();
469 assert_eq!(stacked_mat.height(), 4);
470 assert_eq!(stacked_mat.width(), 3);
471 let stacked_h_mat = transport_matrix_d2h_col_major(
472 &stacked_mat.to_evals(l_skip, &device_ctx).unwrap(),
473 &device_ctx,
474 )
475 .unwrap();
476 assert_eq!(
477 stacked_h_mat.values,
478 [1, 2, 3, 4, 5, 0, 6, 0, 7, 0, 0, 0]
479 .map(F::from_u32)
480 .to_vec()
481 );
482 }
483
484 #[test]
485 fn test_stacked_matrix_manual_strided_1() {
486 let device_ctx = test_ctx();
487 let columns = [vec![1, 2, 3, 4], vec![5, 6], vec![7]]
488 .map(|v| v.into_iter().map(F::from_u32).collect_vec());
489 let mats = columns
490 .into_iter()
491 .map(|c| {
492 transport_matrix_h2d_col_major(&ColMajorMatrix::new(c, 1), &device_ctx).unwrap()
493 })
494 .collect_vec();
495 let mat_refs = mats.iter().collect_vec();
496 let l_skip = 3;
497 let (stacked_mat, _layout) = stacked_matrix(l_skip, 0, &mat_refs, &device_ctx).unwrap();
498 assert_eq!(stacked_mat.height(), 8);
499 assert_eq!(stacked_mat.width(), 3);
500 let stacked_h_mat = transport_matrix_d2h_col_major(
501 &stacked_mat.to_evals(l_skip, &device_ctx).unwrap(),
502 &device_ctx,
503 )
504 .unwrap();
505 assert_eq!(
506 stacked_h_mat.values,
507 [
508 [1, 0, 2, 0, 3, 0, 4, 0],
509 [5, 0, 0, 0, 6, 0, 0, 0],
510 [7, 0, 0, 0, 0, 0, 0, 0]
511 ]
512 .into_iter()
513 .flatten()
514 .map(F::from_u32)
515 .collect_vec()
516 );
517 }
518}