Skip to main content

openvm_cuda_backend/logup_zerocheck/
mod.rs

1//! ## Async frees and peak memory
2//!
3//! Many logup/zerocheck GPU helpers allocate temporary device buffers and launch kernels. Those
4//! buffers are freed via cudaFreeAsync through the VPMM pool, so frees are only stream-ordered. If
5//! new allocations happen before the current stream is synchronized, physical peak memory can
6//! temporarily include both the old and new buffers even though MemTracker only sees the logical
7//! max. Callers that care about physical peak should ensure a current-stream sync after use (e.g.
8//! by calling `to_host_on()` on outputs or synchronizing the owning stream).
9
10use std::{
11    cmp::max,
12    collections::hash_map::Entry,
13    iter::{self, zip},
14    sync::Arc,
15};
16
17use itertools::{izip, Itertools};
18use openvm_cuda_common::{
19    copy::{MemCopyD2H, MemCopyH2D},
20    d_buffer::DeviceBuffer,
21    error::MemCopyError,
22    memory_manager::MemTracker,
23    stream::GpuDeviceCtx,
24};
25use openvm_stark_backend::{
26    air_builders::symbolic::SymbolicConstraints,
27    calculate_n_logup,
28    dft::Radix2BowersSerial,
29    p3_matrix::dense::RowMajorMatrix,
30    poly_common::{
31        eq_uni_poly, eval_eq_mle, eval_eq_sharp_uni, eval_eq_uni, eval_eq_uni_at_one,
32        UnivariatePoly,
33    },
34    proof::{column_openings_by_rot, BatchConstraintProof, GkrProof},
35    prover::{
36        fractional_sumcheck_gkr::Frac, poly::eq_sharp_uni_poly, stacked_pcs::StackedLayout,
37        sumcheck::sumcheck_round0_deg, ColMajorMatrix, DeviceMultiStarkProvingKey,
38        MatrixDimensions, ProvingContext,
39    },
40};
41use p3_dft::TwoAdicSubgroupDft;
42use p3_field::{Field, PrimeCharacteristicRing, TwoAdicField};
43use p3_util::{log2_ceil_usize, log2_strict_usize};
44use rustc_hash::FxHashMap;
45use tracing::{debug, info, info_span, instrument};
46
47use crate::{
48    base::DeviceMatrix,
49    cuda::{
50        logup_zerocheck::{fold_selectors_round0, interpolate_columns_gpu, MainMatrixPtrs},
51        sumcheck::batch_fold_mle,
52    },
53    data_transporter::transport_matrix_d2h_col_major,
54    error::LogupZerocheckError,
55    gpu_backend::GenericGpuBackend,
56    hash_scheme::GpuHashScheme,
57    logup_zerocheck::{
58        batch_mle::evaluate_zerocheck_batched, fold_ple::fold_ple_evals_rotate,
59        gkr_input::TraceInteractionMeta, round0::evaluate_round0_interactions_gpu,
60    },
61    poly::EqEvalLayers,
62    prelude::{EF, F},
63    sponge::GpuFiatShamirTranscript,
64    utils::compute_barycentric_inv_lagrange_denoms,
65};
66
67pub(crate) mod batch_mle;
68pub(crate) mod batch_mle_monomial;
69mod block_ctxs;
70mod errors;
71pub(crate) mod fold_ple;
72/// Fraction sumcheck via GKR
73mod fractional;
74/// Logup interaction evaluations for GKR input
75mod gkr_input;
76mod mle_round;
77mod round0;
78pub(crate) mod rules;
79
80use batch_mle::{evaluate_logup_batched, TraceCtx};
81use batch_mle_monomial::{
82    compute_lambda_combinations, compute_logup_combinations, get_num_monomials,
83    get_zerocheck_rules_len, trace_has_monomials, LogupCombinations, LogupMonomialBatch,
84    ZerocheckMonomialBatch, ZerocheckMonomialParYBatch,
85};
86pub use errors::*;
87pub use fractional::{fractional_sumcheck_gpu, make_synthetic_leaves, FractionalInputSize};
88use gkr_input::{collect_trace_interactions, log_gkr_input_evals};
89use round0::evaluate_round0_constraints_gpu;
90
91/// When `num_monomials >= DAG_FALLBACK_MONOMIAL_RATIO * rules_len`, use DAG evaluation
92/// instead of the monomial kernel for high num_y traces.
93/// This ratio can be tuned. Currently it is set to prefer the monomial kernel except for the
94/// Poseidon2Air where the DAG node size is much smaller than number of monomials.
95const DAG_FALLBACK_MONOMIAL_RATIO: usize = 2;
96// Batch MLE launch overhead dominates in the non-save-memory path, so keep at least this
97// much batching budget there.
98const BATCH_MLE_DEFAULT_MEMORY_FLOOR: usize = 6 << 30; // 6GiB
99
100#[inline]
101fn fractional_gkr_peak_memory_bytes(input_len: usize, peak_work_buffer_bytes: usize) -> usize {
102    input_len
103        .saturating_mul(std::mem::size_of::<Frac<EF>>())
104        .saturating_add(peak_work_buffer_bytes)
105}
106
107#[inline]
108pub(crate) fn air_width_for_mat(need_rot: bool, mat_width: usize) -> u32 {
109    if need_rot {
110        debug_assert_eq!(mat_width % 2, 0, "rotated matrices should have even width");
111        (mat_width / 2) as u32
112    } else {
113        mat_width as u32
114    }
115}
116
117#[allow(clippy::type_complexity)]
118#[instrument(level = "info", skip_all)]
119pub fn prove_zerocheck_and_logup_gpu<HS, TS>(
120    transcript: &mut TS,
121    mpk: &DeviceMultiStarkProvingKey<GenericGpuBackend<HS>>,
122    proving_ctx: &ProvingContext<GenericGpuBackend<HS>>,
123    save_memory: bool,
124    monomial_num_y_threshold: u32,
125    sm_count: u32,
126    device_ctx: &GpuDeviceCtx,
127) -> Result<(GkrProof<HS::SC>, BatchConstraintProof<HS::SC>, Vec<EF>), LogupZerocheckError>
128where
129    HS: GpuHashScheme,
130    TS: GpuFiatShamirTranscript<HS::SC>,
131{
132    let logup_gkr_span = info_span!("prover.rap_constraints.logup_gkr", phase = "prover").entered();
133    let l_skip = mpk.params.l_skip;
134    let constraint_degree = mpk.max_constraint_degree;
135    let num_traces = proving_ctx.per_trace.len();
136
137    // Traces are sorted
138    let n_max =
139        log2_strict_usize(proving_ctx.per_trace[0].1.common_main.height()).saturating_sub(l_skip);
140    // Gather interactions metadata, including interactions stacked layout which depends on trace
141    // heights
142    let mut total_interactions = 0u64;
143    let interactions_meta: Vec<_> = proving_ctx
144        .per_trace
145        .iter()
146        .map(|(air_idx, air_ctx)| {
147            let pk = &mpk.per_air[*air_idx];
148
149            let num_interactions = pk.vk.symbolic_constraints.interactions.len();
150            let height = air_ctx.common_main.height();
151            let log_height = log2_strict_usize(height);
152            let log_lifted_height = log_height.max(l_skip);
153            total_interactions += (num_interactions as u64) << log_lifted_height;
154            (num_interactions, log_lifted_height)
155        })
156        .collect();
157    // Implicitly, the width of this stacking should be 1
158    let n_logup = calculate_n_logup(l_skip, total_interactions);
159    // There's no stride threshold for `interactions_layout` because there's no univariate skip for
160    // GKR
161    let interactions_layout = StackedLayout::new(0, l_skip + n_logup, interactions_meta).unwrap();
162
163    // Grind to increase soundness of random sampling for LogUp
164    let logup_pow_witness = transcript
165        .grind_gpu(mpk.params.logup.pow_bits, device_ctx)
166        .map_err(LogupZerocheckError::Grind)?;
167    let alpha_logup = transcript.sample_ext();
168    let beta_logup = transcript.sample_ext();
169    debug!(%alpha_logup, %beta_logup);
170
171    let has_interactions = !interactions_layout.sorted_cols.is_empty();
172    let mut prover = LogupZerocheckGpu::new(
173        mpk,
174        proving_ctx,
175        n_logup,
176        interactions_layout,
177        alpha_logup,
178        beta_logup,
179        save_memory,
180        monomial_num_y_threshold,
181        sm_count,
182        device_ctx,
183    )?;
184    let n_global = prover.n_global;
185
186    let real_len: usize = total_interactions
187        .try_into()
188        .expect("total interactions should fit in usize");
189    let logical_len = 1 << (l_skip + n_logup);
190    prover
191        .mem
192        .emit_metrics_with_label("prover.before_gkr_input_evals");
193    prover.mem.reset_peak();
194    let sizes = FractionalInputSize::new(real_len, logical_len);
195    let peak_work_buffer_bytes = if has_interactions {
196        sizes.peak_work_buffer_bytes()
197    } else {
198        0
199    };
200    let (inputs, alpha) = if has_interactions {
201        log_gkr_input_evals(
202            &prover.trace_interactions,
203            mpk,
204            proving_ctx,
205            l_skip,
206            alpha_logup,
207            &prover.d_challenges,
208            real_len,
209            peak_work_buffer_bytes,
210            device_ctx,
211        )?
212    } else {
213        (DeviceBuffer::new(), EF::ZERO)
214    };
215    // Set memory limit for batch MLE based on the fractional-GKR peak: input buffer plus
216    // peak work buffers.
217    prover.gkr_mem_contribution =
218        fractional_gkr_peak_memory_bytes(inputs.len(), peak_work_buffer_bytes);
219    prover.memory_limit_bytes = prover.gkr_mem_contribution;
220    if !prover.save_memory {
221        prover.memory_limit_bytes = prover
222            .memory_limit_bytes
223            .max(BATCH_MLE_DEFAULT_MEMORY_FLOOR);
224    }
225    prover.mem.emit_metrics_with_label("prover.gkr_input_evals");
226
227    let (frac_sum_proof, mut xi) = fractional_sumcheck_gpu(
228        transcript,
229        inputs,
230        sizes,
231        alpha,
232        true,
233        &mut prover.mem,
234        device_ctx,
235    )?;
236    while xi.len() != l_skip + n_global {
237        xi.push(transcript.sample_ext());
238    }
239    debug!(?xi);
240    prover.xi = xi;
241
242    logup_gkr_span.exit();
243
244    // Note: this span includes ple_fold, but that function has no cuda synchronization so it does
245    // not include the kernel times for the actual folding
246    let round0_span = info_span!("prover.rap_constraints.round0", phase = "prover").entered();
247    // begin batch sumcheck
248    let mut sumcheck_round_polys = Vec::with_capacity(n_max);
249    let mut r = Vec::with_capacity(n_max + 1);
250    // batching randomness
251    let lambda = transcript.sample_ext();
252    debug!(%lambda);
253
254    let sp_0_polys = prover.sumcheck_uni_round0_polys(proving_ctx, lambda)?;
255    let s_0_cpu_span = info_span!("s'_0 -> s_0 cpu interpolations").entered();
256    let sp_0_deg = sumcheck_round0_deg(l_skip, constraint_degree);
257    let s_deg = constraint_degree + 1;
258    let s_0_deg = sumcheck_round0_deg(l_skip, s_deg);
259    let large_uni_domain = (s_0_deg + 1).next_power_of_two();
260    let dft = Radix2BowersSerial;
261    let s_0_logup_polys = {
262        let eq_sharp_uni = eq_sharp_uni_poly(&prover.xi[..l_skip]);
263        let mut eq_coeffs = eq_sharp_uni.into_coeffs();
264        eq_coeffs.resize(large_uni_domain, EF::ZERO);
265        let eq_evals = dft.dft(eq_coeffs);
266
267        let width = 2 * num_traces;
268        let mut sp_coeffs_mat = EF::zero_vec(width * large_uni_domain);
269        for (i, coeffs) in sp_0_polys[..2 * num_traces].iter().enumerate() {
270            for (j, &c_j) in coeffs.coeffs().iter().enumerate().take(sp_0_deg + 1) {
271                // SAFETY:
272                // - coeffs length is <= sp_0_deg + 1 <= s_0_deg < large_uni_domain
273                // - sp_coeffs_mat allocated for width
274                unsafe {
275                    *sp_coeffs_mat.get_unchecked_mut(j * width + i) = c_j;
276                }
277            }
278        }
279        let mut s_evals = dft.dft_batch(RowMajorMatrix::new(sp_coeffs_mat, width));
280        for (eq, row) in zip(eq_evals, s_evals.values.chunks_mut(width)) {
281            for x in row {
282                *x *= eq;
283            }
284        }
285        dft.idft_batch(s_evals)
286    };
287    let skip_domain_size = F::from_usize(1 << l_skip);
288    // logup sum claims (sum_{\hat p}, sum_{\hat q}) per present AIR
289    let (numerator_term_per_air, denominator_term_per_air): (Vec<_>, Vec<_>) = (0..num_traces)
290        .map(|trace_idx| {
291            let [sum_claim_p, sum_claim_q] = [0, 1].map(|is_denom| {
292                // Compute sum over D of s_0(Z) to get the sum claim
293                (0..=s_0_deg)
294                    .step_by(1 << l_skip)
295                    .map(|j| unsafe {
296                        // SAFETY: matrix is 2 * num_trace x large_uni_domain, s_0_deg <
297                        // large_uni_domain
298                        *s_0_logup_polys
299                            .values
300                            .get_unchecked(j * 2 * num_traces + 2 * trace_idx + is_denom)
301                    })
302                    .sum::<EF>()
303                    * skip_domain_size
304            });
305            transcript.observe_ext(sum_claim_p);
306            transcript.observe_ext(sum_claim_q);
307
308            (sum_claim_p, sum_claim_q)
309        })
310        .unzip();
311
312    let mu = transcript.sample_ext();
313    debug!(%mu);
314    let mu_pows = mu.powers().take(3 * num_traces).collect_vec();
315
316    let s_0_zc_poly = {
317        let eq_uni = eq_uni_poly::<F, _>(l_skip, prover.xi[0]);
318        let mut eq_coeffs = eq_uni.into_coeffs();
319        eq_coeffs.resize(large_uni_domain, EF::ZERO);
320        let eq_evals = dft.dft(eq_coeffs);
321
322        // Algebraically batch
323        let mut sp_coeffs = EF::zero_vec(large_uni_domain);
324        let mus = &mu_pows[2 * num_traces..];
325        let polys = &sp_0_polys[2 * num_traces..];
326        for (j, batch_coeff) in sp_coeffs.iter_mut().enumerate().take(sp_0_deg + 1) {
327            for (&mu, poly) in zip(mus, polys) {
328                *batch_coeff += mu * *poly.coeffs().get(j).unwrap_or(&EF::ZERO);
329            }
330        }
331        let mut s_evals = dft.dft(sp_coeffs);
332        for (eq, x) in zip(eq_evals, &mut s_evals) {
333            *x *= eq;
334        }
335        dft.idft(s_evals)
336    };
337
338    // Algebraically batch
339    let s_0_poly = UnivariatePoly::new(
340        zip(
341            s_0_logup_polys.values.chunks_exact(2 * num_traces),
342            s_0_zc_poly,
343        )
344        .take(s_0_deg + 1)
345        .map(|(logup_row, batched_zc)| {
346            let coeff = batched_zc
347                + zip(&mu_pows, logup_row)
348                    .map(|(&mu_j, &x)| mu_j * x)
349                    .sum::<EF>();
350            transcript.observe_ext(coeff);
351            coeff
352        })
353        .collect(),
354    );
355    drop(s_0_cpu_span);
356
357    let r_0 = transcript.sample_ext();
358    r.push(r_0);
359    debug!(round = 0, r_round = %r_0);
360    prover.prev_s_eval = s_0_poly.eval_at_point(r_0);
361    debug!("s_0(r_0) = {}", prover.prev_s_eval);
362
363    prover.fold_ple_evals(proving_ctx, r_0)?;
364    drop(round0_span);
365
366    // Sumcheck rounds:
367    // - each round the prover needs to compute univariate polynomial `s_round`. This poly is linear
368    //   since we are taking MLE of `evals`.
369    // - at end of each round, sample random `r_round` in `EF`
370    //
371    // `s_round` is degree `s_deg` so we evaluate it at `0, ..., =s_deg`. The prover skips
372    // evaluation at `0` because the verifier can infer it from the previous round's
373    // `s_{round-1}(r)` claim. The degree is constraint_degree + 1, where + 1 is from eq term
374    let mle_rounds_span =
375        info_span!("prover.rap_constraints.mle_rounds", phase = "prover").entered();
376    debug!(%s_deg);
377    for round in 1..=n_max {
378        let sp_round_evals = prover.sumcheck_polys_batch_eval(round, r[round - 1])?;
379        let batch_s = prover.compute_batch_s_poly(sp_round_evals, num_traces, round, &mu_pows);
380        let batch_s_evals = (1..=s_deg)
381            .map(|i| batch_s.eval_at_point(EF::from_usize(i)))
382            .collect_vec();
383        for &eval in &batch_s_evals {
384            transcript.observe_ext(eval);
385        }
386        sumcheck_round_polys.push(batch_s_evals);
387
388        let r_round = transcript.sample_ext();
389        debug!(%round, %r_round);
390        r.push(r_round);
391        prover.prev_s_eval = batch_s.eval_at_point(r_round);
392
393        prover.fold_mle_evals(round, r_round)?;
394    }
395    assert_eq!(r.len(), n_max + 1);
396
397    let column_openings = prover.into_column_openings()?;
398
399    let need_rot_per_trace = proving_ctx
400        .per_trace
401        .iter()
402        .map(|(air_idx, _)| mpk.per_air[*air_idx].vk.params.need_rot)
403        .collect::<Vec<_>>();
404    // Observe common main openings first, and then preprocessed/cached
405    for (need_rot, openings) in need_rot_per_trace.iter().zip(column_openings.iter()) {
406        for (claim, claim_rot) in column_openings_by_rot(&openings[0], *need_rot) {
407            transcript.observe_ext(claim);
408            transcript.observe_ext(claim_rot);
409        }
410    }
411    for (need_rot, openings) in need_rot_per_trace.iter().zip(column_openings.iter()) {
412        for part in openings.iter().skip(1) {
413            for (claim, claim_rot) in column_openings_by_rot(part, *need_rot) {
414                transcript.observe_ext(claim);
415                transcript.observe_ext(claim_rot);
416            }
417        }
418    }
419    drop(mle_rounds_span);
420
421    let batch_constraint_proof = BatchConstraintProof {
422        numerator_term_per_air,
423        denominator_term_per_air,
424        univariate_round_coeffs: s_0_poly.into_coeffs(),
425        sumcheck_round_polys,
426        column_openings,
427    };
428    let gkr_proof = GkrProof {
429        logup_pow_witness,
430        q0_claim: frac_sum_proof.fractional_sum.1,
431        claims_per_layer: frac_sum_proof.claims_per_layer,
432        sumcheck_polys: frac_sum_proof.sumcheck_polys,
433    };
434    Ok((gkr_proof, batch_constraint_proof, r))
435}
436
437pub struct LogupZerocheckGpu<'a, HS: GpuHashScheme> {
438    pub alpha_logup: EF,
439    pub beta_pows: Vec<EF>,
440    // [alpha, beta^0, beta^1, .., beta^max_interaction_len]
441    pub d_challenges: DeviceBuffer<EF>,
442
443    pub l_skip: usize,
444    n_logup: usize,
445    n_global: usize,
446
447    pub omega_skip: F,
448    pub omega_skip_pows: Vec<F>,
449    d_omega_skip_pows: DeviceBuffer<F>,
450
451    pub interactions_layout: StackedLayout,
452    pub constraint_degree: usize,
453    n_per_trace: Vec<isize>,
454    max_num_constraints: usize,
455    pub monomial_num_y_threshold: u32,
456    sm_count: u32,
457    // Available after GKR:
458    pub xi: Vec<EF>,
459    pub lambda_pows: Option<DeviceBuffer<EF>>,
460    /// Precomputed lambda combinations per AIR (indexed by air_idx). Set when lambda is sampled.
461    lambda_combinations: Vec<Option<DeviceBuffer<EF>>>,
462    /// Beta powers on device for logup MLE rounds.
463    d_beta_pows: DeviceBuffer<EF>,
464    /// Precomputed logup combinations per trace (indexed by trace_idx). Set when xi is sampled.
465    logup_combinations: Vec<Option<LogupCombinations>>,
466
467    // n_T => segment tree of eq(xi[j..1+n_T]) for j=1..={n_T-round+1} in _reverse_ layout
468    eq_xis: FxHashMap<usize, EqEvalLayers<EF>>,
469    eq_3b_per_trace: Vec<Vec<EF>>,
470    d_eq_3b_per_trace: Vec<DeviceBuffer<EF>>,
471    // Evaluations on hypercube only, for round 0
472    sels_per_trace_base: Vec<DeviceMatrix<F>>,
473    // After univariate round 0:
474    mat_evals_per_trace: Vec<Vec<DeviceMatrix<EF>>>,
475    sels_per_trace: Vec<DeviceMatrix<EF>>,
476    // Store public_values per trace (similar to CPU's EvalHelper)
477    public_values_per_trace: Vec<DeviceBuffer<F>>,
478    air_indices_per_trace: Vec<usize>,
479    zerocheck_tilde_evals: Vec<EF>,
480    logup_tilde_evals: Vec<[EF; 2]>,
481    needs_next_per_trace: Vec<bool>,
482
483    trace_interactions: Vec<Option<TraceInteractionMeta>>,
484    // round0: Round0Buffers,
485    pk: &'a DeviceMultiStarkProvingKey<GenericGpuBackend<HS>>,
486
487    // In round `j`, contains `s_{j-1}(r_{j-1})`
488    pub(crate) prev_s_eval: EF,
489    pub(crate) eq_ns: Vec<EF>,
490    pub(crate) eq_sharp_ns: Vec<EF>,
491
492    mem: MemTracker,
493    save_memory: bool,
494
495    /// Fractional-GKR peak budget: input buffer plus peak work buffers.
496    gkr_mem_contribution: usize,
497    /// Memory limit for batch MLE intermediate buffers (set after GKR input eval)
498    memory_limit_bytes: usize,
499    device_ctx: GpuDeviceCtx,
500}
501
502impl<'a, HS: GpuHashScheme> LogupZerocheckGpu<'a, HS> {
503    #[allow(clippy::too_many_arguments)]
504    fn new(
505        pk: &'a DeviceMultiStarkProvingKey<GenericGpuBackend<HS>>,
506        proving_ctx: &ProvingContext<GenericGpuBackend<HS>>,
507        n_logup: usize,
508        interactions_layout: StackedLayout,
509        alpha_logup: EF,
510        beta_logup: EF,
511        save_memory: bool,
512        monomial_num_y_threshold: u32,
513        sm_count: u32,
514        device_ctx: &GpuDeviceCtx,
515    ) -> Result<Self, LogupZerocheckError> {
516        let mem = MemTracker::start("prover.logup_zerocheck_prover");
517        let l_skip = pk.params.l_skip;
518        let omega_skip = F::two_adic_generator(l_skip);
519        let omega_skip_pows = omega_skip.powers().take(1 << l_skip).collect_vec();
520        let d_omega_skip_pows = omega_skip_pows.to_device_on(device_ctx)?;
521        let num_airs_present = proving_ctx.per_trace.len();
522
523        let constraint_degree = pk.max_constraint_degree;
524
525        let max_interaction_length = pk
526            .per_air
527            .iter()
528            .map(|air_pk| air_pk.other_data.interaction_rules.max_fields_len)
529            .max()
530            .unwrap_or(0);
531        let beta_pows = beta_logup
532            .powers()
533            .take(max_interaction_length + 1)
534            .collect_vec();
535        let challenges = [&[alpha_logup], &beta_pows[..]].concat();
536        let d_challenges = challenges.to_device_on(device_ctx)?;
537        let d_beta_pows = beta_pows.to_device_on(device_ctx)?;
538
539        let n_per_trace: Vec<isize> = proving_ctx
540            .common_main_traces()
541            .map(|(_, t)| log2_strict_usize(t.height()) as isize - l_skip as isize)
542            .collect();
543        let n_max = n_per_trace[0].max(0) as usize;
544        let n_global = max(n_max, n_logup);
545        info!(%n_global, %n_logup);
546
547        let max_num_constraints = pk
548            .per_air
549            .iter()
550            .map(|air_pk| {
551                air_pk
552                    .vk
553                    .symbolic_constraints
554                    .constraints
555                    .constraint_idx
556                    .len()
557            })
558            .max()
559            .unwrap_or(0);
560
561        // Collect interaction metadata for GPU execution (evaluations still run on CPU for now).
562        let trace_interactions = collect_trace_interactions(pk, proving_ctx, &interactions_layout);
563
564        let needs_next_per_trace = proving_ctx
565            .per_trace
566            .iter()
567            .map(|(air_idx, _)| pk.per_air[*air_idx].vk.params.need_rot)
568            .collect::<Vec<_>>();
569
570        Ok(Self {
571            alpha_logup,
572            beta_pows,
573            d_challenges,
574            l_skip,
575            n_logup,
576            n_global,
577            omega_skip,
578            omega_skip_pows,
579            d_omega_skip_pows,
580            interactions_layout,
581            constraint_degree,
582            n_per_trace,
583            max_num_constraints,
584            sm_count,
585            xi: vec![],
586            lambda_pows: None,
587            lambda_combinations: (0..pk.per_air.len()).map(|_| None).collect(),
588            d_beta_pows,
589            logup_combinations: (0..num_airs_present).map(|_| None).collect(),
590            eq_xis: FxHashMap::default(),
591            eq_3b_per_trace: vec![],
592            d_eq_3b_per_trace: vec![],
593            sels_per_trace_base: vec![],
594            mat_evals_per_trace: vec![],
595            sels_per_trace: vec![],
596            public_values_per_trace: proving_ctx
597                .per_trace
598                .iter()
599                .map(|(_, air_ctx)| {
600                    if air_ctx.public_values.is_empty() {
601                        Ok(DeviceBuffer::new())
602                    } else {
603                        air_ctx.public_values.to_device_on(device_ctx)
604                    }
605                })
606                .collect::<Result<Vec<_>, _>>()?,
607            air_indices_per_trace: proving_ctx
608                .per_trace
609                .iter()
610                .map(|(air_idx, _)| *air_idx)
611                .collect_vec(),
612            zerocheck_tilde_evals: vec![EF::ZERO; num_airs_present],
613            logup_tilde_evals: vec![[EF::ZERO; 2]; num_airs_present],
614            needs_next_per_trace,
615            trace_interactions,
616            pk,
617            prev_s_eval: EF::ZERO,
618            eq_ns: Vec::with_capacity(n_max + 1),
619            eq_sharp_ns: Vec::with_capacity(n_max + 1),
620            mem,
621            save_memory,
622            gkr_mem_contribution: 0,
623            memory_limit_bytes: 0, // Set after GKR input eval
624            monomial_num_y_threshold,
625            device_ctx: device_ctx.clone(),
626        })
627    }
628
629    // PERF[jpw]: we could return evals and batch zerocheck poly by degree before interpolating.
630    // Cannot do it for logup because we need to calculate the sum claims.
631    #[instrument(name = "prover.rap_constraints.ple_round0", level = "info", skip_all)]
632    fn sumcheck_uni_round0_polys(
633        &mut self,
634        ctx: &ProvingContext<GenericGpuBackend<HS>>,
635        lambda: EF,
636    ) -> Result<Vec<UnivariatePoly<EF>>, LogupZerocheckError> {
637        self.mem
638            .emit_metrics_with_label("prover.batch_constraints.before_round0");
639        self.mem.reset_peak();
640        let n_logup = self.n_logup;
641        let l_skip = self.l_skip;
642        let xi = &self.xi;
643        let h_lambda_pows = lambda.powers().take(self.max_num_constraints).collect_vec();
644        self.lambda_pows = Some(if !h_lambda_pows.is_empty() {
645            h_lambda_pows.to_device_on(&self.device_ctx)?
646        } else {
647            DeviceBuffer::new()
648        });
649        // Precompute lambda combinations for all AIRs with monomials
650        let lambda_pows_ref = self.lambda_pows.as_ref().unwrap();
651        for (air_idx, air_pk) in self.pk.per_air.iter().enumerate() {
652            if air_pk.other_data.zerocheck_monomials.is_some() {
653                self.lambda_combinations[air_idx] = Some(
654                    compute_lambda_combinations(
655                        self.pk,
656                        air_idx,
657                        lambda_pows_ref,
658                        &self.device_ctx,
659                    )
660                    .map_err(LogupZerocheckError::LambdaCombinations)?,
661                );
662            }
663        }
664        let num_present_airs = ctx.per_trace.len();
665        debug_assert_eq!(num_present_airs, self.n_per_trace.len());
666
667        self.eq_3b_per_trace = ctx
668            .per_trace
669            .iter()
670            .enumerate()
671            .map(|(trace_idx, (air_idx, _))| {
672                let vk = &self.pk.per_air[*air_idx].vk;
673                let num_interactions = vk.num_interactions();
674                if num_interactions > 0 {
675                    let n = self.n_per_trace[trace_idx];
676                    let n_lift = n.max(0) as usize;
677                    let mut b_vec = vec![F::ZERO; n_logup - n_lift];
678                    let mut weights = Vec::with_capacity(num_interactions);
679                    for interaction_idx in 0..num_interactions {
680                        let stacked_idx = self
681                            .interactions_layout
682                            .get(trace_idx, interaction_idx)
683                            .unwrap()
684                            .row_idx;
685                        let mut b_int = stacked_idx >> (l_skip + n_lift);
686                        for bit in &mut b_vec {
687                            *bit = F::from_bool(b_int & 1 == 1);
688                            b_int >>= 1;
689                        }
690                        let weight =
691                            eval_eq_mle(&self.xi[l_skip + n_lift..l_skip + n_logup], &b_vec);
692                        weights.push(weight);
693                    }
694                    weights
695                } else {
696                    vec![]
697                }
698            })
699            .collect_vec();
700        self.d_eq_3b_per_trace = self
701            .eq_3b_per_trace
702            .iter()
703            .map(|eq_3bs| {
704                if eq_3bs.is_empty() {
705                    Ok(DeviceBuffer::new())
706                } else {
707                    eq_3bs.to_device_on(&self.device_ctx)
708                }
709            })
710            .collect::<Result<Vec<_>, _>>()?;
711
712        // Precompute logup combinations for all traces with interaction monomials
713        for (trace_idx, (air_idx, _)) in ctx.per_trace.iter().enumerate() {
714            let air_pk = &self.pk.per_air[*air_idx];
715            if air_pk.other_data.interaction_monomials.is_some()
716                && !self.eq_3b_per_trace[trace_idx].is_empty()
717            {
718                self.logup_combinations[trace_idx] = Some(
719                    compute_logup_combinations(
720                        self.pk,
721                        *air_idx,
722                        &self.d_beta_pows,
723                        &self.d_eq_3b_per_trace[trace_idx],
724                        &self.eq_3b_per_trace[trace_idx],
725                        &self.beta_pows,
726                        &self.device_ctx,
727                    )
728                    .map_err(LogupZerocheckError::LogupCombinations)?,
729                );
730            }
731        }
732
733        // PERF[jpw]: we could also build the layers for different n in a transposed way using
734        // eq_nonoverlapping_stage_ext, which is more memory efficient
735        let mut eq_xi_one_layer = None;
736        for &n in &self.n_per_trace {
737            let n_lift = n.max(0) as usize;
738            if let Entry::Vacant(entry) = self.eq_xis.entry(n_lift) {
739                let layer_0 = match &eq_xi_one_layer {
740                    Some(layer_0) => Arc::clone(layer_0),
741                    None => {
742                        let layer_0 = EqEvalLayers::one_layer(&self.device_ctx)
743                            .map_err(LogupZerocheckError::EqEvalLayers)?;
744                        eq_xi_one_layer = Some(Arc::clone(&layer_0));
745                        layer_0
746                    }
747                };
748                let layers = EqEvalLayers::new_rev_with_one(
749                    n_lift,
750                    xi[l_skip..l_skip + n_lift].iter().rev(),
751                    layer_0,
752                    &self.device_ctx,
753                )
754                .map_err(LogupZerocheckError::EqEvalLayers)?;
755                entry.insert(layers);
756            }
757        }
758
759        self.sels_per_trace_base = self
760            .n_per_trace
761            .iter()
762            .map(|&n| {
763                let n_lift = n.max(0) as usize;
764                let height = 1 << n_lift;
765                let mut cols = F::zero_vec(3 * height);
766                cols[height..2 * height - 1].fill(F::ONE); // is_transition
767                cols[0] = F::ONE; // is_first
768                cols[2 * height + height - 1] = F::ONE; // is_last
769                let d_cols = cols.to_device_on(&self.device_ctx)?;
770                Ok(DeviceMatrix::new(Arc::new(d_cols), height, 3))
771            })
772            .collect::<Result<Vec<_>, MemCopyError>>()?;
773
774        let selectors_base = self.sels_per_trace_base.clone();
775
776        // All (numer, denom) pairs per present AIR for logup, followed by 1 zerocheck poly per
777        // present AIR
778        let mut batch_sp_poly = vec![UnivariatePoly::new(vec![]); 3 * num_present_airs];
779        let d_lambda_pows = self
780            .lambda_pows
781            .as_ref()
782            .expect("lambda powers must be set before round-0 evaluation");
783
784        // Loop through one AIR at a time; it is more efficient to do everything for one AIR
785        // together
786        for (trace_idx, ((air_idx, air_ctx), &n, selectors_cube, public_values, eq_3bs)) in izip!(
787            &ctx.per_trace,
788            &self.n_per_trace,
789            &selectors_base,
790            &self.public_values_per_trace,
791            &self.eq_3b_per_trace,
792        )
793        .enumerate()
794        {
795            debug!("starting batch constraints for air_idx={air_idx} (trace_idx={trace_idx})");
796            let single_pk = &self.pk.per_air[*air_idx];
797            // Includes both plain AIR constraints and symbolic interactions
798            let single_air_constraints =
799                SymbolicConstraints::from(&single_pk.vk.symbolic_constraints);
800            let local_constraint_deg = single_pk.vk.max_constraint_degree as usize;
801            debug_assert_eq!(
802                single_air_constraints.max_constraint_degree(),
803                local_constraint_deg
804            );
805            assert!(
806                local_constraint_deg <= self.constraint_degree,
807                "Max constraint degree ({local_constraint_deg}) of AIR {air_idx} exceeds the global maximum {}",
808                self.constraint_degree
809            );
810
811            let log_large_domain = log2_ceil_usize(local_constraint_deg << l_skip);
812            let omega_root = F::two_adic_generator(log_large_domain);
813
814            assert!(!xi.is_empty(), "xi vector must not be empty");
815
816            let height = air_ctx.common_main.height();
817            let mut main_parts = Vec::with_capacity(air_ctx.cached_mains.len() + 1);
818            for committed in &air_ctx.cached_mains {
819                main_parts.push(committed.trace.buffer().as_ptr());
820            }
821            main_parts.push(air_ctx.common_main.buffer().as_ptr());
822            let d_main_parts = main_parts.to_device_on(&self.device_ctx)?;
823
824            let n_lift = n.max(0) as usize;
825            let eq_xi_tree = &self.eq_xis[&n_lift];
826            let max_temp_bytes = self.memory_limit_bytes;
827            // local_constraint_deg = 0 means no constraints. The only way that linear constraints
828            // on trace polynomials could vanish on 2^l_skip points is if the constraint polynomial
829            // is identically zero. Thus for local_constraint_deg = 0 or 1, we must have `s'_0 = 0`.
830            let num_cosets_zc = local_constraint_deg.saturating_sub(1);
831            let sum_buffer = evaluate_round0_constraints_gpu(
832                single_pk,
833                selectors_cube.buffer(),
834                &d_main_parts,
835                public_values,
836                eq_xi_tree.get_ptr(n_lift),
837                d_lambda_pows,
838                1 << l_skip,
839                1 << n_lift,
840                height as u32,
841                num_cosets_zc as u32,
842                omega_root,
843                max_temp_bytes,
844                &self.device_ctx,
845            )?;
846            if !sum_buffer.is_empty() {
847                let q_evals = sum_buffer.to_host_on(&self.device_ctx)?;
848                let q = {
849                    // Make q_evals row-major, with columns <> cosets
850                    let mut values = EF::zero_vec(num_cosets_zc << l_skip);
851                    for coset_idx in 0..num_cosets_zc {
852                        for i in 0..1 << l_skip {
853                            values[i * num_cosets_zc + coset_idx] =
854                                q_evals[(coset_idx << l_skip) + i];
855                        }
856                    }
857                    UnivariatePoly::from_geometric_cosets_evals_idft(
858                        RowMajorMatrix::new(values, num_cosets_zc),
859                        omega_root,
860                        omega_root,
861                    )
862                };
863                // sp_0 = (Z^{2^l_skip} - 1) * q
864                let sp_0_deg = sumcheck_round0_deg(l_skip, local_constraint_deg);
865                let coeffs = (0..=sp_0_deg)
866                    .map(|i| {
867                        let mut c = -*q.coeffs().get(i).unwrap_or(&EF::ZERO);
868                        if i >= 1 << l_skip {
869                            c += q.coeffs()[i - (1 << l_skip)];
870                        }
871                        c
872                    })
873                    .collect_vec();
874                debug_assert_eq!(
875                    coeffs.iter().step_by(1 << l_skip).copied().sum::<EF>(),
876                    EF::ZERO,
877                    "Zerocheck sum is not zero for air_id: {}",
878                    ctx.per_trace[trace_idx].0
879                );
880
881                batch_sp_poly[2 * num_present_airs + trace_idx] = UnivariatePoly::new(coeffs);
882            }
883
884            // PERF: we could use an interaction-specific constraint degree here
885            let num_cosets_logup = local_constraint_deg;
886            let sum = evaluate_round0_interactions_gpu(
887                single_pk,
888                &single_air_constraints,
889                selectors_cube.buffer(),
890                &d_main_parts,
891                public_values,
892                eq_xi_tree.get_ptr(n_lift),
893                &self.beta_pows,
894                eq_3bs,
895                1 << l_skip,
896                1 << n_lift,
897                height as u32,
898                num_cosets_logup as u32,
899                omega_root,
900                max_temp_bytes,
901                &self.device_ctx,
902            )?;
903            if !sum.is_empty() {
904                let evals = sum.to_host_on(&self.device_ctx)?;
905                let (mut numer, denom): (Vec<EF>, Vec<EF>) =
906                    evals.into_iter().map(|frac| (frac.p, frac.q)).unzip();
907                if n.is_negative() {
908                    // normalize for lifting
909                    let norm_factor = F::from_u32(1 << n.unsigned_abs()).inverse();
910                    for s in &mut numer {
911                        *s *= norm_factor;
912                    }
913                }
914                let mut numer_values = EF::zero_vec(num_cosets_logup << l_skip);
915                let mut denom_values = EF::zero_vec(num_cosets_logup << l_skip);
916                for coset_idx in 0..num_cosets_logup {
917                    for i in 0..1 << l_skip {
918                        let src = (coset_idx << l_skip) + i;
919                        let dst = i * num_cosets_logup + coset_idx;
920                        numer_values[dst] = numer[src];
921                        denom_values[dst] = denom[src];
922                    }
923                }
924                // Logup uses cosets 1, g^1, g^2, ... (init = 1, shift = omega_root)
925                batch_sp_poly[2 * trace_idx] = UnivariatePoly::from_geometric_cosets_evals_idft(
926                    RowMajorMatrix::new(numer_values, num_cosets_logup),
927                    omega_root,
928                    F::ONE, // init = 1 for identity coset
929                );
930                batch_sp_poly[2 * trace_idx + 1] = UnivariatePoly::from_geometric_cosets_evals_idft(
931                    RowMajorMatrix::new(denom_values, num_cosets_logup),
932                    omega_root,
933                    F::ONE, // init = 1 for identity coset
934                );
935            }
936        }
937        self.mem
938            .emit_metrics_with_label("prover.batch_constraints.round0");
939        Ok(batch_sp_poly)
940    }
941
942    // Note: there are no gpu sync points in this function, so span does not indicate kernel times
943    #[instrument(name = "LogupZerocheck::fold_ple_evals", level = "debug", skip_all)]
944    fn fold_ple_evals(
945        &mut self,
946        ctx: &ProvingContext<GenericGpuBackend<HS>>,
947        r_0: EF,
948    ) -> Result<(), LogupZerocheckError> {
949        let l_skip = self.l_skip;
950        let inv_lagrange_denoms_r0 =
951            compute_barycentric_inv_lagrange_denoms(l_skip, &self.omega_skip_pows, r_0);
952        let d_inv_lagrange_denoms_r0 = inv_lagrange_denoms_r0.to_device_on(&self.device_ctx)?;
953
954        let mut mem_limit = self.gkr_mem_contribution;
955        // GPU folding for mat_evals_per_trace
956        self.mat_evals_per_trace = ctx
957            .per_trace
958            .iter()
959            .map(|(air_idx, air_ctx)| {
960                let air_pk = &self.pk.per_air[*air_idx];
961                let need_rot = air_pk.vk.params.need_rot;
962                let mut results: Vec<DeviceMatrix<EF>> = Vec::new();
963
964                // Preprocessed (if exists)
965                if let Some(committed) = &air_pk.preprocessed_data {
966                    let trace = &committed.trace;
967                    let folded = fold_ple_evals_rotate(
968                        l_skip,
969                        &self.d_omega_skip_pows,
970                        trace,
971                        &d_inv_lagrange_denoms_r0,
972                        need_rot,
973                        &self.device_ctx,
974                    )?;
975                    results.push(folded);
976                }
977
978                // Cached mains
979                for committed in &air_ctx.cached_mains {
980                    let trace = &committed.trace;
981                    let folded = fold_ple_evals_rotate(
982                        l_skip,
983                        &self.d_omega_skip_pows,
984                        trace,
985                        &d_inv_lagrange_denoms_r0,
986                        need_rot,
987                        &self.device_ctx,
988                    )?;
989                    results.push(folded);
990                }
991
992                // Common main
993                let trace = &air_ctx.common_main;
994                let folded = fold_ple_evals_rotate(
995                    l_skip,
996                    &self.d_omega_skip_pows,
997                    trace,
998                    &d_inv_lagrange_denoms_r0,
999                    need_rot,
1000                    &self.device_ctx,
1001                )?;
1002                mem_limit = mem_limit.saturating_sub(folded.buffer().len() * size_of::<EF>());
1003                results.push(folded);
1004
1005                Ok(results)
1006            })
1007            .collect::<Result<Vec<_>, FoldPleError>>()?;
1008        if self.save_memory {
1009            self.memory_limit_bytes = mem_limit;
1010        }
1011
1012        // GPU folding for sels_per_trace (rotate=false, only need offset=0)
1013        self.sels_per_trace = std::mem::take(&mut self.sels_per_trace_base)
1014            .into_iter()
1015            .enumerate()
1016            .map(|(trace_idx, selectors_cube)| {
1017                let n = self.n_per_trace[trace_idx];
1018                let num_x = selectors_cube.height();
1019                debug_assert_eq!(num_x, 1 << n.max(0));
1020                debug_assert_eq!(selectors_cube.width(), 3);
1021                let (l, r) = if n.is_negative() {
1022                    (
1023                        l_skip.wrapping_add_signed(n),
1024                        r_0.exp_power_of_2(-n as usize),
1025                    )
1026                } else {
1027                    (l_skip, r_0)
1028                };
1029                let omega = F::two_adic_generator(l);
1030                let is_first = eval_eq_uni_at_one(l, r);
1031                let is_last = eval_eq_uni_at_one(l, r * omega);
1032                let folded_buf = DeviceBuffer::<EF>::with_capacity_on(num_x * 3, &self.device_ctx);
1033                unsafe {
1034                    fold_selectors_round0(
1035                        folded_buf.as_mut_ptr(),
1036                        selectors_cube.buffer().as_ptr(),
1037                        is_first,
1038                        is_last,
1039                        num_x,
1040                        self.device_ctx.stream.as_raw(),
1041                    )
1042                    .map_err(LogupZerocheckError::FoldSelectorsRound0)?;
1043                }
1044                Ok(DeviceMatrix::new(Arc::new(folded_buf), num_x, 3))
1045            })
1046            .collect::<Result<Vec<_>, LogupZerocheckError>>()?;
1047
1048        // GPU scalar multiplication for eq_xi_per_trace and eq_sharp_per_trace
1049        // Compute scalars on CPU (small computation)
1050        let eq_r0 = eval_eq_uni(l_skip, self.xi[0], r_0);
1051        let eq_sharp_r0 = eval_eq_sharp_uni(&self.omega_skip_pows, &self.xi[..l_skip], r_0);
1052        self.eq_ns.push(eq_r0);
1053        self.eq_sharp_ns.push(eq_sharp_r0);
1054        for tree in self.eq_xis.values_mut() {
1055            // trim the back (which corresponds to r_{j-1}) because we don't need it anymore
1056            if tree.layers.len() > 1 {
1057                tree.layers.pop();
1058            }
1059        }
1060
1061        self.mem
1062            .emit_metrics_with_label("prover.batch_constraints.fold_ple_evals");
1063        Ok(())
1064    }
1065
1066    #[instrument(
1067        name = "LogupZerocheck::sumcheck_polys_batch_eval",
1068        level = "info",
1069        skip_all,
1070        fields(round = round)
1071    )]
1072    fn sumcheck_polys_batch_eval(
1073        &mut self,
1074        round: usize,
1075        r_prev: EF,
1076    ) -> Result<Vec<Vec<EF>>, LogupZerocheckError> {
1077        let sp_deg = self.constraint_degree;
1078
1079        // Per-trace outputs (filled as we go)
1080        let mut zc_out: Vec<Vec<EF>> = vec![vec![EF::ZERO; sp_deg]; self.n_per_trace.len()];
1081        let mut logup_out: Vec<[Vec<EF>; 2]> =
1082            vec![[vec![EF::ZERO; sp_deg], vec![EF::ZERO; sp_deg]]; self.n_per_trace.len()];
1083
1084        // Keep early interpolations alive for duration of kernels
1085        let mut _keepalive_interpolated: Vec<DeviceMatrix<EF>> = Vec::new();
1086
1087        let mut late_eval: Vec<TraceCtx> = Vec::new(); // round == n_lift + 1
1088        let mut early_eval: Vec<TraceCtx> = Vec::new(); // round <= n_lift
1089
1090        // First, handle traces in original order and split into cases
1091        for (trace_idx, (&n, mats, sels, eq_3bs, public_vals, &air_idx)) in izip!(
1092            self.n_per_trace.iter(),
1093            self.mat_evals_per_trace.iter(),
1094            self.sels_per_trace.iter(),
1095            self.d_eq_3b_per_trace.iter(),
1096            self.public_values_per_trace.iter(),
1097            self.air_indices_per_trace.iter()
1098        )
1099        .enumerate()
1100        {
1101            let pk = &self.pk.per_air[air_idx];
1102            let dag = &pk.vk.symbolic_constraints;
1103            let has_constraints = dag.constraints.num_constraints() > 0;
1104            let has_interactions = !dag.interactions.is_empty();
1105            if !has_constraints && !has_interactions {
1106                continue;
1107            }
1108
1109            let n_lift = n.max(0) as usize;
1110            let norm_factor_denom = 1 << (-n).max(0);
1111            let norm_factor = F::from_usize(norm_factor_denom).inverse();
1112            let has_preprocessed = pk.preprocessed_data.is_some();
1113            let need_rot = pk.vk.params.need_rot;
1114            let first_main_idx = usize::from(has_preprocessed);
1115            let eq_xi_tree = &self.eq_xis[&n_lift];
1116
1117            if round > n_lift {
1118                // Case A
1119                if round == n_lift + 1 {
1120                    // A.1: evaluate directly at (num_x=1, num_y=1)
1121                    let prep_ptr = if has_preprocessed {
1122                        MainMatrixPtrs {
1123                            data: mats[0].buffer().as_ptr(),
1124                            air_width: air_width_for_mat(need_rot, mats[0].width()),
1125                        }
1126                    } else {
1127                        MainMatrixPtrs {
1128                            data: std::ptr::null(),
1129                            air_width: 0,
1130                        }
1131                    };
1132                    let main_ptrs: Vec<MainMatrixPtrs<EF>> = mats[first_main_idx..]
1133                        .iter()
1134                        .map(|m| MainMatrixPtrs {
1135                            data: m.buffer().as_ptr(),
1136                            air_width: air_width_for_mat(need_rot, m.width()),
1137                        })
1138                        .collect_vec();
1139                    let main_ptrs_dev = main_ptrs.to_device_on(&self.device_ctx)?;
1140
1141                    late_eval.push(TraceCtx {
1142                        trace_idx,
1143                        air_idx,
1144                        n_lift,
1145                        num_y: 1,
1146                        has_constraints,
1147                        has_interactions,
1148                        norm_factor,
1149                        eq_xi_ptr: eq_xi_tree.get_ptr(0),
1150                        sels_ptr: sels.buffer().as_ptr(),
1151                        prep_ptr,
1152                        main_ptrs_dev,
1153                        public_ptr: public_vals.as_ptr(),
1154                        eq_3bs_ptr: eq_3bs.as_ptr(),
1155                    });
1156                } else {
1157                    // A.2: scale tilde evals only
1158                    if has_constraints {
1159                        let tilde_eval = &mut self.zerocheck_tilde_evals[trace_idx];
1160                        *tilde_eval *= r_prev;
1161                        // zc_out not set, will be handled directly from tilde eval in
1162                        // compute_batch_s
1163                    }
1164                    if has_interactions {
1165                        for x in self.logup_tilde_evals[trace_idx].iter_mut() {
1166                            *x *= r_prev;
1167                        }
1168                        // logup_out not set, will be handled directly from tilde eval in
1169                        // compute_batch_s
1170                    }
1171                }
1172            } else {
1173                // Case B: interpolate columns and evaluate (num_x = s_deg, num_y = height/2)
1174                let log_num_y = n_lift - round;
1175                let num_y = 1 << log_num_y;
1176                let height = 2 * num_y;
1177                debug_assert_eq!(height, mats[0].height());
1178
1179                let mut columns: Vec<*const EF> = Vec::new();
1180                columns.extend(
1181                    iter::once(sels)
1182                        .chain(mats.iter())
1183                        .flat_map(|m| {
1184                            assert_eq!(m.height(), height);
1185                            (0..m.width())
1186                                .map(|col| m.buffer().as_ptr().wrapping_add(col * m.height()))
1187                        })
1188                        .collect_vec(),
1189                );
1190                let interpolated = DeviceMatrix::<EF>::with_capacity_on(
1191                    sp_deg * num_y,
1192                    columns.len(),
1193                    &self.device_ctx,
1194                );
1195                let d_columns = columns.to_device_on(&self.device_ctx)?;
1196                unsafe {
1197                    interpolate_columns_gpu(
1198                        interpolated.buffer(),
1199                        &d_columns,
1200                        sp_deg,
1201                        num_y,
1202                        self.device_ctx.stream.as_raw(),
1203                    )
1204                    .map_err(|e| LogupZerocheckError::InterpolateColumns(e.into()))?;
1205                }
1206
1207                let interpolated_height = interpolated.height();
1208                let mut widths_so_far = 0usize;
1209                let sels_ptr = interpolated
1210                    .buffer()
1211                    .as_ptr()
1212                    .wrapping_add(widths_so_far * interpolated_height);
1213                widths_so_far += 3;
1214                let prep_ptr = if has_preprocessed {
1215                    MainMatrixPtrs {
1216                        data: interpolated
1217                            .buffer()
1218                            .as_ptr()
1219                            .wrapping_add(widths_so_far * interpolated_height),
1220                        air_width: air_width_for_mat(need_rot, mats[0].width()),
1221                    }
1222                } else {
1223                    MainMatrixPtrs {
1224                        data: std::ptr::null(),
1225                        air_width: 0,
1226                    }
1227                };
1228                if has_preprocessed {
1229                    widths_so_far += mats[0].width();
1230                }
1231                let main_ptrs: Vec<MainMatrixPtrs<EF>> = mats[first_main_idx..]
1232                    .iter()
1233                    .map(|m| {
1234                        let main_ptr = MainMatrixPtrs {
1235                            data: interpolated
1236                                .buffer()
1237                                .as_ptr()
1238                                .wrapping_add(widths_so_far * interpolated_height),
1239                            air_width: air_width_for_mat(need_rot, m.width()),
1240                        };
1241                        widths_so_far += m.width();
1242                        main_ptr
1243                    })
1244                    .collect_vec();
1245                debug_assert_eq!(widths_so_far, interpolated.width());
1246                let main_ptrs_dev = main_ptrs.to_device_on(&self.device_ctx)?;
1247
1248                _keepalive_interpolated.push(interpolated);
1249                let eq_xi_ptr = eq_xi_tree.get_ptr(log_num_y);
1250
1251                early_eval.push(TraceCtx {
1252                    trace_idx,
1253                    air_idx,
1254                    n_lift,
1255                    num_y: num_y as u32,
1256                    has_constraints,
1257                    has_interactions,
1258                    norm_factor,
1259                    eq_xi_ptr,
1260                    sels_ptr,
1261                    prep_ptr,
1262                    main_ptrs_dev,
1263                    public_ptr: public_vals.as_ptr(),
1264                    eq_3bs_ptr: eq_3bs.as_ptr(),
1265                });
1266            }
1267        }
1268
1269        let d_challenges_ptr = self.d_challenges.as_ptr();
1270
1271        // Late traces (num_y=1): always use monomial
1272        let late_logup_traces: Vec<_> = late_eval.iter().filter(|t| t.has_interactions).collect();
1273        if !late_logup_traces.is_empty() {
1274            let logup_combs: Vec<_> = late_logup_traces
1275                .iter()
1276                .map(|t| {
1277                    self.logup_combinations[t.trace_idx]
1278                        .as_ref()
1279                        .expect("missing logup monomial combinations for late trace")
1280                })
1281                .collect();
1282            let batch = LogupMonomialBatch::new(
1283                late_logup_traces.iter().copied(),
1284                self.pk,
1285                &logup_combs,
1286                &self.device_ctx,
1287            )?;
1288            let out = batch
1289                .evaluate(1)
1290                .map_err(LogupZerocheckError::MleInteractionEval)?;
1291            let host = out.to_host_on(&self.device_ctx)?;
1292            for (i, trace_idx) in batch.trace_indices().enumerate() {
1293                self.logup_tilde_evals[trace_idx][0] = host[i].p * late_logup_traces[i].norm_factor;
1294                self.logup_tilde_evals[trace_idx][1] = host[i].q;
1295            }
1296        }
1297        let late_mono_traces: Vec<_> = late_eval
1298            .iter()
1299            .filter(|t| trace_has_monomials(t, self.pk))
1300            .collect();
1301        if !late_mono_traces.is_empty() {
1302            let lambda_combs: Vec<_> = late_mono_traces
1303                .iter()
1304                .map(|t| self.lambda_combinations[t.air_idx].as_ref().unwrap())
1305                .collect();
1306            let batch = ZerocheckMonomialBatch::new(
1307                late_mono_traces,
1308                self.pk,
1309                &lambda_combs,
1310                &self.device_ctx,
1311            )?;
1312            let out = batch
1313                .evaluate(1)
1314                .map_err(LogupZerocheckError::MleConstraintEval)?;
1315            let host = out.to_host_on(&self.device_ctx)?;
1316            for (i, trace_idx) in batch.trace_indices().enumerate() {
1317                self.zerocheck_tilde_evals[trace_idx] = host[i];
1318                // zc_out not set for num_x=1, handled from tilde_eval in compute_batch_s
1319            }
1320        }
1321
1322        // Logup for early traces: partition by num_y threshold
1323        if !early_eval.is_empty() {
1324            evaluate_logup_batched(
1325                &early_eval,
1326                self.pk,
1327                d_challenges_ptr,
1328                sp_deg as u32,
1329                self.monomial_num_y_threshold,
1330                &self.logup_combinations,
1331                &mut logup_out,
1332                &mut self.logup_tilde_evals,
1333                self.memory_limit_bytes,
1334                &self.device_ctx,
1335            )
1336            .map_err(LogupZerocheckError::MleInteractionEval)?;
1337        }
1338
1339        // Early traces (num_y>1): partition by threshold for zerocheck path
1340        let (low_early, high_early): (Vec<&TraceCtx>, Vec<&TraceCtx>) = early_eval
1341            .iter()
1342            .filter(|t| t.has_constraints)
1343            .partition(|t| t.num_y <= self.monomial_num_y_threshold);
1344
1345        // Partition high num_y traces by monomial-to-rules ratio
1346        // (traces without monomials are skipped - they contribute zero)
1347        let (high_dag_traces, high_mono_traces): (Vec<&TraceCtx>, Vec<&TraceCtx>) =
1348            high_early.iter().partition(|t| {
1349                let num_monomials = get_num_monomials(t, self.pk);
1350                let rules_len = get_zerocheck_rules_len(t, self.pk);
1351                // Use DAG when monomial expansion significantly increased the term count
1352                num_monomials as usize >= DAG_FALLBACK_MONOMIAL_RATIO * rules_len
1353            });
1354
1355        // DAG evaluation for high num_y traces with high monomial-to-rules ratio
1356        if !high_dag_traces.is_empty() {
1357            let lambda_pows = self.lambda_pows.as_ref().unwrap();
1358            evaluate_zerocheck_batched(
1359                high_dag_traces,
1360                self.pk,
1361                lambda_pows,
1362                sp_deg as u32,
1363                &mut zc_out,
1364                self.memory_limit_bytes,
1365                &self.device_ctx,
1366            )
1367            .map_err(LogupZerocheckError::MleConstraintEval)?;
1368        }
1369
1370        // Par-Y monomial kernel for high num_y traces
1371        if !high_mono_traces.is_empty() {
1372            let lambda_combs: Vec<_> = high_mono_traces
1373                .iter()
1374                .map(|t| self.lambda_combinations[t.air_idx].as_ref().unwrap())
1375                .collect();
1376            let batch = ZerocheckMonomialParYBatch::new(
1377                high_mono_traces,
1378                self.pk,
1379                &lambda_combs,
1380                self.sm_count,
1381                sp_deg as u32,
1382                None,
1383                &self.device_ctx,
1384            )?;
1385            let out = batch
1386                .evaluate(sp_deg as u32)
1387                .map_err(LogupZerocheckError::MleConstraintEval)?;
1388            let host = out.to_host_on(&self.device_ctx)?;
1389            for (i, trace_idx) in batch.trace_indices().enumerate() {
1390                zc_out[trace_idx].copy_from_slice(&host[(i * sp_deg)..((i + 1) * sp_deg)]);
1391            }
1392        }
1393
1394        // Monomial zerocheck for low num_y traces
1395        let low_mono_traces = low_early;
1396        if !low_mono_traces.is_empty() {
1397            let lambda_combs: Vec<_> = low_mono_traces
1398                .iter()
1399                .map(|t| self.lambda_combinations[t.air_idx].as_ref().unwrap())
1400                .collect();
1401            let batch = ZerocheckMonomialBatch::new(
1402                low_mono_traces,
1403                self.pk,
1404                &lambda_combs,
1405                &self.device_ctx,
1406            )?;
1407            let out = batch
1408                .evaluate(sp_deg as u32)
1409                .map_err(LogupZerocheckError::MleConstraintEval)?;
1410            let host = out.to_host_on(&self.device_ctx)?;
1411            for (i, trace_idx) in batch.trace_indices().enumerate() {
1412                zc_out[trace_idx].copy_from_slice(&host[(i * sp_deg)..((i + 1) * sp_deg)]);
1413            }
1414        }
1415
1416        Ok(logup_out.into_iter().flatten().chain(zc_out).collect())
1417    }
1418
1419    #[instrument(level = "debug", skip_all, fields(round = round))]
1420    fn compute_batch_s_poly(
1421        &mut self,
1422        sp_round_evals: Vec<Vec<EF>>,
1423        num_traces: usize,
1424        round: usize,
1425        mu_pows: &[EF],
1426    ) -> UnivariatePoly<EF> {
1427        debug_assert_eq!(sp_round_evals.len(), 3 * num_traces);
1428        debug_assert_eq!(sp_round_evals.len(), mu_pows.len());
1429        let constraint_degree = self.constraint_degree;
1430        let mut sp_head_zc = vec![EF::ZERO; constraint_degree];
1431        let mut sp_head_logup = vec![EF::ZERO; constraint_degree];
1432        let mut sp_tail = EF::ZERO;
1433        for (trace_idx, &n) in self.n_per_trace.iter().enumerate() {
1434            let n_lift = n.max(0) as usize;
1435            let zc_idx = 2 * num_traces + trace_idx;
1436            let numer_idx = 2 * trace_idx;
1437            let denom_idx = numer_idx + 1;
1438            if round == n_lift + 1 {
1439                let eq_r_acc = *self.eq_ns.last().unwrap();
1440                let eq_sharp_r_acc = *self.eq_sharp_ns.last().unwrap();
1441                self.zerocheck_tilde_evals[trace_idx] *= eq_r_acc;
1442                self.logup_tilde_evals[trace_idx][0] *= eq_sharp_r_acc;
1443                self.logup_tilde_evals[trace_idx][1] *= eq_sharp_r_acc;
1444            }
1445            if round <= n_lift {
1446                for i in 0..constraint_degree {
1447                    sp_head_zc[i] += mu_pows[zc_idx] * sp_round_evals[zc_idx][i];
1448                    sp_head_logup[i] += mu_pows[numer_idx] * sp_round_evals[numer_idx][i]
1449                        + mu_pows[denom_idx] * sp_round_evals[denom_idx][i];
1450                }
1451            } else {
1452                sp_tail += mu_pows[zc_idx] * self.zerocheck_tilde_evals[trace_idx]
1453                    + mu_pows[numer_idx] * self.logup_tilde_evals[trace_idx][0]
1454                    + mu_pows[denom_idx] * self.logup_tilde_evals[trace_idx][1];
1455            }
1456        }
1457        let s_deg = constraint_degree + 1;
1458        let l_skip = self.l_skip;
1459        // With eq(xi,r) contributions
1460        let mut sp_head_evals = vec![EF::ZERO; s_deg];
1461        for i in 0..constraint_degree {
1462            sp_head_evals[i + 1] = self.eq_ns[round - 1] * sp_head_zc[i]
1463                + self.eq_sharp_ns[round - 1] * sp_head_logup[i];
1464        }
1465        // We need to derive s'(0).
1466        // We use that s_j(0) + s_j(1) = s_{j-1}(r_{j-1})
1467        let xi_cur = self.xi[l_skip + round - 1];
1468        {
1469            let eq_xi_0 = EF::ONE - xi_cur;
1470            let eq_xi_1 = xi_cur;
1471            sp_head_evals[0] =
1472                (self.prev_s_eval - eq_xi_1 * sp_head_evals[1] - sp_tail) * eq_xi_0.inverse();
1473        }
1474        // s' has degree s_deg - 1
1475        let sp_head = UnivariatePoly::lagrange_interpolate(
1476            &(0..s_deg).map(F::from_usize).collect_vec(),
1477            &sp_head_evals,
1478        );
1479        // eq(xi, X) = (2 * xi - 1) * X + (1 - xi)
1480        // Compute s(X) = eq(xi, X) * s'_head(X) + s'_tail * X (s'_head now contains eq(..,r))
1481        // s(X) has degree s_deg
1482        let mut coeffs = sp_head.into_coeffs();
1483        coeffs.push(EF::ZERO);
1484        let b = EF::ONE - xi_cur;
1485        let a = xi_cur - b;
1486        for i in (0..s_deg).rev() {
1487            coeffs[i + 1] = a * coeffs[i] + b * coeffs[i + 1];
1488        }
1489        coeffs[0] *= b;
1490        coeffs[1] += sp_tail;
1491        UnivariatePoly::new(coeffs)
1492    }
1493
1494    #[instrument(name = "LogupZerocheck::fold_mle_evals", level = "debug", skip_all, fields(round = round))]
1495    fn fold_mle_evals(&mut self, round: usize, r_round: EF) -> Result<(), LogupZerocheckError> {
1496        // Assumes that input_mats are sorted by height
1497        let batch_fold = |input_mats: Vec<DeviceMatrix<EF>>| -> Result<Vec<DeviceMatrix<EF>>, LogupZerocheckError> {
1498            let num_matrices = input_mats.partition_point(|mat| mat.height() > 1);
1499            let mut max_output_cells = 0;
1500            let (log_output_heights, widths, mut output_mats): (Vec<_>, Vec<_>, Vec<_>) =
1501                input_mats
1502                    .iter()
1503                    .take(num_matrices)
1504                    .map(|mat| {
1505                        let height = mat.height();
1506                        let width = mat.width();
1507                        let output_height = height >> 1;
1508                        max_output_cells = max(max_output_cells, output_height * width);
1509                        let output_mat =
1510                            DeviceMatrix::<EF>::with_capacity_on(output_height, width, &self.device_ctx);
1511                        (output_height.ilog2() as u8, width as u32, output_mat)
1512                    })
1513                    .multiunzip();
1514
1515            let input_ptrs = input_mats
1516                .iter()
1517                .take(num_matrices)
1518                .map(|mat| mat.buffer().as_ptr())
1519                .collect_vec();
1520            let output_ptrs = output_mats
1521                .iter()
1522                .map(|mat| mat.buffer().as_mut_ptr())
1523                .collect_vec();
1524
1525            let d_input_ptrs = input_ptrs.to_device_on(&self.device_ctx)?;
1526            let d_output_ptrs = output_ptrs.to_device_on(&self.device_ctx)?;
1527            let d_log_output_heights = log_output_heights.to_device_on(&self.device_ctx)?;
1528            let d_widths = widths.to_device_on(&self.device_ctx)?;
1529
1530            unsafe {
1531                batch_fold_mle(
1532                    &d_input_ptrs,
1533                    &d_output_ptrs,
1534                    &d_widths,
1535                    num_matrices.try_into().unwrap(),
1536                    &d_log_output_heights,
1537                    max_output_cells.try_into().unwrap(),
1538                    r_round,
1539                    self.device_ctx.stream.as_raw(),
1540                )
1541                .map_err(LogupZerocheckError::BatchFoldMle)?;
1542            }
1543            output_mats.extend_from_slice(&input_mats[num_matrices..]);
1544            Ok(output_mats)
1545        };
1546
1547        // Fold mat_evals_per_trace: Vec<Vec<DeviceMatrix<EF>>>
1548        self.mat_evals_per_trace = {
1549            let lengths = self
1550                .mat_evals_per_trace
1551                .iter()
1552                .map(|v| v.len())
1553                .collect_vec();
1554            let input_mats = std::mem::take(&mut self.mat_evals_per_trace)
1555                .into_iter()
1556                .flatten()
1557                .collect_vec();
1558            let mut output_mats = batch_fold(input_mats)?.into_iter();
1559            lengths
1560                .into_iter()
1561                .map(|len| output_mats.by_ref().take(len).collect())
1562                .collect()
1563        };
1564        if self.save_memory {
1565            self.memory_limit_bytes = self.gkr_mem_contribution.saturating_sub(
1566                self.mat_evals_per_trace
1567                    .iter()
1568                    .flatten()
1569                    .map(|m| m.buffer().len() * size_of::<EF>())
1570                    .sum(),
1571            );
1572        }
1573
1574        // Fold sels_per_trace: Vec<DeviceMatrix<EF>>
1575        self.sels_per_trace = batch_fold(std::mem::take(&mut self.sels_per_trace))?;
1576
1577        for tree in self.eq_xis.values_mut() {
1578            // trim the back (which corresponds to r_{j-1}) because we don't need it anymore
1579            if tree.layers.len() > 1 {
1580                tree.layers.pop();
1581            }
1582        }
1583        let xi = self.xi[self.l_skip + round - 1];
1584        let eq_r = eval_eq_mle(&[xi], &[r_round]);
1585        self.eq_ns.push(self.eq_ns[round - 1] * eq_r);
1586        self.eq_sharp_ns.push(self.eq_sharp_ns[round - 1] * eq_r);
1587        Ok(())
1588    }
1589
1590    #[instrument(
1591        name = "LogupZerocheck::into_column_openings",
1592        level = "debug",
1593        skip_all
1594    )]
1595    fn into_column_openings(mut self) -> Result<Vec<Vec<Vec<EF>>>, LogupZerocheckError> {
1596        let num_airs_present = self.mat_evals_per_trace.len();
1597        let mut column_openings = Vec::with_capacity(num_airs_present);
1598
1599        // At the end, we've folded all MLEs so they only have one row equal to evaluation at `\vec
1600        // r`.
1601        for (&need_rot, mat_evals) in self
1602            .needs_next_per_trace
1603            .iter()
1604            .zip(std::mem::take(&mut self.mat_evals_per_trace))
1605        {
1606            // GPU matrices are doubled-width (original + rotated), so we need to split them
1607            // First, copy all matrices to host and split them
1608            let mut split_mats: Vec<Option<ColMajorMatrix<EF>>> = mat_evals
1609                .into_iter()
1610                .map(|mat| {
1611                    let mat_host = transport_matrix_d2h_col_major(&mat, &self.device_ctx)?;
1612                    let width = mat_host.width();
1613                    let height = mat_host.height();
1614                    debug_assert_eq!(height, 1, "Matrices should have height=1 after folding");
1615                    let air_width = if need_rot {
1616                        debug_assert_eq!(
1617                            width % 2,
1618                            0,
1619                            "GPU matrices should have doubled width (original + rotated)"
1620                        );
1621                        width / 2
1622                    } else {
1623                        width
1624                    };
1625
1626                    // Split doubled-width matrix into original and rotated parts
1627                    let values = &mat_host.values;
1628                    let orig: Vec<EF> = (0..air_width)
1629                        .map(|col| values[col * height]) // height=1, so values[col]
1630                        .collect();
1631                    let rot: Option<Vec<EF>> = if need_rot {
1632                        Some(
1633                            (air_width..width)
1634                                .map(|col| values[col * height]) // height=1, so values[col]
1635                                .collect(),
1636                        )
1637                    } else {
1638                        None
1639                    };
1640
1641                    Ok(vec![
1642                        Some(ColMajorMatrix::new(orig, air_width)),
1643                        rot.map(|mat| ColMajorMatrix::new(mat, air_width)),
1644                    ])
1645                })
1646                .collect::<Result<Vec<_>, MemCopyError>>()?
1647                .into_iter()
1648                .flatten()
1649                .collect();
1650
1651            // Order of mats after splitting is:
1652            // - preprocessed (if has_preprocessed),
1653            // - preprocessed_rot (if has_preprocessed),
1654            // - cached(0), cached(0)_rot, ...
1655            // - common_main
1656            // - common_main_rot
1657            // For column openings, we pop common_main, common_main_rot and put it at the front
1658            assert_eq!(
1659                split_mats.len() % 2,
1660                0,
1661                "Should have even number of matrices after splitting"
1662            );
1663            let common_main_rot = split_mats.pop().unwrap();
1664            let common_main = split_mats.pop().unwrap();
1665
1666            let openings_of_air = iter::once(&[common_main, common_main_rot] as &[_])
1667                .chain(split_mats.chunks_exact(2))
1668                .map(|pair| {
1669                    let plains = pair[0].as_ref().unwrap();
1670                    if let Some(rots) = pair[1].as_ref() {
1671                        std::iter::zip(plains.columns(), rots.columns())
1672                            .flat_map(|(claim, claim_rot)| {
1673                                assert_eq!(claim.len(), 1);
1674                                assert_eq!(claim_rot.len(), 1);
1675                                [claim[0], claim_rot[0]]
1676                            })
1677                            .collect_vec()
1678                    } else {
1679                        plains
1680                            .columns()
1681                            .map(|claim| {
1682                                assert_eq!(claim.len(), 1);
1683                                claim[0]
1684                            })
1685                            .collect_vec()
1686                    }
1687                })
1688                .collect_vec();
1689            column_openings.push(openings_of_air);
1690        }
1691        Ok(column_openings)
1692    }
1693}