1use 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;
72mod fractional;
74mod 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
91const DAG_FALLBACK_MONOMIAL_RATIO: usize = 2;
96const BATCH_MLE_DEFAULT_MEMORY_FLOOR: usize = 6 << 30; #[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 let n_max =
139 log2_strict_usize(proving_ctx.per_trace[0].1.common_main.height()).saturating_sub(l_skip);
140 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 let n_logup = calculate_n_logup(l_skip, total_interactions);
159 let interactions_layout = StackedLayout::new(0, l_skip + n_logup, interactions_meta).unwrap();
162
163 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 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 let round0_span = info_span!("prover.rap_constraints.round0", phase = "prover").entered();
247 let mut sumcheck_round_polys = Vec::with_capacity(n_max);
249 let mut r = Vec::with_capacity(n_max + 1);
250 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 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 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 (0..=s_0_deg)
294 .step_by(1 << l_skip)
295 .map(|j| unsafe {
296 *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 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 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 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 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 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 pub xi: Vec<EF>,
459 pub lambda_pows: Option<DeviceBuffer<EF>>,
460 lambda_combinations: Vec<Option<DeviceBuffer<EF>>>,
462 d_beta_pows: DeviceBuffer<EF>,
464 logup_combinations: Vec<Option<LogupCombinations>>,
466
467 eq_xis: FxHashMap<usize, EqEvalLayers<EF>>,
469 eq_3b_per_trace: Vec<Vec<EF>>,
470 d_eq_3b_per_trace: Vec<DeviceBuffer<EF>>,
471 sels_per_trace_base: Vec<DeviceMatrix<F>>,
473 mat_evals_per_trace: Vec<Vec<DeviceMatrix<EF>>>,
475 sels_per_trace: Vec<DeviceMatrix<EF>>,
476 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 pk: &'a DeviceMultiStarkProvingKey<GenericGpuBackend<HS>>,
486
487 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 gkr_mem_contribution: usize,
497 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 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, monomial_num_y_threshold,
625 device_ctx: device_ctx.clone(),
626 })
627 }
628
629 #[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 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 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 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); cols[0] = F::ONE; cols[2 * height + height - 1] = F::ONE; 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 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 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 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 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 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 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 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 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 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, );
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, );
935 }
936 }
937 self.mem
938 .emit_metrics_with_label("prover.batch_constraints.round0");
939 Ok(batch_sp_poly)
940 }
941
942 #[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 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 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 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 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 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 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 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 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 let mut _keepalive_interpolated: Vec<DeviceMatrix<EF>> = Vec::new();
1086
1087 let mut late_eval: Vec<TraceCtx> = Vec::new(); let mut early_eval: Vec<TraceCtx> = Vec::new(); 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 if round == n_lift + 1 {
1120 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 if has_constraints {
1159 let tilde_eval = &mut self.zerocheck_tilde_evals[trace_idx];
1160 *tilde_eval *= r_prev;
1161 }
1164 if has_interactions {
1165 for x in self.logup_tilde_evals[trace_idx].iter_mut() {
1166 *x *= r_prev;
1167 }
1168 }
1171 }
1172 } else {
1173 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 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 }
1320 }
1321
1322 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 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 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 num_monomials as usize >= DAG_FALLBACK_MONOMIAL_RATIO * rules_len
1353 });
1354
1355 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 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 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 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 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 let sp_head = UnivariatePoly::lagrange_interpolate(
1476 &(0..s_deg).map(F::from_usize).collect_vec(),
1477 &sp_head_evals,
1478 );
1479 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 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 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 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 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 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 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 let values = &mat_host.values;
1628 let orig: Vec<EF> = (0..air_width)
1629 .map(|col| values[col * height]) .collect();
1631 let rot: Option<Vec<EF>> = if need_rot {
1632 Some(
1633 (air_width..width)
1634 .map(|col| values[col * height]) .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 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}