1use std::{array::from_fn, cmp::max, ffi::c_void, iter::zip, mem, sync::Arc};
2
3use itertools::{zip_eq, Itertools};
4use openvm_cuda_common::{
5 copy::{cuda_memcpy_on, MemCopyD2H, MemCopyH2D},
6 d_buffer::DeviceBuffer,
7 error::MemCopyError,
8 memory_manager::MemTracker,
9 stream::GpuDeviceCtx,
10};
11use openvm_stark_backend::{
12 dft::Radix2BowersSerial,
13 p3_matrix::dense::RowMajorMatrix,
14 poly_common::{
15 eq_uni_poly, eval_eq_mle, eval_eq_uni, eval_eq_uni_at_one, eval_in_uni, Squarable,
16 UnivariatePoly,
17 },
18 proof::StackingProof,
19 prover::{
20 stacked_pcs::StackedLayout, sumcheck::sumcheck_round0_deg, DeviceMultiStarkProvingKey,
21 MatrixDimensions, ProvingContext,
22 },
23};
24use p3_dft::TwoAdicSubgroupDft;
25use p3_field::{PrimeCharacteristicRing, TwoAdicField};
26use tracing::{debug, info_span, instrument};
27
28use crate::{
29 base::DeviceMatrix,
30 cuda::{
31 batch_ntt_small::ensure_device_ntt_twiddles_initialized,
32 poly::vector_scalar_multiply_ext,
33 stacked_reduction::{
34 _stacked_reduction_r0_required_temp_buffer_size, initialize_k_rot_from_eq_segments,
35 stacked_reduction_fold_ple, stacked_reduction_sumcheck_mle_round,
36 stacked_reduction_sumcheck_mle_round_degenerate, stacked_reduction_sumcheck_round0,
37 NUM_G,
38 },
39 sumcheck::{fold_mle, triangular_fold_mle},
40 },
41 gpu_backend::GenericGpuBackend,
42 hash_scheme::GpuHashScheme,
43 poly::EqEvalSegments,
44 prelude::{Digest, D_EF, EF, F},
45 sponge::GpuFiatShamirTranscript,
46 stacked_pcs::StackedPcsDataGpu,
47 utils::{compute_barycentric_inv_lagrange_denoms, reduce_raw_u64_to_ef},
48 GpuDevice, StackedReductionError,
49};
50
51pub const STACKED_REDUCTION_S_DEG: usize = 2;
53
54pub struct StackedReductionGpu<D = Digest> {
55 device_ctx: GpuDeviceCtx,
56 sm_count: u32,
57
58 l_skip: usize,
59 n_stack: usize,
60
61 omega_skip: F,
62 omega_skip_pows: Vec<F>,
63 d_omega_skip_pows: DeviceBuffer<F>,
64
65 r_0: EF,
66 d_lambda_pows: DeviceBuffer<EF>,
67 eq_const: EF,
68
69 pub(crate) stacked_per_commit: Vec<StackedPcsData2<D>>,
70 d_q_widths: DeviceBuffer<u32>,
71 q_width_max: u32,
72 d_q_eval_ptrs: DeviceBuffer<*const EF>,
73
74 trace_ptrs: Vec<(
75 *const F, usize, usize, )>,
79 unstacked_cols: Vec<UnstackedSlice>,
80 d_unstacked_cols: DeviceBuffer<UnstackedSlice>,
81 ht_diff_idxs: Vec<usize>,
85 n_max: usize,
86
87 eq_r_ns: EqEvalSegments<EF>,
90
91 q_evals: Vec<DeviceBuffer<EF>>, eq_stable: Vec<EF>,
97 k_rot_stable: Vec<EF>,
98
99 k_rot_ns: EqEvalSegments<EF>,
103 eq_ub_per_trace: Vec<EF>,
105 d_eq_ub: DeviceBuffer<EF>,
106
107 d_block_sums: DeviceBuffer<EF>,
108 d_accum: DeviceBuffer<u64>,
109 d_input_ptrs: DeviceBuffer<*const EF>,
110 d_output_ptrs: DeviceBuffer<*mut EF>,
111
112 mem: MemTracker,
113}
114
115pub struct StackedPcsData2<D = Digest> {
123 pub(crate) inner: Arc<StackedPcsDataGpu<F, D>>,
124 pub(crate) traces: Vec<DeviceMatrix<F>>,
126}
127
128impl<D> StackedPcsData2<D> {
129 pub unsafe fn from_raw(
132 pcs_data: Arc<StackedPcsDataGpu<F, D>>,
133 traces: Vec<DeviceMatrix<F>>,
134 ) -> Self {
135 Self {
136 inner: pcs_data,
137 traces,
138 }
139 }
140
141 pub fn layout(&self) -> &StackedLayout {
142 &self.inner.layout
143 }
144}
145
146#[repr(C)]
154#[derive(Clone, Copy, Debug)]
155pub(crate) struct UnstackedSlice {
156 commit_idx: u32,
157 log_height: u32,
158 stacked_row_idx: u32,
159 stacked_col_idx: u32,
160}
161
162impl<D> StackedReductionGpu<D> {
163 fn log_stacked_height(&self, round: usize) -> usize {
164 self.n_stack - (round - 1)
165 }
166
167 fn stacked_height(&self, round: usize) -> usize {
168 1 << self.log_stacked_height(round)
169 }
170
171 fn cur_max_n(&self, round: usize) -> usize {
173 self.n_max - (round - 1)
174 }
175}
176
177#[allow(clippy::type_complexity)]
182#[instrument(
183 name = "prover.openings.stacked_reduction",
184 level = "info",
185 skip_all,
186 fields(phase = "prover")
187)]
188pub fn prove_stacked_opening_reduction_gpu<HS, TS>(
189 device: &GpuDevice,
190 transcript: &mut TS,
191 mpk: &DeviceMultiStarkProvingKey<GenericGpuBackend<HS>>,
192 ctx: ProvingContext<GenericGpuBackend<HS>>,
193 common_main_pcs_data: StackedPcsDataGpu<F, HS::Digest>,
194 r: &[EF],
195) -> Result<
196 (
197 StackingProof<HS::SC>,
198 Vec<EF>,
199 Vec<StackedPcsData2<HS::Digest>>,
200 ),
201 StackedReductionError,
202>
203where
204 HS: GpuHashScheme,
205 TS: GpuFiatShamirTranscript<HS::SC>,
206{
207 let n_stack = device.params().n_stack;
208 let lambda = transcript.sample_ext();
210
211 let _round0_span =
212 info_span!("prover.openings.stacked_reduction.round0", phase = "prover").entered();
213 let mut prover = StackedReductionGpu::new::<HS>(
214 mpk,
215 ctx,
216 common_main_pcs_data,
217 r,
218 lambda,
219 device.sm_count(),
220 device.device_ctx.clone(),
221 )?;
222
223 let s_0 = prover.batch_sumcheck_uni_round0_poly()?;
225 for &coeff in s_0.coeffs() {
226 transcript.observe_ext(coeff);
227 }
228
229 let mut u_vec = Vec::with_capacity(n_stack + 1);
230 let u_0 = transcript.sample_ext();
231 u_vec.push(u_0);
232 debug!(round = 0, u_round = %u_0);
233
234 prover.fold_ple_evals(u_0)?;
235 drop(_round0_span);
236 let mut sumcheck_round_polys = Vec::with_capacity(n_stack);
239
240 let _mle_rounds_span = info_span!(
242 "prover.openings.stacked_reduction.mle_rounds",
243 phase = "prover"
244 )
245 .entered();
246 #[allow(clippy::needless_range_loop)]
247 for round in 1..=n_stack {
248 let batch_s_evals = prover.batch_sumcheck_poly_eval(round, u_vec[round - 1])?;
249
250 for &eval in &batch_s_evals {
251 transcript.observe_ext(eval);
252 }
253 sumcheck_round_polys.push(batch_s_evals);
254
255 let u_round = transcript.sample_ext();
256 u_vec.push(u_round);
257 debug!(%round, %u_round);
258
259 prover.fold_mle_evals(round, u_round)?;
260 }
261 let stacking_openings = prover.get_stacked_openings()?;
262 for claims_for_com in &stacking_openings {
263 for &claim in claims_for_com {
264 transcript.observe_ext(claim);
265 }
266 }
267 drop(_mle_rounds_span);
268 let proof = StackingProof {
269 univariate_round_coeffs: s_0.into_coeffs(),
270 sumcheck_round_polys,
271 stacking_openings,
272 };
273 Ok((proof, u_vec, prover.stacked_per_commit))
274}
275
276impl<D: Copy + Clone + Send + Sync + 'static> StackedReductionGpu<D> {
277 #[instrument("stacked_reduction_new", level = "debug", skip_all)]
278 fn new<HS: GpuHashScheme<Digest = D>>(
279 mpk: &DeviceMultiStarkProvingKey<GenericGpuBackend<HS>>,
280 proving_ctx: ProvingContext<GenericGpuBackend<HS>>,
281 common_main_pcs_data: StackedPcsDataGpu<F, D>,
282 r: &[EF],
283 lambda: EF,
284 sm_count: u32,
285 device_ctx: GpuDeviceCtx,
286 ) -> Result<Self, StackedReductionError> {
287 ensure_device_ntt_twiddles_initialized().map_err(StackedReductionError::InitNttTwiddles)?;
288 let mem = MemTracker::start("prover.stacked_reduction_new");
289 let l_skip = mpk.params.l_skip;
290 let n_stack = mpk.params.n_stack;
291
292 let omega_skip = F::two_adic_generator(l_skip);
293 let omega_skip_pows = omega_skip.powers().take(1 << l_skip).collect_vec();
294 let d_omega_skip_pows = omega_skip_pows.to_device_on(&device_ctx)?;
295
296 let common_main_traces = proving_ctx
301 .per_trace
302 .iter()
303 .map(|(_, air_ctx)| air_ctx.common_main.clone())
304 .collect_vec();
305 let common_main_stacked = unsafe {
307 StackedPcsData2::from_raw(Arc::new(common_main_pcs_data), common_main_traces)
308 };
309 let mut stacked_per_commit = vec![common_main_stacked];
310 for (air_idx, air_ctx) in proving_ctx.per_trace.iter() {
311 for committed in mpk.per_air[*air_idx]
312 .preprocessed_data
313 .iter()
314 .chain(air_ctx.cached_mains.iter())
315 {
316 let stacked = unsafe {
318 StackedPcsData2::from_raw(committed.data.clone(), vec![committed.trace.clone()])
319 };
320 stacked_per_commit.push(stacked);
321 }
322 }
323
324 debug_assert!(stacked_per_commit
325 .iter()
326 .all(|d| d.layout().height() == 1 << (l_skip + n_stack)));
327
328 let need_rot_per_trace = proving_ctx
329 .per_trace
330 .iter()
331 .map(|(air_idx, _)| mpk.per_air[*air_idx].vk.params.need_rot)
332 .collect_vec();
333 let mut need_rot_per_commit = vec![need_rot_per_trace];
334 for (air_idx, air_ctx) in proving_ctx.per_trace.iter() {
335 let need_rot = mpk.per_air[*air_idx].vk.params.need_rot;
336 if mpk.per_air[*air_idx].preprocessed_data.is_some() {
337 need_rot_per_commit.push(vec![need_rot]);
338 }
339 for _ in &air_ctx.cached_mains {
340 need_rot_per_commit.push(vec![need_rot]);
341 }
342 }
343 let q_widths = stacked_per_commit
344 .iter()
345 .map(|d| d.layout().width() as u32)
346 .collect_vec();
347 let q_width_max = *q_widths.iter().max().unwrap();
348 let d_q_widths = q_widths.to_device_on(&device_ctx)?;
349
350 let total_num_cols: usize = stacked_per_commit
351 .iter()
352 .map(|d| d.layout().sorted_cols.len())
353 .sum();
354 let mut unstacked_cols = Vec::with_capacity(total_num_cols);
355 let mut need_rot_per_col = Vec::with_capacity(total_num_cols);
356 let mut ht_diff_idxs = Vec::new();
357 let mut trace_ptrs = Vec::new();
358 for (commit_idx, stacked) in stacked_per_commit.iter().enumerate() {
359 let layout = stacked.layout();
360 let need_rot_for_commit = &need_rot_per_commit[commit_idx];
361 debug_assert_eq!(need_rot_for_commit.len(), layout.mat_starts.len());
362 for (mat_idx, (trace, &idx)) in zip_eq(&stacked.traces, &layout.mat_starts).enumerate()
363 {
364 debug_assert_ne!(trace.width(), 0);
365 debug_assert_ne!(trace.height(), 0);
366 ht_diff_idxs.push(unstacked_cols.len());
367 trace_ptrs.push((trace.buffer().as_ptr(), trace.height(), trace.width()));
368 let need_rot = need_rot_for_commit[mat_idx];
369 for j in 0..trace.width() {
370 let (_, _j, s) = layout.sorted_cols[idx + j];
371 debug_assert_eq!(_j, j);
372 debug_assert_eq!(1 << s.log_height(), trace.height());
373 unstacked_cols.push(UnstackedSlice {
374 commit_idx: commit_idx as u32,
375 log_height: s.log_height() as u32,
376 stacked_row_idx: s.row_idx as u32,
377 stacked_col_idx: s.col_idx as u32,
378 });
379 need_rot_per_col.push(need_rot);
380 }
381 }
382 }
383 debug_assert_eq!(unstacked_cols.len(), total_num_cols);
384 ht_diff_idxs.push(unstacked_cols.len());
385
386 let lambda_pows_used = lambda.powers().take(total_num_cols * 2).collect_vec();
387 let mut lambda_pows = vec![EF::ZERO; total_num_cols * 2];
388 for (col_idx, need_rot) in need_rot_per_col.into_iter().enumerate() {
389 let lambda_eq_idx = 2 * col_idx;
390 let lambda_rot_idx = 2 * col_idx + 1;
391 lambda_pows[lambda_eq_idx] = lambda_pows_used[lambda_eq_idx];
392 if need_rot {
393 lambda_pows[lambda_rot_idx] = lambda_pows_used[lambda_rot_idx];
394 }
395 }
396 let d_lambda_pows = lambda_pows.to_device_on(&device_ctx)?;
397
398 let d_unstacked_cols = unstacked_cols.to_device_on(&device_ctx)?;
399 let num_windows = ht_diff_idxs.len().saturating_sub(1).max(1);
400
401 let n_max = r.len() - 1;
403 debug_assert_eq!(
404 n_max,
405 stacked_per_commit
406 .iter()
407 .map(|d| d.layout().sorted_cols[0].2.log_height())
408 .max()
409 .unwrap_or(0)
410 .saturating_sub(l_skip)
411 );
412 let eq_r_ns = EqEvalSegments::new(&r[1..], &device_ctx)
413 .map_err(StackedReductionError::EqEvalSegments)?;
414
415 let eq_const = eval_eq_uni_at_one(l_skip, r[0] * omega_skip);
416 let eq_ub_per_trace = vec![EF::ONE; unstacked_cols.len()];
417 let d_q_eval_ptrs = if stacked_per_commit.is_empty() {
418 DeviceBuffer::new()
419 } else {
420 DeviceBuffer::with_capacity_on(stacked_per_commit.len(), &device_ctx)
421 };
422 let d_input_ptrs = if stacked_per_commit.is_empty() {
423 DeviceBuffer::new()
424 } else {
425 DeviceBuffer::with_capacity_on(stacked_per_commit.len(), &device_ctx)
426 };
427 let d_output_ptrs = if stacked_per_commit.is_empty() {
428 DeviceBuffer::new()
429 } else {
430 DeviceBuffer::with_capacity_on(stacked_per_commit.len(), &device_ctx)
431 };
432 let d_accum = DeviceBuffer::<u64>::with_capacity_on(
433 num_windows * STACKED_REDUCTION_S_DEG * D_EF,
434 &device_ctx,
435 );
436 let d_eq_ub = if unstacked_cols.is_empty() {
437 DeviceBuffer::new()
438 } else {
439 DeviceBuffer::with_capacity_on(unstacked_cols.len(), &device_ctx)
440 };
441
442 Ok(Self {
443 device_ctx,
444 sm_count,
445 l_skip,
446 n_stack,
447 omega_skip,
448 omega_skip_pows,
449 d_omega_skip_pows,
450 r_0: r[0],
451 d_lambda_pows,
452 eq_const,
453 stacked_per_commit,
454 d_q_widths,
455 q_width_max,
456 d_q_eval_ptrs,
457 trace_ptrs,
458 unstacked_cols,
459 d_unstacked_cols,
460 ht_diff_idxs,
461 n_max,
462 eq_r_ns,
463 q_evals: vec![],
464 eq_stable: vec![],
465 k_rot_stable: vec![],
466 k_rot_ns: unsafe { EqEvalSegments::from_raw_parts(DeviceBuffer::new(), 0) },
468 eq_ub_per_trace,
469 d_eq_ub,
470 d_block_sums: DeviceBuffer::new(),
471 d_accum,
472 d_input_ptrs,
473 d_output_ptrs,
474 mem,
475 })
476 }
477
478 #[instrument(
484 "stacked_reduction_sumcheck",
485 level = "debug",
486 skip_all,
487 fields(round = 0)
488 )]
489 fn batch_sumcheck_uni_round0_poly(
490 &mut self,
491 ) -> Result<UnivariatePoly<EF>, StackedReductionError> {
492 let l_skip = self.l_skip;
493 let skip_domain = 1 << l_skip;
494 let s_0_deg = sumcheck_round0_deg(l_skip, STACKED_REDUCTION_S_DEG);
495
496 let mut d_g_pos =
499 DeviceBuffer::<EF>::with_capacity_on(NUM_G * skip_domain, &self.device_ctx);
500 d_g_pos
501 .fill_zero_on(&self.device_ctx)
502 .map_err(StackedReductionError::FillZero)?;
503
504 let mut d_g_neg: Vec<DeviceBuffer<EF>> = (0..l_skip)
506 .map(|_| {
507 let b = DeviceBuffer::with_capacity_on(NUM_G * skip_domain, &self.device_ctx);
508 b.fill_zero_on(&self.device_ctx)
509 .map_err(StackedReductionError::FillZero)?;
510 Ok(b)
511 })
512 .collect::<Result<Vec<_>, StackedReductionError>>()?;
513
514 for ((trace_ptr, trace_height, trace_width), window) in zip(
516 mem::take(&mut self.trace_ptrs),
517 self.ht_diff_idxs.windows(2),
518 ) {
519 debug_assert_eq!(window[1] - window[0], trace_width);
520 let log_height = trace_height.ilog2();
521 let n = log_height as isize - l_skip as isize;
522
523 let d_g_output = if n >= 0 {
525 &mut d_g_pos
526 } else {
527 &mut d_g_neg[(-n - 1) as usize]
528 };
529
530 let block_sums_len = unsafe {
532 _stacked_reduction_r0_required_temp_buffer_size(
533 trace_height as u32,
534 trace_width as u32,
535 l_skip as u32,
536 )
537 } as usize;
538
539 if block_sums_len > self.d_block_sums.len() {
540 self.d_block_sums =
541 DeviceBuffer::<EF>::with_capacity_on(block_sums_len, &self.device_ctx);
542 }
543
544 unsafe {
545 let lambda_pows_ptr = self.d_lambda_pows.as_ptr().add(2 * window[0]);
547
548 stacked_reduction_sumcheck_round0(
549 &self.eq_r_ns,
550 trace_ptr,
551 lambda_pows_ptr,
552 &mut self.d_block_sums,
553 d_g_output,
554 trace_height,
555 trace_width,
556 l_skip,
557 self.device_ctx.stream.as_raw(),
558 )
559 .map_err(StackedReductionError::SumcheckRound0)?;
560 };
561 }
562
563 let s_0 = self.reconstruct_s0_from_g(d_g_pos, d_g_neg, s_0_deg)?;
565 self.mem.tracing_info("stacked_reduction_sumcheck round 0");
566
567 Ok(s_0)
568 }
569
570 fn reconstruct_s0_from_g(
579 &self,
580 d_g_pos: DeviceBuffer<EF>,
581 d_g_neg: Vec<DeviceBuffer<EF>>,
582 s_0_deg: usize,
583 ) -> Result<UnivariatePoly<EF>, StackedReductionError> {
584 let l_skip = self.l_skip;
585 let skip_domain = 1 << l_skip;
586 let large_uni_domain = (s_0_deg + 1).next_power_of_two(); let dft = Radix2BowersSerial;
588
589 let mut s_0_coeffs = vec![EF::ZERO; large_uni_domain];
591
592 let g_pos = d_g_pos.to_host_on(&self.device_ctx)?;
594 if !g_pos.iter().all(|&x| x == EF::ZERO) {
595 let e0 = eq_uni_poly::<F, EF>(l_skip, self.r_0);
597 let e1 = eq_uni_poly::<F, EF>(l_skip, self.r_0 * self.omega_skip);
598 let e2 = eq_uni_at_one_poly(l_skip, self.eq_const);
599
600 Self::ntt_multiply_and_add(
602 &dft,
603 large_uni_domain,
604 [e0.coeffs(), e1.coeffs(), e2.coeffs()],
605 [
606 &g_pos[0..skip_domain],
607 &g_pos[skip_domain..2 * skip_domain],
608 &g_pos[2 * skip_domain..3 * skip_domain],
609 ],
610 &mut s_0_coeffs,
611 );
612 }
613
614 for (bucket_idx, d_g_neg_bucket) in d_g_neg.into_iter().enumerate() {
616 let n_abs = bucket_idx + 1;
617 let g_neg = d_g_neg_bucket.to_host_on(&self.device_ctx)?;
618 if g_neg.iter().all(|&x| x == EF::ZERO) {
619 continue;
620 }
621
622 let l = l_skip - n_abs;
624 let omega_l = self.omega_skip.exp_power_of_2(n_abs);
625 let r_uni = self.r_0.exp_power_of_2(n_abs);
626
627 let ind = build_indicator_poly(l_skip, -(n_abs as isize));
629 let e0_base = eq_uni_poly::<F, EF>(l, r_uni);
630 let e1_base = eq_uni_poly::<F, EF>(l, r_uni * omega_l);
631 let e2_base = eq_uni_at_one_poly(l, self.eq_const);
632
633 let e0_neg = poly_multiply_ntt(&dft, e0_base.coeffs(), ind.coeffs(), skip_domain);
635 let e1_neg = poly_multiply_ntt(&dft, e1_base.coeffs(), ind.coeffs(), skip_domain);
636 let e2_neg = poly_multiply_ntt(&dft, e2_base.coeffs(), ind.coeffs(), skip_domain);
637
638 Self::ntt_multiply_and_add(
639 &dft,
640 large_uni_domain,
641 [&e0_neg, &e1_neg, &e2_neg],
642 [
643 &g_neg[0..skip_domain],
644 &g_neg[skip_domain..2 * skip_domain],
645 &g_neg[2 * skip_domain..3 * skip_domain],
646 ],
647 &mut s_0_coeffs,
648 );
649 }
650
651 s_0_coeffs.truncate(s_0_deg + 1);
652 Ok(UnivariatePoly::new(s_0_coeffs))
653 }
654
655 fn ntt_multiply_and_add(
658 dft: &Radix2BowersSerial,
659 domain_size: usize,
660 e_coeffs: [&[EF]; 3],
661 g_evals: [&[EF]; 3], out: &mut [EF],
663 ) {
664 let g_coeffs: [Vec<EF>; 3] = std::array::from_fn(|i| dft.idft(g_evals[i].to_vec()));
666
667 let mut e_padded = vec![EF::ZERO; domain_size * 3];
669 let mut g_padded = vec![EF::ZERO; domain_size * 3];
670 for i in 0..3 {
671 for (j, &c) in e_coeffs[i].iter().enumerate() {
672 e_padded[j * 3 + i] = c;
673 }
674 for (j, &c) in g_coeffs[i].iter().enumerate() {
675 g_padded[j * 3 + i] = c;
676 }
677 }
678
679 let e_evals_mat = dft.dft_batch(RowMajorMatrix::new(e_padded, 3));
681 let g_evals_mat = dft.dft_batch(RowMajorMatrix::new(g_padded, 3));
682
683 let mut s_evals = vec![EF::ZERO; domain_size];
685 for (j, s_j) in s_evals.iter_mut().enumerate() {
686 for i in 0..3 {
687 *s_j += e_evals_mat.values[j * 3 + i] * g_evals_mat.values[j * 3 + i];
688 }
689 }
690
691 let s_coeffs = dft.idft(s_evals);
693
694 for (o, c) in out.iter_mut().zip(s_coeffs) {
696 *o += c;
697 }
698 }
699
700 #[instrument("stacked_reduction_fold_ple", level = "debug", skip_all)]
701 fn fold_ple_evals(&mut self, u_0: EF) -> Result<(), StackedReductionError> {
702 let l_skip = self.l_skip;
703 let n_stack = self.n_stack;
704 let r_0 = self.r_0;
705 let omega_skip = self.omega_skip;
706 let n_max = self.n_max;
707 self.q_evals.clear();
708
709 let skip_domain = 1 << l_skip;
711 let inv_lagrange_denoms =
712 compute_barycentric_inv_lagrange_denoms(l_skip, &self.omega_skip_pows, u_0);
713 let d_inv_lagrange_denoms = inv_lagrange_denoms.to_device_on(&self.device_ctx)?;
714
715 for stacked in &self.stacked_per_commit {
716 let layout = stacked.layout();
717 let num_x = 1 << n_stack;
718 let stacked_width = layout.width();
719 debug_assert_eq!(layout.height(), 1 << (l_skip + n_stack));
720 let folded_evals =
721 DeviceBuffer::<EF>::with_capacity_on(num_x * stacked_width, &self.device_ctx);
722 folded_evals
724 .fill_zero_on(&self.device_ctx)
725 .map_err(StackedReductionError::FillZero)?;
726 let mut dst_offset = 0;
727 for trace in &stacked.traces {
728 if trace.width() == 0 || trace.height() == 0 {
729 continue;
730 }
731 let new_height = max(trace.height(), skip_domain) / skip_domain;
732
733 unsafe {
742 let dst = folded_evals.as_mut_ptr().add(dst_offset);
743 stacked_reduction_fold_ple(
744 trace.buffer().as_ptr(),
745 dst,
746 &self.d_omega_skip_pows,
747 &d_inv_lagrange_denoms,
748 trace.height(),
749 trace.width(),
750 l_skip,
751 self.device_ctx.stream.as_raw(),
752 )
753 .map_err(StackedReductionError::FoldPle)?;
754 }
755
756 dst_offset += new_height * trace.width();
757 }
758 self.q_evals.push(folded_evals);
759 }
760
761 let eq_uni_u0r0 = eval_eq_uni(l_skip, u_0, r_0);
763 let eq_uni_u0r0_rot = eval_eq_uni(l_skip, u_0, r_0 * omega_skip);
764 let eq_uni_u01 = eval_eq_uni_at_one(l_skip, u_0);
765 debug_assert_eq!(self.eq_r_ns.buffer.len(), 2 << n_max);
766 self.k_rot_ns.buffer = DeviceBuffer::with_capacity_on(2 << n_max, &self.device_ctx);
767 [EF::ZERO].copy_to_on(&mut self.k_rot_ns.buffer, &self.device_ctx)?;
768 unsafe {
769 initialize_k_rot_from_eq_segments(
772 &self.eq_r_ns,
773 &mut self.k_rot_ns.buffer,
774 eq_uni_u0r0_rot,
775 self.eq_const * eq_uni_u01,
776 n_max as u32,
777 self.device_ctx.stream.as_raw(),
778 )
779 .map_err(StackedReductionError::InitKRot)?;
780 }
781 vector_scalar_multiply_ext(
782 &mut self.eq_r_ns.buffer,
783 eq_uni_u0r0,
784 self.device_ctx.stream.as_raw(),
785 )
786 .map_err(StackedReductionError::VectorScalarMul)?;
787
788 (self.eq_stable, self.k_rot_stable) =
792 zip(r_0.exp_powers_of_2(), omega_skip.exp_powers_of_2())
793 .enumerate()
794 .skip(1)
795 .take(l_skip)
796 .map(|(n_abs, (r, omega_l))| {
797 let l = l_skip - n_abs;
798 let eq_uni = eval_eq_uni(l, u_0, r);
799 let eq_uni_rot = eval_eq_uni(l, u_0, r * omega_l);
800 let ind = eval_in_uni(l_skip, -(n_abs as isize), u_0);
801 (ind * eq_uni, ind * eq_uni_rot)
802 })
803 .unzip();
804 self.eq_stable.reverse();
805 self.k_rot_stable.reverse();
806 Ok(())
807 }
808
809 #[instrument("stacked_reduction_sumcheck", level = "debug", skip_all, fields(round = round))]
810 fn batch_sumcheck_poly_eval(
811 &mut self,
812 round: usize,
813 _u_prev: EF,
814 ) -> Result<[EF; STACKED_REDUCTION_S_DEG], StackedReductionError> {
815 let l_skip = self.l_skip;
816
817 let q_eval_ptrs = self.q_evals.iter().map(|q| q.as_ptr()).collect_vec();
818 q_eval_ptrs.copy_to_on(&mut self.d_q_eval_ptrs, &self.device_ctx)?;
819
820 if self.n_max >= (round - 1) {
821 let mut tmp = [EF::ZERO];
823 debug_assert_eq!(self.eq_stable.len(), l_skip + round - 1);
824 debug_assert_eq!(self.k_rot_stable.len(), l_skip + round - 1);
825 debug_assert!(self.eq_r_ns.buffer.len() > 1);
826 debug_assert!(self.k_rot_ns.buffer.len() > 1);
827 unsafe {
829 cuda_memcpy_on::<true, false>(
831 tmp.as_mut_ptr() as *mut c_void,
832 self.eq_r_ns.get_ptr(0) as *const c_void,
833 size_of::<EF>(),
834 &self.device_ctx,
835 )?;
836
837 self.eq_stable.push(tmp[0]);
838
839 cuda_memcpy_on::<true, false>(
841 tmp.as_mut_ptr() as *mut c_void,
842 self.k_rot_ns.get_ptr(0) as *const c_void,
843 size_of::<EF>(),
844 &self.device_ctx,
845 )?;
846
847 self.k_rot_stable.push(tmp[0]);
848 }
849 }
850 let accum_stride = STACKED_REDUCTION_S_DEG * D_EF;
851 let num_windows = self.ht_diff_idxs.len() - 1;
852 debug_assert!(self.d_accum.len() >= num_windows * accum_stride);
853
854 self.d_accum
855 .fill_zero_on(&self.device_ctx)
856 .map_err(StackedReductionError::FillZero)?;
857
858 let has_degenerate_window = self.ht_diff_idxs.windows(2).any(|window| {
859 let log_height = self.unstacked_cols[window[0]].log_height as usize;
860 log_height < l_skip + round
861 });
862 if has_degenerate_window {
863 self.eq_ub_per_trace
864 .copy_to_on(&mut self.d_eq_ub, &self.device_ctx)?;
865 }
866
867 for (window_idx, window) in self.ht_diff_idxs.windows(2).enumerate() {
868 let window_len = window[1] - window[0];
869 let unstacked_cols_ptr = unsafe { self.d_unstacked_cols.as_ptr().add(window[0]) };
871 let lambda_pows_ptr = unsafe { self.d_lambda_pows.as_ptr().add(2 * window[0]) };
874 let output_ptr = unsafe { self.d_accum.as_mut_ptr().add(window_idx * accum_stride) };
876
877 let log_height = self.unstacked_cols[window[0]].log_height as usize;
878
879 if log_height < l_skip + round {
880 let eq_r = self.eq_stable[log_height];
885 let k_rot_r = self.k_rot_stable[log_height];
886 let eq_ub_ptr = unsafe { self.d_eq_ub.as_ptr().add(window[0]) };
889 let stacked_height = self.stacked_height(round);
890 unsafe {
891 stacked_reduction_sumcheck_mle_round_degenerate(
892 &self.d_q_eval_ptrs,
893 eq_ub_ptr,
894 eq_r,
895 k_rot_r,
896 unstacked_cols_ptr,
897 lambda_pows_ptr,
898 output_ptr,
899 stacked_height,
900 window_len,
901 l_skip,
902 round,
903 self.device_ctx.stream.as_raw(),
904 )
905 .map_err(StackedReductionError::SumcheckMleRoundDegenerate)?;
906 }
907 } else {
908 let hypercube_dim = log_height - l_skip - round;
909 let num_y = 1 << hypercube_dim;
910 let stacked_height = self.stacked_height(round);
914 unsafe {
915 stacked_reduction_sumcheck_mle_round(
916 &self.d_q_eval_ptrs,
917 &self.eq_r_ns,
918 &self.k_rot_ns,
919 unstacked_cols_ptr,
920 lambda_pows_ptr,
921 output_ptr,
922 stacked_height,
923 window_len,
924 num_y,
925 self.sm_count,
926 self.device_ctx.stream.as_raw(),
927 )
928 .map_err(StackedReductionError::SumcheckMleRound)?;
929 };
930 }
931 }
932
933 let h_accum = self.d_accum.to_host_on(&self.device_ctx)?;
935 let s_evals_batch = h_accum[..num_windows * accum_stride]
936 .chunks_exact(accum_stride)
937 .map(reduce_raw_u64_to_ef)
938 .collect_vec();
939
940 Ok(from_fn(|i| {
941 s_evals_batch.iter().map(|evals| evals[i]).sum::<EF>()
942 }))
943 }
944
945 #[instrument("stacked_reduction_fold_mle", level = "debug", skip_all, fields(round = round))]
946 fn fold_mle_evals(&mut self, round: usize, u_round: EF) -> Result<(), StackedReductionError> {
947 debug_assert!(round <= self.n_stack);
948 let l_skip = self.l_skip;
949 let (folded_q_evals, input_ptrs, output_ptrs): (Vec<_>, Vec<_>, Vec<_>) = self
950 .q_evals
951 .iter()
952 .map(|q| {
953 let folded = DeviceBuffer::with_capacity_on(q.len() >> 1, &self.device_ctx);
954 let output_ptr = folded.as_mut_ptr();
955 (folded, q.as_ptr(), output_ptr)
956 })
957 .multiunzip();
958 input_ptrs.copy_to_on(&mut self.d_input_ptrs, &self.device_ctx)?;
959 output_ptrs.copy_to_on(&mut self.d_output_ptrs, &self.device_ctx)?;
960
961 let output_height = self.stacked_height(round + 1) as u32;
967 unsafe {
968 fold_mle(
969 &self.d_input_ptrs,
970 &self.d_output_ptrs,
971 &self.d_q_widths,
972 self.q_evals.len().try_into().unwrap(),
973 self.stacked_height(round + 1) as u32,
974 self.q_width_max * output_height,
975 u_round,
976 self.device_ctx.stream.as_raw(),
977 )
978 .map_err(StackedReductionError::FoldMle)?;
979 }
980 self.q_evals = folded_q_evals;
981
982 if self.n_max >= (round - 1) {
983 let input_max_n = self.cur_max_n(round);
984 let output_max_n = input_max_n.saturating_sub(1);
985 let output_len = 1 << input_max_n;
986
987 let mut buffer = DeviceBuffer::<EF>::with_capacity_on(output_len, &self.device_ctx);
988 [EF::ZERO].copy_to_on(&mut buffer, &self.device_ctx)?;
989 unsafe {
993 let mut output = EqEvalSegments::from_raw_parts(buffer, output_max_n);
994 if input_max_n != 0 {
995 triangular_fold_mle(
996 &mut output,
997 &self.eq_r_ns,
998 u_round,
999 output_max_n,
1000 self.device_ctx.stream.as_raw(),
1001 )
1002 .map_err(StackedReductionError::TriangularFoldMle)?;
1003 }
1004 self.eq_r_ns = output;
1005 }
1006
1007 let mut buffer = DeviceBuffer::<EF>::with_capacity_on(output_len, &self.device_ctx);
1008 [EF::ZERO].copy_to_on(&mut buffer, &self.device_ctx)?;
1009 unsafe {
1013 let mut output = EqEvalSegments::from_raw_parts(buffer, output_max_n);
1014 if input_max_n != 0 {
1015 triangular_fold_mle(
1016 &mut output,
1017 &self.k_rot_ns,
1018 u_round,
1019 output_max_n,
1020 self.device_ctx.stream.as_raw(),
1021 )
1022 .map_err(StackedReductionError::TriangularFoldMle)?;
1023 }
1024 self.k_rot_ns = output;
1025 }
1026 } else {
1027 assert_eq!(self.eq_r_ns.buffer.len(), 1);
1028 assert_eq!(self.k_rot_ns.buffer.len(), 1);
1029 }
1030 for (s, eq_ub) in zip(&self.unstacked_cols, &mut self.eq_ub_per_trace) {
1031 if round + l_skip > s.log_height as usize {
1032 debug_assert_eq!(s.stacked_row_idx % (1 << s.log_height), 0);
1035 let b = (s.stacked_row_idx >> (l_skip + round - 1)) & 1;
1036 *eq_ub *= eval_eq_mle(&[u_round], &[F::from_bool(b == 1)]);
1037 }
1038 }
1039 Ok(())
1040 }
1041
1042 #[instrument(level = "debug", skip_all)]
1043 fn get_stacked_openings(&self) -> Result<Vec<Vec<EF>>, StackedReductionError> {
1044 let lengths = self.q_evals.iter().map(DeviceBuffer::len).collect_vec();
1045 let total_len = lengths.iter().sum();
1046 let mut host = EF::zero_vec(total_len);
1047
1048 let mut offset = 0;
1049 for (q, &len) in zip(&self.q_evals, &lengths) {
1050 unsafe {
1051 cuda_memcpy_on::<true, false>(
1052 host.as_mut_ptr().add(offset) as *mut c_void,
1053 q.as_ptr() as *const c_void,
1054 len * size_of::<EF>(),
1055 &self.device_ctx,
1056 )?;
1057 }
1058 offset += len;
1059 }
1060 self.device_ctx
1061 .stream
1062 .to_host_sync()
1063 .map_err(MemCopyError::from)?;
1064
1065 let mut offset = 0;
1066 Ok(lengths
1067 .into_iter()
1068 .map(|len| {
1069 let next = offset + len;
1070 let values = host[offset..next].to_vec();
1071 offset = next;
1072 values
1073 })
1074 .collect())
1075 }
1076}
1077
1078fn build_indicator_poly(l_skip: usize, n: isize) -> UnivariatePoly<EF> {
1080 let n_abs = (-n) as usize;
1081 let l = l_skip - n_abs;
1082 let scale = EF::ONE.halve().exp_u64(n_abs as u64);
1083 let mut coeffs = vec![EF::ZERO; 1 << l_skip];
1084 for k in 0..(1 << n_abs) {
1085 coeffs[k * (1 << l)] = scale;
1086 }
1087 UnivariatePoly::new(coeffs)
1088}
1089
1090fn eq_uni_at_one_poly(l: usize, scale: EF) -> UnivariatePoly<EF> {
1094 let n_inv = F::ONE.halve().exp_u64(l as u64);
1095 UnivariatePoly::new(vec![EF::from(n_inv) * scale; 1 << l])
1096}
1097
1098fn poly_multiply_ntt(dft: &Radix2BowersSerial, a: &[EF], b: &[EF], min_size: usize) -> Vec<EF> {
1100 let size = (a.len() + b.len() - 1).max(min_size).next_power_of_two();
1101 let mut a_pad = a.to_vec();
1102 a_pad.resize(size, EF::ZERO);
1103 let mut b_pad = b.to_vec();
1104 b_pad.resize(size, EF::ZERO);
1105 let a_evals = dft.dft(a_pad);
1106 let b_evals = dft.dft(b_pad);
1107 let c_evals: Vec<EF> = a_evals
1108 .into_iter()
1109 .zip(b_evals)
1110 .map(|(a, b)| a * b)
1111 .collect();
1112 dft.idft(c_evals)
1113}