Skip to main content

openvm_cuda_backend/
stacked_reduction.rs

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
51/// Degree of the sumcheck polynomial for stacked reduction.
52pub 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, /* trace_ptr */
76        usize,    /* height */
77        usize,    /* width */
78    )>,
79    unstacked_cols: Vec<UnstackedSlice>,
80    d_unstacked_cols: DeviceBuffer<UnstackedSlice>,
81    // boundary indices where heights change. all columns between two boundaries must have the same
82    // height. We can have multiple chunks of the same height if that is needed to reduce peak GPU
83    // memory.
84    ht_diff_idxs: Vec<usize>,
85    n_max: usize,
86
87    // Initially holds eq(r[1..=n], H_n) for n=0..=n_max but gets updated after each sumcheck round
88    // by some custom folding
89    eq_r_ns: EqEvalSegments<EF>,
90
91    // == After round 0 ==
92    q_evals: Vec<DeviceBuffer<EF>>, // get width from stacked_per_commit
93    // Stores folded eq values that won't change anymore (no more folding)
94    // Corresponds to log_height in 0..l_skip+round-1 _before_ round `round`. Gets updated with one
95    // new element after each round.
96    eq_stable: Vec<EF>,
97    k_rot_stable: Vec<EF>,
98
99    /// Stores the folded k_rot evaluations for `\kappa_\rot(x, r) = eq_n(rot^{-1}(x), r)` for each
100    /// `n` after each round. We use the [EqEvalSegments] type to guard the segment-based
101    /// memory layout.
102    k_rot_ns: EqEvalSegments<EF>,
103    /// Stores eq(u[1+n_T..round-1], b_{T,j}[..round-n_T-1])
104    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
115/// A struct for holding stacked pcs data. We only need the `MerkleTreeGpu` from `StackedPcsDataGpu`
116/// but we wrap the latter in an `Arc` to provide uniformity in dealing with common main and
117/// preprocessed/cached traces. This struct stores the unstacked `traces`, in prismalinear
118/// evaluation form, corresponding to the stacked pcs data.
119///
120/// Generic over the Merkle digest type `D`.  The default `D = Digest` preserves the existing
121/// BabyBear-Poseidon2 behaviour.
122pub struct StackedPcsData2<D = Digest> {
123    pub(crate) inner: Arc<StackedPcsDataGpu<F, D>>,
124    /// The unstacked traces corresponding to `inner`'s commitment.
125    pub(crate) traces: Vec<DeviceMatrix<F>>,
126}
127
128impl<D> StackedPcsData2<D> {
129    /// # Safety
130    /// `traces` must be the traces that were committed to in `pcs_data`.
131    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/// Pointer with length to location in a big device buffer. The device buffer is identified by
147/// `commit_idx`. Due to pecuarlities with `l_skip`, the length of the slice is defined as
148/// `max(2^log_height, 2^l_skip)`. In other words, this is a pointer to the slice
149/// `q[commit_idx].column(stacked_col_idx)[stacked_row_idx..stacked_row_idx + len]`.
150///
151/// # Safety
152/// - this type is `repr(C)` as it will cross FFI boundaries for CUDA usage.
153#[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    /// Current maximum `n` supported by eq_r_ns
172    fn cur_max_n(&self, round: usize) -> usize {
173        self.n_max - (round - 1)
174    }
175}
176
177/// Batch sumcheck to reduce trace openings, including rotations, to stacked matrix opening.
178///
179/// The `stacked_matrix, stacked_layout` should be the result of stacking the `traces` with
180/// parameters `l_skip` and `n_stack`.
181#[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    // Batching randomness
209    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    // Round 0: univariate sumcheck
224    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    // end round 0
237
238    let mut sumcheck_round_polys = Vec::with_capacity(n_stack);
239
240    // Rounds 1..=n_stack: MLE sumcheck
241    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        // NOTE: DeviceMatrix contains an Arc, so clone is lightweight [for now].
297        // PERF[jpw]: stack the traces all at once. To save memory, we could either stack and drop
298        // traces as we go or even at commit time only store the stacked matrix and drop the traces
299        // much earlier.
300        let common_main_traces = proving_ctx
301            .per_trace
302            .iter()
303            .map(|(_, air_ctx)| air_ctx.common_main.clone())
304            .collect_vec();
305        // SAFETY: common_main_traces commits to common_main_pcs_data
306        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                // SAFETY: committed.trace commits to committed.data
317                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        // layout per commit is sorted, first height is largest
402        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            // SAFETY: This is unused in round 0 and will be initialized properly after round 0.
467            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    /// SP_DEG=1 round 0: computes G0, G1, G2 on identity coset, then reconstructs s_0 on CPU.
479    ///
480    /// Key insight: Instead of computing s₀(Z) directly on 2 cosets with in-kernel NTT,
481    /// we compute three partial sums G0, G1, G2 on the identity coset only, then
482    /// reconstruct s₀ via CPU-side NTT-based polynomial multiplication.
483    #[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        // Accumulation buffers for G0, G1, G2 (on identity coset)
497        // d_g_pos: for n >= 0 traces
498        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        // d_g_neg[k]: for traces with |n| = k+1, where n < 0 and k in 0..l_skip
505        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        // Process each trace - call kernel for each, accumulating into appropriate bucket
515        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            // Select output bucket based on n
524            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            // Allocate block_sums buffer for intermediate reduction
531            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                // 2 per column for (eq, k_rot) - coeff_eq and coeff_rot
546                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        // CPU reconstruction: s₀(Z) = E0(Z)*G0(Z) + E1(Z)*G1(Z) + E2(Z)*G2(Z)
564        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    /// Reconstructs s_0 from G0, G1, G2 using NTT-based polynomial multiplication.
571    ///
572    /// For n >= 0:
573    /// - E0(Z) = eq_uni(l_skip, Z, r0)
574    /// - E1(Z) = eq_uni(l_skip, Z, r0*ω_skip)
575    /// - E2(Z) = eq_const * eq_uni_at_one(l_skip, Z)
576    ///
577    /// For n < 0: multiply each E by ind(Z) = eval_in_uni(l_skip, n, Z)
578    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(); // 2 * skip_domain
587        let dft = Radix2BowersSerial;
588
589        // Accumulate s_0 coefficients across all buckets
590        let mut s_0_coeffs = vec![EF::ZERO; large_uni_domain];
591
592        // --- Process n >= 0 bucket ---
593        let g_pos = d_g_pos.to_host_on(&self.device_ctx)?;
594        if !g_pos.iter().all(|&x| x == EF::ZERO) {
595            // Build E polynomials for n >= 0
596            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            // NTT-based multiplication: s_0 += E0*G0 + E1*G1 + E2*G2
601            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        // --- Process n < 0 buckets ---
615        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            // Adjusted parameters for n < 0
623            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            // Build E polynomials with indicator factor
628            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            // E_neg = E_base * ind (polynomial multiplication)
634            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    /// NTT-based polynomial multiplication following logup pattern.
656    /// Computes: out += sum_i E[i] * G[i]
657    fn ntt_multiply_and_add(
658        dft: &Radix2BowersSerial,
659        domain_size: usize,
660        e_coeffs: [&[EF]; 3],
661        g_evals: [&[EF]; 3], // G evaluations on identity coset
662        out: &mut [EF],
663    ) {
664        // 1. iDFT G evaluations to get G coefficients
665        let g_coeffs: [Vec<EF>; 3] = std::array::from_fn(|i| dft.idft(g_evals[i].to_vec()));
666
667        // 2. Prepare coefficient matrices, resize to domain_size
668        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        // 3. DFT batch to evaluation domain
680        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        // 4. Pointwise multiply and sum: s[j] = sum_i e[j][i] * g[j][i]
684        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        // 5. iDFT to get product coefficients
692        let s_coeffs = dft.idft(s_evals);
693
694        // 6. Add to output
695        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        // Precompute Lagrange denominators once (shared across all traces)
710        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            // We must fill with zeros because some parts will be left empty due to stacking
723            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                // Launch single-trace kernel for this trace
734                // SAFETY:
735                // - `trace.buffer()` is a valid device pointer for `trace.height() * trace.width()`
736                //   elements
737                // - `folded_evals` at `dst_offset` is valid for `new_height * trace.width()`
738                //   elements since we allocated `num_x * stacked_width` and traces fill
739                //   contiguously
740                // - `d_omega_skip_pows` and `d_inv_lagrange_denoms` have length `>= skip_domain`
741                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        // fold PLEs into MLEs for \eq and \kappa_\rot, using u_0
762        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            // SAFETY:
770            // - We allocated `k_rot_ns` with same capacity as `eq_r_ns` above.
771            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        // Compute the special eq values for n = -l_skip..0
789        // First in order -n = 1..=l_skip, then reverse to order n = -l_skip..0 corresponding to
790        // log_height = 0..l_skip
791        (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            // Move stable eq, k_rot to stable vectors
822            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            // SAFETY: size of eq_r_ns, k_rot_ns is currently 2 * 2^{n_max - round + 1}
828            unsafe {
829                // D2H copy of single EF element
830                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                // D2H copy of single EF element
840                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            // SAFETY: in bounds by construction of ht_diff_idxs
870            let unstacked_cols_ptr = unsafe { self.d_unstacked_cols.as_ptr().add(window[0]) };
871            // 2 per column for (eq, k_rot)
872            // SAFETY: in bounds by construction of lambda_pows
873            let lambda_pows_ptr = unsafe { self.d_lambda_pows.as_ptr().add(2 * window[0]) };
874            // SAFETY: `d_accum` has one accumulator slot per window.
875            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                // We are in the eq, k_rot stable regime
881                // This includes all n < 0 cases
882                // In this case, the `s` poly contribution is a constant and we don't need to
883                // interpolate
884                let eq_r = self.eq_stable[log_height];
885                let k_rot_r = self.k_rot_stable[log_height];
886                // SAFETY: the full per-trace buffer was copied above, so this window slice stays
887                // valid while all queued kernels run asynchronously.
888                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                // Allow the CUDA launcher to auto-tune grid.y (thread_window_stride) based on
911                // (num_y, window_len) and device SM count.
912
913                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        // D2H copy and reduce modulo P once after all window kernels have queued.
934        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        // SAFETY:
962        // - `d_input_ptrs` points to matrices with widths specified by `d_q_widths` and heights
963        //   `stacked_height(round) = stacked_height(round + 1) * 2`.
964        // - `d_output_ptrs` points to matrices just allocated with widths specified by `d_q_widths`
965        //   and heights `stacked_height(round + 1)`.
966        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            // SAFETY:
990            // - eq_r_ns has max_n equal to input_max_n
991            // - we allocate output for half the size of eq_r_ns
992            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            // SAFETY:
1010            // - k_rot_ns has max_n equal to input_max_n
1011            // - we allocate output for half the size of eq_r_ns
1012            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                // Folding above did nothing, and we update the eq(u[1+n_T..=round],
1033                // b_{T,j}[..=round-n_T-1]) value
1034                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
1078/// Build indicator polynomial: ind(Z) = sum_{k=0}^{2^{n_abs}-1} Z^{k * 2^l} / 2^{n_abs}
1079fn 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
1090/// eq_uni_at_one polynomial: eq_D(Z, 1) as a polynomial in Z
1091///
1092/// All coefficients are n_inv * scale where n_inv = 1 / 2^l
1093fn 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
1098/// NTT-based polynomial multiplication
1099fn 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}