1use std::{array::from_fn, convert::TryInto, env, ffi::c_void, mem::transmute};
2
3use openvm_cuda_common::{
4 copy::{cuda_memcpy_on, MemCopyD2H},
5 d_buffer::DeviceBuffer,
6 memory_manager::MemTracker,
7 stream::GpuDeviceCtx,
8};
9use openvm_stark_backend::{
10 poly_common::{eval_eq_mle, interpolate_linear_at_01, interpolate_quadratic_at_012},
11 proof::GkrLayerClaims,
12 prover::fractional_sumcheck_gkr::{Frac, FracSumcheckProof},
13 FiatShamirTranscript, StarkProtocolConfig,
14};
15use p3_field::{Field, PrimeCharacteristicRing};
16use p3_util::log2_strict_usize;
17use tracing::{debug_span, instrument};
18
19use super::errors::FractionalSumcheckError;
20use crate::{
21 cuda::{
22 logup_zerocheck::{
23 _frac_compute_round_temp_buffer_size, fold_ef_frac_columns,
24 fold_ef_frac_columns_inplace, frac_build_tree_layer, frac_build_tree_two_layers,
25 frac_compute_round, frac_compute_round_and_fold, frac_compute_round_and_fold_inplace,
26 frac_compute_round_and_revert, frac_multifold_raw, frac_precompute_m_build_raw,
27 frac_precompute_m_eval_round_raw,
28 },
29 ntt::{bit_rev_frac_ext, bit_rev_frac_ext_build_k2},
30 },
31 poly::SqrtEqLayers,
32 prelude::EF,
33};
34
35const GKR_S_DEG: usize = 3;
36const GKR_WINDOW_SIZE: usize = 3;
37const GKR_WINDOW_DEFAULT_MIN_N: usize = 22;
38
39#[derive(Debug, Clone, Copy, PartialEq, Eq)]
40pub struct FractionalInputSize {
41 pub real_len: usize,
42 pub logical_len: usize,
43}
44
45impl FractionalInputSize {
46 pub fn new(real_len: usize, logical_len: usize) -> Self {
47 debug_assert!(real_len <= logical_len);
48 debug_assert!(logical_len.is_power_of_two() || real_len == logical_len);
49 Self {
50 real_len,
51 logical_len,
52 }
53 }
54
55 pub fn dense(len: usize) -> Self {
56 Self {
57 real_len: len,
58 logical_len: len,
59 }
60 }
61
62 pub fn peak_work_buffer_bytes(&self) -> usize {
85 let s_frac = std::mem::size_of::<Frac<EF>>();
86 let fold_eval = (self.logical_len / 4) * s_frac;
87
88 let s_ef = std::mem::size_of::<EF>();
89 let w = GKR_WINDOW_SIZE;
90 let precompute_f =
91 (self.logical_len >> (1 + w)).max(1 << GKR_WINDOW_DEFAULT_MIN_N) * s_frac;
92 let precompute_ef = ((1 << (2 * w + 1)) + (1 << (w + 1))) * s_ef;
96
97 fold_eval.max(precompute_f + precompute_ef)
98 }
99}
100
101#[derive(Debug, Clone, Copy)]
103enum BufferTarget {
104 LayerToWork,
106 WorkToLayer,
108 InPlaceLayer,
110 InPlaceWork,
112}
113
114struct BufferScheduler {
119 data_in_work_buffer: bool,
121 work_buffer_cap: usize,
123}
124
125#[derive(Debug, Clone, Copy, PartialEq, Eq)]
126enum GkrRoundStrategy {
127 FoldEval,
128 PrecomputeM,
129}
130
131const PRECOMPUTE_M_TAIL_TILE: usize = 4096;
132const PRECOMPUTE_M_MIN_TAIL_TILE: usize = 256;
133const PRECOMPUTE_M_DEFAULT_MIN_BLOCKS: usize = 64;
134const PRECOMPUTE_M_DEFAULT_TARGET_BLOCKS: usize = 1024;
135
136impl BufferScheduler {
137 fn new(work_buffer_cap: usize) -> Self {
139 Self {
140 data_in_work_buffer: false,
141 work_buffer_cap,
142 }
143 }
144
145 fn can_pingpong(&self, post_fold_size: usize) -> bool {
147 post_fold_size <= self.work_buffer_cap
148 }
149
150 fn next_target(&mut self, post_fold_size: usize, last_outer_round: bool) -> BufferTarget {
155 let can_pingpong = self.can_pingpong(post_fold_size);
156
157 if last_outer_round {
158 if can_pingpong {
159 if self.data_in_work_buffer {
161 self.data_in_work_buffer = false;
162 BufferTarget::WorkToLayer
163 } else {
164 self.data_in_work_buffer = true;
165 BufferTarget::LayerToWork
166 }
167 } else {
168 debug_assert!(
170 !self.data_in_work_buffer,
171 "in-place path requires data in layer"
172 );
173 BufferTarget::InPlaceLayer
174 }
175 } else {
176 if self.data_in_work_buffer {
178 BufferTarget::InPlaceWork
180 } else {
181 self.data_in_work_buffer = true;
183 BufferTarget::LayerToWork
184 }
185 }
186 }
187
188 fn final_fold_target(&self, last_outer_round: bool) -> BufferTarget {
190 if last_outer_round {
191 if self.data_in_work_buffer {
192 BufferTarget::InPlaceWork
193 } else {
194 BufferTarget::InPlaceLayer
195 }
196 } else {
197 if self.data_in_work_buffer {
199 BufferTarget::InPlaceWork
200 } else {
201 BufferTarget::LayerToWork
202 }
203 }
204 }
205}
206
207fn precompute_m_enabled() -> bool {
208 !matches!(
209 env::var("SWIRL_CUDA_GKR_PRECOMPUTE_M"),
210 Ok(val) if matches!(val.as_str(), "0" | "false" | "FALSE" | "no" | "NO")
211 )
212}
213
214fn precompute_m_min_blocks_threshold() -> usize {
215 env::var("SWIRL_CUDA_GKR_PRECOMPUTE_M_MIN_BLOCKS")
216 .ok()
217 .and_then(|val| val.parse::<usize>().ok())
218 .unwrap_or(PRECOMPUTE_M_DEFAULT_MIN_BLOCKS)
219 .max(1)
220}
221
222fn precompute_m_num_tail_blocks(rem_n: usize, w: usize, tail_tile: usize) -> usize {
223 let tail_n = rem_n - w;
224 (1usize << tail_n).div_ceil(tail_tile)
225}
226
227fn precompute_m_target_blocks() -> usize {
228 env::var("SWIRL_CUDA_GKR_PRECOMPUTE_M_TARGET_BLOCKS")
229 .ok()
230 .and_then(|val| val.parse::<usize>().ok())
231 .unwrap_or(PRECOMPUTE_M_DEFAULT_TARGET_BLOCKS)
232 .max(1)
233}
234
235fn precompute_m_tail_tile_override() -> Option<usize> {
236 env::var("SWIRL_CUDA_GKR_PRECOMPUTE_M_TAIL_TILE")
237 .ok()
238 .and_then(|val| val.parse::<usize>().ok())
239 .map(|v| v.clamp(PRECOMPUTE_M_MIN_TAIL_TILE, PRECOMPUTE_M_TAIL_TILE))
240}
241
242fn precompute_m_min_n() -> usize {
243 env::var("SWIRL_CUDA_GKR_PRECOMPUTE_M_MIN_N")
244 .ok()
245 .and_then(|val| val.parse::<usize>().ok())
246 .unwrap_or(GKR_WINDOW_DEFAULT_MIN_N)
247}
248
249fn precompute_m_build_tail_tile(
250 rem_n: usize,
251 w: usize,
252 min_blocks_threshold: usize,
253 target_blocks: usize,
254 tail_tile_override: Option<usize>,
255) -> usize {
256 if let Some(tile) = tail_tile_override {
257 return tile;
258 }
259 let tail_n = rem_n - w;
260 let k = 1usize << tail_n;
261 let target_blocks = target_blocks.max(min_blocks_threshold).max(1);
262 let desired_tile = k.div_ceil(target_blocks).max(1);
263 desired_tile.clamp(PRECOMPUTE_M_MIN_TAIL_TILE, PRECOMPUTE_M_TAIL_TILE)
264}
265
266fn choose_precompute_m_window_w(
267 rem_n: usize,
268 rounds_left: usize,
269 min_blocks_threshold: usize,
270 target_blocks: usize,
271 tail_tile_override: Option<usize>,
272 min_n: usize,
273) -> Option<usize> {
274 if rem_n < min_n || rounds_left < GKR_WINDOW_SIZE {
275 return None;
276 }
277 let w = GKR_WINDOW_SIZE;
278 let tail_tile = precompute_m_build_tail_tile(
279 rem_n,
280 w,
281 min_blocks_threshold,
282 target_blocks,
283 tail_tile_override,
284 );
285 (precompute_m_num_tail_blocks(rem_n, w, tail_tile) >= min_blocks_threshold).then_some(w)
286}
287
288fn choose_round_strategy(
289 round: usize,
290 precompute_m_env: bool,
291 precompute_m_min_blocks_threshold: usize,
292 precompute_m_target_blocks: usize,
293 precompute_m_tail_tile_override: Option<usize>,
294 precompute_m_min_n: usize,
295) -> GkrRoundStrategy {
296 if !precompute_m_env {
297 return GkrRoundStrategy::FoldEval;
298 }
299 let start_base = 1usize;
300 let stop = round.div_ceil(2);
301 let rem_n = round - start_base;
302 let rounds_left = stop - start_base;
303 if choose_precompute_m_window_w(
304 rem_n,
305 rounds_left,
306 precompute_m_min_blocks_threshold,
307 precompute_m_target_blocks,
308 precompute_m_tail_tile_override,
309 precompute_m_min_n,
310 )
311 .is_none()
312 {
313 return GkrRoundStrategy::FoldEval;
314 }
315
316 GkrRoundStrategy::PrecomputeM
317}
318
319fn eval_mle_table(points: &[EF], out: &mut [EF]) {
320 let n = points.len();
322 let size = 1usize << n;
323 debug_assert!(out.len() >= size);
324 for (bits, dst) in out.iter_mut().enumerate().take(size) {
325 let mut acc = EF::ONE;
326 for (i, &x) in points.iter().enumerate() {
327 let bit = ((bits >> (n - 1 - i)) & 1) == 1;
328 acc *= if bit { x } else { EF::ONE - x };
329 }
330 *dst = acc;
331 }
332}
333
334fn eq_tail_ptrs(
337 eq_buffer: &SqrtEqLayers,
338 drop_count: usize,
339) -> (*const EF, *const EF, usize, usize) {
340 let mut high_n = eq_buffer.high_n();
341 let mut low_n = eq_buffer.low_n();
342 let total_n = high_n + low_n;
343 if drop_count >= total_n {
344 return (std::ptr::null(), std::ptr::null(), 1, 0);
345 }
346 if drop_count <= high_n {
347 high_n -= drop_count;
348 } else {
349 low_n -= drop_count - high_n;
350 high_n = 0;
351 }
352 let tail_n = low_n + high_n;
353 if tail_n == 0 {
354 (std::ptr::null(), std::ptr::null(), 1, 0)
355 } else {
356 (
357 eq_buffer.low.get_ptr(low_n),
358 eq_buffer.high.get_ptr(high_n),
359 1 << low_n,
360 tail_n,
361 )
362 }
363}
364
365fn copy_to_device_ptr<T: Copy>(
366 dst: *mut T,
367 src: &[T],
368 device_ctx: &GpuDeviceCtx,
369) -> Result<(), FractionalSumcheckError> {
370 if src.is_empty() {
371 return Ok(());
372 }
373 unsafe {
374 cuda_memcpy_on::<false, true>(
375 dst as *mut c_void,
376 src.as_ptr() as *const c_void,
377 std::mem::size_of_val(src),
378 device_ctx,
379 )?;
380 }
381 Ok(())
382}
383
384fn bit_reverse_usize(value: usize, bits: usize) -> usize {
385 if bits == 0 {
386 0
387 } else {
388 value.reverse_bits() >> (usize::BITS as usize - bits)
389 }
390}
391
392fn virtual_padding_q(alpha: EF, subtree_len: usize) -> EF {
393 let mut q = alpha;
394 let mut len = subtree_len;
395 while len > 1 {
396 q *= q;
397 len >>= 1;
398 }
399 q
400}
401
402fn folded_virtual_support_len(source_support_len: usize) -> usize {
403 if source_support_len == 0 {
404 return 0;
405 }
406 let last = source_support_len - 1;
410 2 * (last / 4) + if last.is_multiple_of(4) { 1 } else { 2 }
411}
412
413#[allow(clippy::too_many_arguments)]
414fn copy_compact_node_from_device(
415 layer: &DeviceBuffer<Frac<EF>>,
416 dense_idx: usize,
417 active_size: usize,
418 real_len: usize,
419 logical_len: usize,
420 alpha: EF,
421 copy_scratch: &mut DeviceBuffer<Frac<EF>>,
422 device_ctx: &GpuDeviceCtx,
423) -> Result<Frac<EF>, FractionalSumcheckError> {
424 let subtree_len = logical_len / active_size;
425 let start = bit_reverse_usize(dense_idx, log2_strict_usize(active_size)) * subtree_len;
426 if start >= real_len {
427 return Ok(Frac {
428 p: EF::ZERO,
429 q: virtual_padding_q(alpha, subtree_len),
430 });
431 }
432 let physical_idx = if dense_idx < real_len {
433 dense_idx
434 } else {
435 start
436 };
437 copy_from_device(layer, physical_idx, copy_scratch, device_ctx)
438}
439
440#[allow(clippy::too_many_arguments)]
442fn observe_and_update<SC, TS>(
443 d_sum_evals: &DeviceBuffer<EF>,
444 transcript: &mut TS,
445 round_polys_eval: &mut Vec<[EF; GKR_S_DEG]>,
446 r_vec: &mut Vec<EF>,
447 prev_s_eval: &mut EF,
448 xi_j: EF,
449 eq_r_acc: &mut EF,
450 device_ctx: &GpuDeviceCtx,
451) -> Result<EF, FractionalSumcheckError>
452where
453 SC: StarkProtocolConfig<EF = EF>,
454 TS: FiatShamirTranscript<SC>,
455{
456 let (s_evals, sp_evals) =
457 reconstruct_s_evals(d_sum_evals, *prev_s_eval, xi_j, *eq_r_acc, device_ctx)?;
458
459 for &eval in &s_evals {
460 transcript.observe_ext(eval);
461 }
462 round_polys_eval.push(s_evals);
463
464 let r = transcript.sample_ext();
465 r_vec.push(r);
466
467 let eq_r = eval_eq_mle(&[xi_j], &[r]);
468 *eq_r_acc *= eq_r;
469 *prev_s_eval = eq_r * interpolate_quadratic_at_012(&sp_evals, r);
470
471 Ok(r)
472}
473
474#[allow(clippy::too_many_arguments)]
478fn do_sumcheck_round_and_revert<SC, TS>(
479 eq_buffer: &mut SqrtEqLayers,
480 layer: &mut DeviceBuffer<Frac<EF>>,
481 pq_size: usize,
482 total_leaves: usize,
483 lambda: EF,
484 alpha: EF,
485 transcript: &mut TS,
486 d_sum_evals: &mut DeviceBuffer<EF>,
487 tmp_block_sums: &mut DeviceBuffer<EF>,
488 round_polys_eval: &mut Vec<[EF; GKR_S_DEG]>,
489 r_vec: &mut Vec<EF>,
490 prev_s_eval: &mut EF,
491 xi_j: EF,
492 eq_r_acc: &mut EF,
493 device_ctx: &GpuDeviceCtx,
494) -> Result<EF, FractionalSumcheckError>
495where
496 SC: StarkProtocolConfig<EF = EF>,
497 TS: FiatShamirTranscript<SC>,
498{
499 let stream = device_ctx.stream.as_raw();
500 unsafe {
501 frac_compute_round_and_revert(
502 eq_buffer,
503 layer,
504 pq_size / 2,
505 total_leaves,
506 lambda,
507 alpha,
508 d_sum_evals,
509 tmp_block_sums,
510 stream,
511 )
512 .map_err(FractionalSumcheckError::ComputeRound)?;
513 }
514 eq_buffer.drop_layer();
515 observe_and_update(
516 d_sum_evals,
517 transcript,
518 round_polys_eval,
519 r_vec,
520 prev_s_eval,
521 xi_j,
522 eq_r_acc,
523 device_ctx,
524 )
525}
526
527#[allow(clippy::too_many_arguments)]
532fn do_fused_sumcheck_round<SC, TS>(
533 eq_buffer: &mut SqrtEqLayers,
534 src_pq_buffer: &DeviceBuffer<Frac<EF>>,
535 dst_pq_buffer: &mut DeviceBuffer<Frac<EF>>,
536 src_pq_size: usize,
537 src_real_len: usize,
538 total_leaves: usize,
539 lambda: EF,
540 r_prev: EF,
541 alpha: EF,
542 transcript: &mut TS,
543 d_sum_evals: &mut DeviceBuffer<EF>,
544 tmp_block_sums: &mut DeviceBuffer<EF>,
545 round_polys_eval: &mut Vec<[EF; GKR_S_DEG]>,
546 r_vec: &mut Vec<EF>,
547 prev_s_eval: &mut EF,
548 xi_j: EF,
549 eq_r_acc: &mut EF,
550 device_ctx: &GpuDeviceCtx,
551) -> Result<EF, FractionalSumcheckError>
552where
553 SC: StarkProtocolConfig<EF = EF>,
554 TS: FiatShamirTranscript<SC>,
555{
556 let stream = device_ctx.stream.as_raw();
557 unsafe {
558 frac_compute_round_and_fold(
559 eq_buffer,
560 src_pq_buffer,
561 dst_pq_buffer,
562 src_pq_size,
563 src_real_len,
564 total_leaves,
565 lambda,
566 r_prev,
567 alpha,
568 d_sum_evals,
569 tmp_block_sums,
570 stream,
571 )
572 .map_err(FractionalSumcheckError::ComputeRound)?;
573 }
574 eq_buffer.drop_layer();
575 observe_and_update(
576 d_sum_evals,
577 transcript,
578 round_polys_eval,
579 r_vec,
580 prev_s_eval,
581 xi_j,
582 eq_r_acc,
583 device_ctx,
584 )
585}
586
587#[allow(clippy::too_many_arguments)]
589fn do_fused_sumcheck_round_inplace<SC, TS>(
590 eq_buffer: &mut SqrtEqLayers,
591 pq_buffer: &mut DeviceBuffer<Frac<EF>>,
592 src_pq_size: usize,
593 src_real_len: usize,
594 src_logical_len: usize,
595 dst_real_len: usize,
596 dst_logical_len: usize,
597 lambda: EF,
598 r_prev: EF,
599 alpha: EF,
600 transcript: &mut TS,
601 d_sum_evals: &mut DeviceBuffer<EF>,
602 tmp_block_sums: &mut DeviceBuffer<EF>,
603 round_polys_eval: &mut Vec<[EF; GKR_S_DEG]>,
604 r_vec: &mut Vec<EF>,
605 prev_s_eval: &mut EF,
606 xi_j: EF,
607 eq_r_acc: &mut EF,
608 device_ctx: &GpuDeviceCtx,
609) -> Result<EF, FractionalSumcheckError>
610where
611 SC: StarkProtocolConfig<EF = EF>,
612 TS: FiatShamirTranscript<SC>,
613{
614 let stream = device_ctx.stream.as_raw();
615 unsafe {
616 frac_compute_round_and_fold_inplace(
617 eq_buffer,
618 pq_buffer,
619 src_pq_size,
620 src_real_len,
621 src_logical_len,
622 dst_real_len,
623 dst_logical_len,
624 lambda,
625 r_prev,
626 alpha,
627 d_sum_evals,
628 tmp_block_sums,
629 stream,
630 )
631 .map_err(FractionalSumcheckError::ComputeRound)?;
632 }
633 eq_buffer.drop_layer();
634 observe_and_update(
635 d_sum_evals,
636 transcript,
637 round_polys_eval,
638 r_vec,
639 prev_s_eval,
640 xi_j,
641 eq_r_acc,
642 device_ctx,
643 )
644}
645
646#[instrument(skip_all)]
649pub fn fractional_sumcheck_gpu<SC, TS>(
650 transcript: &mut TS,
651 leaves: DeviceBuffer<Frac<EF>>,
652 sizes: FractionalInputSize,
653 alpha: EF,
654 assert_zero: bool,
655 mem: &mut MemTracker,
656 device_ctx: &GpuDeviceCtx,
657) -> Result<(FracSumcheckProof<SC>, Vec<EF>), FractionalSumcheckError>
658where
659 SC: StarkProtocolConfig<EF = EF>,
660 TS: FiatShamirTranscript<SC>,
661{
662 let mut layer = leaves;
663 if layer.is_empty() {
664 return Ok((
665 FracSumcheckProof {
666 fractional_sum: (EF::ZERO, EF::ONE),
667 claims_per_layer: vec![],
668 sumcheck_polys: vec![],
669 },
670 vec![],
671 ));
672 };
673 let stream = device_ctx.stream.as_raw();
674 let real_len = sizes.real_len;
675 let total_leaves = sizes.logical_len;
676 assert_eq!(
677 layer.len(),
678 real_len,
679 "fractional_sumcheck_gpu input length must equal real_len"
680 );
681 assert!(real_len > 0, "real_len must be nonzero");
682 assert!(
683 total_leaves.is_power_of_two(),
684 "logical_len must be a power of two"
685 );
686 assert!(
687 real_len <= total_leaves,
688 "real_len must not exceed logical_len"
689 );
690 assert!(
691 total_leaves / 2 <= real_len,
692 "virtual padding requires logical_len / 2 <= real_len"
693 );
694 let total_rounds = log2_strict_usize(total_leaves);
696 assert!(total_rounds > 0, "n_logup > 0 when there are interactions");
697 let virtual_input = real_len < total_leaves;
706 let start_layer_i = if total_leaves > 1024 {
707 unsafe {
708 let buf = transmute::<&DeviceBuffer<Frac<EF>>, &DeviceBuffer<(EF, EF)>>(&layer);
710 bit_rev_frac_ext_build_k2(buf, real_len, total_rounds as u32, alpha, stream)
711 .map_err(FractionalSumcheckError::BitReversal)?;
712 }
713 2 } else {
715 unsafe {
717 if !virtual_input {
718 let buf = transmute::<&DeviceBuffer<Frac<EF>>, &DeviceBuffer<(EF, EF)>>(&layer);
719 bit_rev_frac_ext(
720 buf,
721 buf,
722 total_rounds as u32,
723 total_leaves.try_into().unwrap(),
724 1,
725 stream,
726 )
727 .map_err(FractionalSumcheckError::BitReversal)?;
728 }
729 frac_build_tree_layer(
730 &mut layer,
731 total_leaves,
732 total_leaves,
733 false,
734 alpha,
735 true,
736 stream,
737 )
738 .map_err(FractionalSumcheckError::SegmentTree)?;
739 use crate::cuda::logup_zerocheck::frac_add_alpha;
740 if !virtual_input {
741 let half = total_leaves / 2;
742 let second_half_ptr = layer.as_mut_raw_ptr() as *mut Frac<EF>;
743 let second_half_buf =
744 DeviceBuffer::<Frac<EF>>::from_raw_parts(second_half_ptr.add(half), half);
745 frac_add_alpha(&second_half_buf, alpha, stream)
746 .map_err(FractionalSumcheckError::SegmentTree)?;
747 std::mem::forget(second_half_buf);
748 }
749 }
750 1 };
752
753 let mut i = start_layer_i;
755 while i + 1 < total_rounds {
756 let half_i1 = total_leaves >> (i + 2);
757 unsafe {
758 frac_build_tree_two_layers(&mut layer, half_i1, total_leaves, alpha, stream)
759 .map_err(FractionalSumcheckError::SegmentTree)?;
760 }
761 i += 2;
762 }
763 if i < total_rounds {
765 unsafe {
766 frac_build_tree_layer(
767 &mut layer,
768 total_leaves >> i,
769 total_leaves,
770 false,
771 alpha,
772 false,
773 stream,
774 )
775 .map_err(FractionalSumcheckError::SegmentTree)?;
776 }
777 }
778 mem.emit_metrics_with_label("frac_sumcheck.segment_tree");
779 mem.tracing_info("fractional_sumcheck_gkr: after building segment tree");
780 let mut copy_scratch = DeviceBuffer::<Frac<EF>>::with_capacity_on(1, device_ctx);
781 let root = copy_from_device(&layer, 0, &mut copy_scratch, device_ctx)?;
782 unsafe {
783 frac_build_tree_layer(&mut layer, 2, total_leaves, true, alpha, false, stream)
784 .map_err(FractionalSumcheckError::SegmentTree)?;
785 }
786 if assert_zero {
787 if root.p != EF::ZERO {
788 return Err(FractionalSumcheckError::NonzeroRootSum {
789 p: root.p,
790 q: root.q,
791 });
792 }
793 } else {
794 transcript.observe_ext(root.p);
795 }
796 transcript.observe_ext(root.q);
797
798 let mut claims_per_layer = Vec::with_capacity(total_rounds);
799 let mut sumcheck_polys = Vec::with_capacity(total_rounds);
800
801 let first_left = copy_compact_node_from_device(
802 &layer,
803 0,
804 2,
805 real_len,
806 total_leaves,
807 alpha,
808 &mut copy_scratch,
809 device_ctx,
810 )?;
811 let first_right = copy_compact_node_from_device(
812 &layer,
813 1,
814 2,
815 real_len,
816 total_leaves,
817 alpha,
818 &mut copy_scratch,
819 device_ctx,
820 )?;
821 claims_per_layer.push(GkrLayerClaims {
822 p_xi_0: first_left.p,
823 q_xi_0: first_left.q,
824 p_xi_1: first_right.p,
825 q_xi_1: first_right.q,
826 });
827 for value in [
828 claims_per_layer[0].p_xi_0,
829 claims_per_layer[0].q_xi_0,
830 claims_per_layer[0].p_xi_1,
831 claims_per_layer[0].q_xi_1,
832 ] {
833 transcript.observe_ext(value);
834 }
835 let mu_1 = transcript.sample_ext();
836 let mut xi_prev = vec![mu_1];
837 let mut d_sum_evals = DeviceBuffer::<EF>::with_capacity_on(2, device_ctx);
838
839 let precompute_m_env = precompute_m_enabled();
840
841 let max_work_size = if total_rounds > 2 {
848 if precompute_m_env {
849 (total_leaves >> (1 + GKR_WINDOW_SIZE)).max(1 << GKR_WINDOW_DEFAULT_MIN_N)
850 } else {
851 total_leaves >> 2
852 }
853 } else {
854 0
855 };
856 let mut work_buffer = if max_work_size > 0 {
857 DeviceBuffer::<Frac<EF>>::with_capacity_on(max_work_size, device_ctx)
858 } else {
859 DeviceBuffer::new()
860 };
861 let max_tmp_buffer_capacity = if total_rounds > 1 {
862 (unsafe { _frac_compute_round_temp_buffer_size((1 << (total_rounds - 1)) as u32) }) as usize
863 } else {
864 0
865 };
866 let mut tmp_block_sums = if max_tmp_buffer_capacity > 0 {
867 DeviceBuffer::<EF>::with_capacity_on(max_tmp_buffer_capacity, device_ctx)
868 } else {
869 DeviceBuffer::new()
870 };
871 let mut final_fold_buffer = DeviceBuffer::<Frac<EF>>::new();
872 let precompute_m_min_blocks_threshold = precompute_m_min_blocks_threshold();
873 let precompute_m_target_blocks = precompute_m_target_blocks();
874 let precompute_m_tail_tile_override = precompute_m_tail_tile_override();
875 let precompute_m_min_n = precompute_m_min_n();
876 let mut m_buffer = DeviceBuffer::<EF>::new();
877 let mut m_partial_buffer = DeviceBuffer::<EF>::new();
878 let mut eq_r_prefix_buffer = DeviceBuffer::<EF>::new();
879 let mut eq_suffix_buffer = DeviceBuffer::<EF>::new();
880 for round in 1..total_rounds {
883 let gkr_round_span = debug_span!("GKR", round).entered();
884
885 debug_assert_eq!(xi_prev.len(), round);
888 let mut eq_buffer = SqrtEqLayers::from_xi(&xi_prev[1..], device_ctx)
891 .map_err(FractionalSumcheckError::EvalEqHypercube)?;
892
893 let mut round_polys_eval = Vec::with_capacity(round);
894 let mut r_vec = Vec::with_capacity(round);
895 let mut pq_size = 2 << round;
896
897 let lambda = transcript.sample_ext();
898
899 let tmp_buffer_capacity =
900 unsafe { _frac_compute_round_temp_buffer_size((1 << round) as u32) } as usize;
901 if tmp_buffer_capacity > tmp_block_sums.len() {
902 tmp_block_sums = DeviceBuffer::<EF>::with_capacity_on(tmp_buffer_capacity, device_ctx);
903 }
904
905 let last_outer_round = round == total_rounds - 1;
906 debug_assert!(round > 0);
907 let backend = choose_round_strategy(
908 round,
909 precompute_m_env,
910 precompute_m_min_blocks_threshold,
911 precompute_m_target_blocks,
912 precompute_m_tail_tile_override,
913 precompute_m_min_n,
914 );
915
916 let (numer_claim, denom_claim) =
918 reduce_to_single_evaluation(claims_per_layer.last().unwrap(), xi_prev[0]);
919 let mut prev_s_eval = numer_claim + lambda * denom_claim;
920 let mut eq_r_acc = EF::ONE;
921
922 let r0 = do_sumcheck_round_and_revert(
926 &mut eq_buffer,
927 &mut layer,
928 pq_size,
929 total_leaves,
930 lambda,
931 alpha,
932 transcript,
933 &mut d_sum_evals,
934 &mut tmp_block_sums,
935 &mut round_polys_eval,
936 &mut r_vec,
937 &mut prev_s_eval,
938 xi_prev[0],
939 &mut eq_r_acc,
940 device_ctx,
941 )?;
942
943 let mut prev_r = r0;
945 let active: &mut DeviceBuffer<Frac<EF>>;
946
947 match backend {
948 GkrRoundStrategy::FoldEval => {
949 let mut scheduler = BufferScheduler::new(max_work_size);
951 let mut source_real_len = real_len;
952 let mut source_logical_len = total_leaves;
953 for &xi_j in xi_prev.iter().skip(1) {
954 let src_pq_size = pq_size;
955 let post_fold_size = pq_size >> 1;
956 let target = scheduler.next_target(post_fold_size, last_outer_round);
957 let source_support_len = if source_logical_len == total_leaves && virtual_input
958 {
959 let subtree_len = total_leaves / src_pq_size;
960 real_len.div_ceil(subtree_len)
961 } else {
962 source_real_len
963 };
964 let compact_inplace_layer = matches!(target, BufferTarget::InPlaceLayer)
965 && source_logical_len == total_leaves
966 && virtual_input;
967 let dst_real_len = if compact_inplace_layer {
968 folded_virtual_support_len(source_support_len)
969 } else {
970 post_fold_size
971 };
972 let dst_logical_len = post_fold_size;
973
974 let r = match target {
975 BufferTarget::LayerToWork => do_fused_sumcheck_round(
976 &mut eq_buffer,
977 &layer,
978 &mut work_buffer,
979 src_pq_size,
980 source_real_len,
981 source_logical_len,
982 lambda,
983 prev_r,
984 alpha,
985 transcript,
986 &mut d_sum_evals,
987 &mut tmp_block_sums,
988 &mut round_polys_eval,
989 &mut r_vec,
990 &mut prev_s_eval,
991 xi_j,
992 &mut eq_r_acc,
993 device_ctx,
994 )?,
995 BufferTarget::WorkToLayer => do_fused_sumcheck_round(
996 &mut eq_buffer,
997 &work_buffer,
998 &mut layer,
999 src_pq_size,
1000 source_real_len,
1001 source_logical_len,
1002 lambda,
1003 prev_r,
1004 alpha,
1005 transcript,
1006 &mut d_sum_evals,
1007 &mut tmp_block_sums,
1008 &mut round_polys_eval,
1009 &mut r_vec,
1010 &mut prev_s_eval,
1011 xi_j,
1012 &mut eq_r_acc,
1013 device_ctx,
1014 )?,
1015 BufferTarget::InPlaceLayer => do_fused_sumcheck_round_inplace(
1016 &mut eq_buffer,
1017 &mut layer,
1018 src_pq_size,
1019 source_real_len,
1020 source_logical_len,
1021 dst_real_len,
1022 dst_logical_len,
1023 lambda,
1024 prev_r,
1025 alpha,
1026 transcript,
1027 &mut d_sum_evals,
1028 &mut tmp_block_sums,
1029 &mut round_polys_eval,
1030 &mut r_vec,
1031 &mut prev_s_eval,
1032 xi_j,
1033 &mut eq_r_acc,
1034 device_ctx,
1035 )?,
1036 BufferTarget::InPlaceWork => do_fused_sumcheck_round_inplace(
1037 &mut eq_buffer,
1038 &mut work_buffer,
1039 src_pq_size,
1040 source_real_len,
1041 source_logical_len,
1042 dst_real_len,
1043 dst_logical_len,
1044 lambda,
1045 prev_r,
1046 alpha,
1047 transcript,
1048 &mut d_sum_evals,
1049 &mut tmp_block_sums,
1050 &mut round_polys_eval,
1051 &mut r_vec,
1052 &mut prev_s_eval,
1053 xi_j,
1054 &mut eq_r_acc,
1055 device_ctx,
1056 )?,
1057 };
1058
1059 pq_size >>= 1;
1060 prev_r = r;
1061 source_real_len = dst_real_len;
1062 source_logical_len = dst_logical_len;
1063 }
1064
1065 let compact_virtual_final_fold = source_real_len < source_logical_len;
1067 active = match scheduler.final_fold_target(last_outer_round) {
1068 BufferTarget::InPlaceWork if compact_virtual_final_fold => {
1069 let output_len = pq_size / 2;
1070 if final_fold_buffer.len() < output_len {
1071 final_fold_buffer =
1072 DeviceBuffer::<Frac<EF>>::with_capacity_on(output_len, device_ctx);
1073 }
1074 unsafe {
1075 fold_ef_frac_columns(
1076 &work_buffer,
1077 &mut final_fold_buffer,
1078 pq_size,
1079 source_real_len,
1080 source_logical_len,
1081 prev_r,
1082 alpha,
1083 stream,
1084 )
1085 .map_err(FractionalSumcheckError::FoldColumns)?;
1086 }
1087 &mut final_fold_buffer
1088 }
1089 BufferTarget::InPlaceLayer if compact_virtual_final_fold => {
1090 let output_len = pq_size / 2;
1091 if final_fold_buffer.len() < output_len {
1092 final_fold_buffer =
1093 DeviceBuffer::<Frac<EF>>::with_capacity_on(output_len, device_ctx);
1094 }
1095 unsafe {
1096 fold_ef_frac_columns(
1097 &layer,
1098 &mut final_fold_buffer,
1099 pq_size,
1100 source_real_len,
1101 source_logical_len,
1102 prev_r,
1103 alpha,
1104 stream,
1105 )
1106 .map_err(FractionalSumcheckError::FoldColumns)?;
1107 }
1108 &mut final_fold_buffer
1109 }
1110 BufferTarget::InPlaceWork => {
1111 unsafe {
1112 fold_ef_frac_columns_inplace(
1113 &mut work_buffer,
1114 pq_size,
1115 source_real_len,
1116 source_logical_len,
1117 prev_r,
1118 alpha,
1119 stream,
1120 )
1121 .map_err(FractionalSumcheckError::FoldColumns)?;
1122 }
1123 &mut work_buffer
1124 }
1125 BufferTarget::InPlaceLayer => {
1126 unsafe {
1127 fold_ef_frac_columns_inplace(
1128 &mut layer,
1129 pq_size,
1130 source_real_len,
1131 source_logical_len,
1132 prev_r,
1133 alpha,
1134 stream,
1135 )
1136 .map_err(FractionalSumcheckError::FoldColumns)?;
1137 }
1138 &mut layer
1139 }
1140 BufferTarget::LayerToWork => {
1141 unsafe {
1142 fold_ef_frac_columns(
1143 &layer,
1144 &mut work_buffer,
1145 pq_size,
1146 source_real_len,
1147 source_logical_len,
1148 prev_r,
1149 alpha,
1150 stream,
1151 )
1152 .map_err(FractionalSumcheckError::FoldColumns)?;
1153 }
1154 &mut work_buffer
1155 }
1156 BufferTarget::WorkToLayer => unreachable!(),
1157 };
1158 pq_size >>= 1;
1159 }
1160 GkrRoundStrategy::PrecomputeM => {
1161 let base = 1usize;
1162 let stop = round.div_ceil(2);
1163
1164 let mut pending_fold = true;
1170 let layer_read_ptr = layer.as_ptr();
1171 let active_pq = if last_outer_round && !virtual_input {
1176 &mut layer
1177 } else {
1178 &mut work_buffer
1179 };
1180 let mut active_real_len = 0usize;
1181 let mut active_logical_len = 0usize;
1182
1183 let mut eq_r_window_host = vec![EF::ZERO; 1 << (GKR_WINDOW_SIZE + 1)];
1186 let mut eq_r_prefix_host = vec![EF::ZERO; 1 << GKR_WINDOW_SIZE];
1187 let mut eq_suffix_host = vec![EF::ZERO; 1 << GKR_WINDOW_SIZE];
1188
1189 if eq_r_prefix_buffer.is_empty() {
1190 eq_r_prefix_buffer =
1191 DeviceBuffer::<EF>::with_capacity_on(1usize << GKR_WINDOW_SIZE, device_ctx);
1192 }
1193 if eq_suffix_buffer.is_empty() {
1194 eq_suffix_buffer =
1195 DeviceBuffer::<EF>::with_capacity_on(1usize << GKR_WINDOW_SIZE, device_ctx);
1196 }
1197
1198 let mut base = base;
1199 while base < stop {
1200 let rem_n = round - base;
1201 let rounds_left = stop - base;
1202 let Some(w) = choose_precompute_m_window_w(
1203 rem_n,
1204 rounds_left,
1205 precompute_m_min_blocks_threshold,
1206 precompute_m_target_blocks,
1207 precompute_m_tail_tile_override,
1208 precompute_m_min_n,
1209 ) else {
1210 break;
1211 };
1212 if m_buffer.is_empty() {
1213 let max_m_len = 1usize << (2 * GKR_WINDOW_SIZE);
1214 m_buffer = DeviceBuffer::<EF>::with_capacity_on(max_m_len, device_ctx);
1215 }
1216 let m_ptr = m_buffer.as_mut_ptr();
1217 let max_eq_r_window_len = 1usize << (GKR_WINDOW_SIZE + 1);
1221 debug_assert!(
1222 tmp_block_sums.len() >= max_eq_r_window_len,
1223 "tmp_block_sums too small for eq_r_window: {} < {}",
1224 tmp_block_sums.len(),
1225 max_eq_r_window_len,
1226 );
1227 let d_eq_r_window = tmp_block_sums.as_mut_ptr();
1228
1229 let tail_tile = precompute_m_build_tail_tile(
1233 rem_n,
1234 w,
1235 precompute_m_min_blocks_threshold,
1236 precompute_m_target_blocks,
1237 precompute_m_tail_tile_override,
1238 );
1239 let num_blocks = precompute_m_num_tail_blocks(rem_n, w, tail_tile);
1240 let m_len = (1usize << w) * (1usize << w);
1241 let partial_len = num_blocks * m_len;
1242 if partial_len > m_partial_buffer.len() {
1243 m_partial_buffer =
1244 DeviceBuffer::<EF>::with_capacity_on(partial_len, device_ctx);
1245 }
1246
1247 let (eq_tail_low, eq_tail_high, eq_low_cap, _) =
1248 eq_tail_ptrs(&eq_buffer, w - 1);
1249
1250 let r_fold = prev_r; let build_src = if pending_fold {
1252 layer_read_ptr
1253 } else {
1254 active_pq.as_ptr()
1255 };
1256 let build_real_len = if pending_fold {
1257 real_len
1258 } else {
1259 active_real_len
1260 };
1261 let build_logical_len = if pending_fold {
1262 total_leaves
1263 } else {
1264 active_logical_len
1265 };
1266 unsafe {
1267 frac_precompute_m_build_raw(
1268 build_src,
1269 build_real_len,
1270 build_logical_len,
1271 rem_n,
1272 w,
1273 lambda,
1274 r_fold,
1275 alpha,
1276 pending_fold, eq_tail_low,
1278 eq_tail_high,
1279 eq_low_cap,
1280 tail_tile,
1281 m_partial_buffer.as_mut_ptr(),
1282 partial_len,
1283 m_ptr,
1284 stream,
1285 )
1286 .map_err(FractionalSumcheckError::ComputeRound)?;
1287 }
1288
1289 let mut window_rs = Vec::with_capacity(w);
1290 for t in 0..w {
1291 let prefix_bits = t;
1292 let suffix_bits = w - t - 1;
1293 eval_mle_table(&window_rs, &mut eq_r_prefix_host);
1294 eval_mle_table(&xi_prev[base + t + 1..base + w], &mut eq_suffix_host);
1295
1296 copy_to_device_ptr(
1297 eq_r_prefix_buffer.as_mut_ptr(),
1298 &eq_r_prefix_host[..(1usize << prefix_bits)],
1299 device_ctx,
1300 )?;
1301 copy_to_device_ptr(
1302 eq_suffix_buffer.as_mut_ptr(),
1303 &eq_suffix_host[..(1usize << suffix_bits)],
1304 device_ctx,
1305 )?;
1306 unsafe {
1307 frac_precompute_m_eval_round_raw(
1308 m_ptr,
1309 w,
1310 t,
1311 eq_r_prefix_buffer.as_ptr(),
1312 eq_suffix_buffer.as_ptr(),
1313 d_sum_evals.as_mut_ptr(),
1314 stream,
1315 )
1316 .map_err(FractionalSumcheckError::ComputeRound)?;
1317 }
1318 eq_buffer.drop_layer();
1319 let r = observe_and_update(
1320 &d_sum_evals,
1321 transcript,
1322 &mut round_polys_eval,
1323 &mut r_vec,
1324 &mut prev_s_eval,
1325 xi_prev[base + t],
1326 &mut eq_r_acc,
1327 device_ctx,
1328 )?;
1329 prev_r = r;
1330 window_rs.push(r);
1331 }
1332
1333 let (buf_vars, w_fold) = if pending_fold {
1335 let mut all_rs = Vec::with_capacity(w + 1);
1336 all_rs.push(r_fold);
1337 all_rs.extend_from_slice(&window_rs);
1338 eval_mle_table(&all_rs, &mut eq_r_window_host);
1339 copy_to_device_ptr(
1340 d_eq_r_window,
1341 &eq_r_window_host[..(1 << (w + 1))],
1342 device_ctx,
1343 )?;
1344 (rem_n + 1, w + 1)
1345 } else {
1346 eval_mle_table(&window_rs, &mut eq_r_window_host);
1347 copy_to_device_ptr(
1348 d_eq_r_window,
1349 &eq_r_window_host[..(1 << w)],
1350 device_ctx,
1351 )?;
1352 (rem_n, w)
1353 };
1354
1355 let multifold_src = if pending_fold {
1356 layer_read_ptr
1357 } else {
1358 active_pq.as_ptr()
1359 };
1360 let multifold_real_len = if pending_fold {
1361 real_len
1362 } else {
1363 active_real_len
1364 };
1365 let multifold_logical_len = if pending_fold {
1366 total_leaves
1367 } else {
1368 active_logical_len
1369 };
1370 debug_assert!(
1371 active_pq.len() >= (pq_size >> w_fold),
1372 "active_pq too small for multifold output: {} < {}",
1373 active_pq.len(),
1374 pq_size >> w_fold
1375 );
1376 unsafe {
1377 frac_multifold_raw(
1378 multifold_src,
1379 active_pq.as_mut_ptr(),
1380 multifold_real_len,
1381 multifold_logical_len,
1382 buf_vars,
1383 w_fold,
1384 alpha,
1385 d_eq_r_window,
1386 stream,
1387 )
1388 .map_err(FractionalSumcheckError::FoldColumns)?;
1389 }
1390 pq_size >>= w_fold;
1391 active_real_len = pq_size;
1392 active_logical_len = pq_size;
1393 pending_fold = false;
1394 base += w;
1395 }
1396
1397 if base < round {
1398 unsafe {
1401 frac_compute_round(
1402 &eq_buffer,
1403 active_pq,
1404 pq_size / 2,
1405 lambda,
1406 &mut d_sum_evals,
1407 &mut tmp_block_sums,
1408 stream,
1409 )
1410 .map_err(FractionalSumcheckError::ComputeRound)?;
1411 }
1412 eq_buffer.drop_layer();
1413 prev_r = observe_and_update(
1414 &d_sum_evals,
1415 transcript,
1416 &mut round_polys_eval,
1417 &mut r_vec,
1418 &mut prev_s_eval,
1419 xi_prev[base],
1420 &mut eq_r_acc,
1421 device_ctx,
1422 )?;
1423
1424 for &xi_j in xi_prev.iter().skip(base + 1) {
1425 let src_pq_size = pq_size;
1426 prev_r = do_fused_sumcheck_round_inplace(
1427 &mut eq_buffer,
1428 active_pq,
1429 src_pq_size,
1430 pq_size,
1431 pq_size,
1432 pq_size >> 1,
1433 pq_size >> 1,
1434 lambda,
1435 prev_r,
1436 alpha,
1437 transcript,
1438 &mut d_sum_evals,
1439 &mut tmp_block_sums,
1440 &mut round_polys_eval,
1441 &mut r_vec,
1442 &mut prev_s_eval,
1443 xi_j,
1444 &mut eq_r_acc,
1445 device_ctx,
1446 )?;
1447 pq_size >>= 1;
1448 }
1449 }
1450
1451 unsafe {
1452 fold_ef_frac_columns_inplace(
1453 active_pq, pq_size, pq_size, pq_size, prev_r, alpha, stream,
1454 )
1455 .map_err(FractionalSumcheckError::FoldColumns)?;
1456 }
1457 active = active_pq;
1458 pq_size >>= 1;
1459 }
1460 }
1461
1462 let pq_host = [
1463 copy_from_device(active, 0, &mut copy_scratch, device_ctx)?,
1464 copy_from_device(active, pq_size / 2, &mut copy_scratch, device_ctx)?,
1465 ];
1466
1467 claims_per_layer.push(GkrLayerClaims {
1468 p_xi_0: pq_host[0].p,
1469 q_xi_0: pq_host[0].q,
1470 p_xi_1: pq_host[1].p,
1471 q_xi_1: pq_host[1].q,
1472 });
1473
1474 transcript.observe_ext(claims_per_layer[round].p_xi_0);
1475 transcript.observe_ext(claims_per_layer[round].q_xi_0);
1476 transcript.observe_ext(claims_per_layer[round].p_xi_1);
1477 transcript.observe_ext(claims_per_layer[round].q_xi_1);
1478
1479 let mu = transcript.sample_ext();
1480 xi_prev = [vec![mu], r_vec].concat();
1481
1482 sumcheck_polys.push(round_polys_eval);
1483 gkr_round_span.exit();
1484 }
1485 mem.emit_metrics_with_label("frac_sumcheck.gkr_rounds");
1486 mem.tracing_info("after_fractional_sumcheck_gkr");
1487
1488 Ok((
1489 FracSumcheckProof {
1490 fractional_sum: (root.p, root.q),
1491 claims_per_layer,
1492 sumcheck_polys,
1493 },
1494 xi_prev,
1495 ))
1496}
1497
1498fn copy_from_device<T: Copy>(
1499 buf: &DeviceBuffer<T>,
1500 index: usize,
1501 scratch: &mut DeviceBuffer<T>,
1502 device_ctx: &GpuDeviceCtx,
1503) -> Result<T, FractionalSumcheckError> {
1504 debug_assert!(!scratch.is_empty());
1505 unsafe {
1506 cuda_memcpy_on::<true, true>(
1507 scratch.as_mut_raw_ptr(),
1508 buf.as_ptr().add(index) as *const std::ffi::c_void,
1509 std::mem::size_of::<T>(),
1510 device_ctx,
1511 )?;
1512 }
1513 let host = scratch.to_host_on(device_ctx)?;
1514 Ok(host[0])
1515}
1516
1517fn reduce_to_single_evaluation<SC: StarkProtocolConfig<EF = EF>>(
1519 claims: &GkrLayerClaims<SC>,
1520 mu: EF,
1521) -> (EF, EF) {
1522 let numer = interpolate_linear_at_01(&[claims.p_xi_0, claims.p_xi_1], mu);
1523 let denom = interpolate_linear_at_01(&[claims.q_xi_0, claims.q_xi_1], mu);
1524 (numer, denom)
1525}
1526
1527fn reconstruct_s_evals(
1534 d_sum_evals: &DeviceBuffer<EF>,
1535 prev_s_eval: EF,
1536 xi_j: EF,
1537 eq_r_acc: EF,
1538 device_ctx: &GpuDeviceCtx,
1539) -> Result<([EF; GKR_S_DEG], [EF; GKR_S_DEG]), FractionalSumcheckError> {
1540 let sp_vec = d_sum_evals.to_host_on(device_ctx)?;
1541 debug_assert_eq!(sp_vec.len(), GKR_S_DEG - 1);
1542
1543 let mut sp_evals = [EF::ZERO; GKR_S_DEG];
1545 sp_evals[1] = sp_vec[0] * eq_r_acc;
1546 sp_evals[2] = sp_vec[1] * eq_r_acc;
1547
1548 let eq_xi_0 = EF::ONE - xi_j;
1555 debug_assert_ne!(eq_xi_0, EF::ZERO);
1556 let eq_xi_1 = xi_j;
1557 sp_evals[0] = (prev_s_eval - eq_xi_1 * sp_evals[1]) * eq_xi_0.inverse();
1558
1559 let s_evals: [EF; GKR_S_DEG] = from_fn(|i| {
1560 let x = EF::from_usize(i + 1);
1562 let sp_eval = if i < GKR_S_DEG - 1 {
1563 sp_evals[i + 1]
1564 } else {
1565 interpolate_quadratic_at_012(&sp_evals, x)
1566 };
1567 eval_eq_mle(&[xi_j], &[x]) * sp_eval
1568 });
1569
1570 Ok((s_evals, sp_evals))
1571}
1572
1573pub fn make_synthetic_leaves(
1575 n: usize,
1576 device_ctx: &GpuDeviceCtx,
1577) -> Result<DeviceBuffer<Frac<EF>>, FractionalSumcheckError> {
1578 use openvm_cuda_common::copy::cuda_memcpy_on;
1579 use rand::{rngs::StdRng, Rng, SeedableRng};
1580
1581 let size = 1usize << n;
1582 let mut rng = StdRng::seed_from_u64(42);
1583 let host: Vec<(EF, EF)> = (0..size)
1584 .map(|_| (rng.random::<EF>(), rng.random::<EF>()))
1585 .collect();
1586 let d_leaves = DeviceBuffer::<Frac<EF>>::with_capacity_on(size, device_ctx);
1587 unsafe {
1588 cuda_memcpy_on::<false, true>(
1589 d_leaves.as_mut_raw_ptr(),
1590 host.as_ptr() as *const std::ffi::c_void,
1591 std::mem::size_of_val(host.as_slice()),
1592 device_ctx,
1593 )?;
1594 }
1595 Ok(d_leaves)
1596}
1597
1598#[cfg(test)]
1599mod tests {
1600 use openvm_cuda_common::{
1601 common::get_device,
1602 copy::MemCopyH2D,
1603 memory_manager::MemTracker,
1604 stream::{CudaStream, GpuDeviceCtx, StreamGuard},
1605 };
1606 use p3_field::PrimeCharacteristicRing;
1607 use rand::{rngs::StdRng, Rng, SeedableRng};
1608
1609 use super::{
1610 fractional_sumcheck_gpu, make_synthetic_leaves, Frac, FractionalInputSize,
1611 FractionalSumcheckError, GkrRoundStrategy, EF,
1612 };
1613 use crate::{prelude::SC, sponge::DuplexSpongeGpu};
1614
1615 fn test_ctx() -> GpuDeviceCtx {
1616 GpuDeviceCtx {
1617 device_id: get_device().unwrap() as u32,
1618 stream: StreamGuard::new(CudaStream::new_non_blocking().unwrap()),
1619 }
1620 }
1621
1622 fn run_with_strategy(
1624 n: usize,
1625 strategy: GkrRoundStrategy,
1626 ) -> Result<(super::FracSumcheckProof<SC>, Vec<EF>), FractionalSumcheckError> {
1627 let enable_precompute_m = matches!(strategy, GkrRoundStrategy::PrecomputeM);
1629 unsafe {
1630 std::env::set_var(
1631 "SWIRL_CUDA_GKR_PRECOMPUTE_M",
1632 if enable_precompute_m { "1" } else { "0" },
1633 );
1634 }
1635 let device_ctx = test_ctx();
1636 let mut transcript = DuplexSpongeGpu::default();
1637 let leaves = make_synthetic_leaves(n, &device_ctx)?;
1638 let mut mem = MemTracker::start("test.precompute_m");
1639 let result = fractional_sumcheck_gpu(
1640 &mut transcript,
1641 leaves,
1642 FractionalInputSize::dense(1usize << n),
1643 EF::ZERO,
1644 false,
1645 &mut mem,
1646 &device_ctx,
1647 )?;
1648 device_ctx.stream.synchronize().expect("sync");
1649 Ok(result)
1650 }
1651
1652 fn assert_proofs_equal(
1653 a: &(super::FracSumcheckProof<SC>, Vec<EF>),
1654 b: &(super::FracSumcheckProof<SC>, Vec<EF>),
1655 ) {
1656 assert_proofs_equal_with_context(a, b, "proof");
1657 }
1658
1659 fn assert_proofs_equal_with_context(
1660 a: &(super::FracSumcheckProof<SC>, Vec<EF>),
1661 b: &(super::FracSumcheckProof<SC>, Vec<EF>),
1662 context: &str,
1663 ) {
1664 assert_eq!(
1665 a.0.fractional_sum, b.0.fractional_sum,
1666 "{context}: fractional_sum mismatch"
1667 );
1668 assert_eq!(
1669 a.0.claims_per_layer, b.0.claims_per_layer,
1670 "{context}: claims_per_layer mismatch"
1671 );
1672 assert_eq!(
1673 a.0.sumcheck_polys, b.0.sumcheck_polys,
1674 "{context}: sumcheck_polys mismatch"
1675 );
1676 assert_eq!(a.1, b.1, "{context}: final randomness mismatch");
1677 }
1678
1679 fn assert_virtual_matches_dense(
1680 real: &[Frac<EF>],
1681 real_len: usize,
1682 logical_len: usize,
1683 alpha: EF,
1684 strategy: GkrRoundStrategy,
1685 ) -> Result<(), FractionalSumcheckError> {
1686 let mut dense = real.to_vec();
1687 dense.resize(logical_len, Frac::default());
1688
1689 let virtual_proof = run_from_host(real, real_len, logical_len, alpha, strategy)?;
1690 let dense_proof = run_from_host(&dense, logical_len, logical_len, alpha, strategy)?;
1691 let context =
1692 format!("strategy={strategy:?}, real_len={real_len}, logical_len={logical_len}");
1693 assert_proofs_equal_with_context(&virtual_proof, &dense_proof, &context);
1694 Ok(())
1695 }
1696
1697 fn make_host_leaves(len: usize) -> Vec<Frac<EF>> {
1698 let mut rng = StdRng::seed_from_u64(20260429);
1699 (0..len)
1700 .map(|_| Frac {
1701 p: rng.random::<EF>(),
1702 q: rng.random::<EF>(),
1703 })
1704 .collect()
1705 }
1706
1707 fn virtual_padding_test_alpha() -> EF {
1708 EF::from_u32(7)
1709 }
1710
1711 fn run_from_host(
1712 host: &[Frac<EF>],
1713 real_len: usize,
1714 logical_len: usize,
1715 alpha: EF,
1716 strategy: GkrRoundStrategy,
1717 ) -> Result<(super::FracSumcheckProof<SC>, Vec<EF>), FractionalSumcheckError> {
1718 let enable_precompute_m = matches!(strategy, GkrRoundStrategy::PrecomputeM);
1720 unsafe {
1721 std::env::set_var(
1722 "SWIRL_CUDA_GKR_PRECOMPUTE_M",
1723 if enable_precompute_m { "1" } else { "0" },
1724 );
1725 if enable_precompute_m {
1726 std::env::set_var("SWIRL_CUDA_GKR_PRECOMPUTE_M_MIN_BLOCKS", "1");
1727 std::env::set_var("SWIRL_CUDA_GKR_PRECOMPUTE_M_MIN_N", "4");
1728 } else {
1729 std::env::remove_var("SWIRL_CUDA_GKR_PRECOMPUTE_M_MIN_BLOCKS");
1730 std::env::remove_var("SWIRL_CUDA_GKR_PRECOMPUTE_M_MIN_N");
1731 }
1732 }
1733 let device_ctx = test_ctx();
1734 let leaves = host.to_device_on(&device_ctx)?;
1735 let mut transcript = DuplexSpongeGpu::default();
1736 let mut mem = MemTracker::start("test.virtual_input_padding");
1737 let result = fractional_sumcheck_gpu(
1738 &mut transcript,
1739 leaves,
1740 FractionalInputSize::new(real_len, logical_len),
1741 alpha,
1742 false,
1743 &mut mem,
1744 &device_ctx,
1745 )?;
1746 device_ctx.stream.synchronize().expect("sync");
1747 Ok(result)
1748 }
1749
1750 #[test]
1751 fn test_virtual_input_padding_matches_dense_padding() -> Result<(), FractionalSumcheckError> {
1752 let small_cases = [4, 8, 16, 32].into_iter().flat_map(|logical_len| {
1753 (logical_len / 2..logical_len).map(move |real_len| (real_len, logical_len))
1754 });
1755 for (real_len, logical_len) in small_cases.chain([(1500, 2048)]) {
1756 let real = make_host_leaves(real_len);
1757 assert_virtual_matches_dense(
1758 &real,
1759 real_len,
1760 logical_len,
1761 virtual_padding_test_alpha(),
1762 GkrRoundStrategy::FoldEval,
1763 )?;
1764 }
1765 Ok(())
1766 }
1767
1768 #[test]
1769 fn test_virtual_input_padding_matches_dense_padding_precompute_m(
1770 ) -> Result<(), FractionalSumcheckError> {
1771 for (real_len, logical_len) in [
1772 (33, 64),
1773 (47, 64),
1774 (63, 64),
1775 (65, 128),
1776 (96, 128),
1777 (127, 128),
1778 (1025, 2048),
1779 (1500, 2048),
1780 (2047, 2048),
1781 (32769, 65536),
1782 (49153, 65536),
1783 (65535, 65536),
1784 ] {
1785 let real = make_host_leaves(real_len);
1786 assert_virtual_matches_dense(
1787 &real,
1788 real_len,
1789 logical_len,
1790 virtual_padding_test_alpha(),
1791 GkrRoundStrategy::PrecomputeM,
1792 )?;
1793 }
1794 Ok(())
1795 }
1796
1797 #[test]
1802 fn test_precompute_m_matches_fused() -> Result<(), FractionalSumcheckError> {
1803 for n in [24, 25, 26] {
1804 eprintln!("--- testing n={n} ---");
1805 let fused = run_with_strategy(n, GkrRoundStrategy::FoldEval)?;
1806 let precompute = run_with_strategy(n, GkrRoundStrategy::PrecomputeM)?;
1807 assert_proofs_equal(&fused, &precompute);
1808 }
1809 Ok(())
1810 }
1811
1812 #[test]
1816 fn test_precompute_m_multi_window_matches_fused() -> Result<(), FractionalSumcheckError> {
1817 unsafe {
1818 std::env::set_var("SWIRL_CUDA_GKR_PRECOMPUTE_M_MIN_N", "8");
1819 std::env::set_var("SWIRL_CUDA_GKR_PRECOMPUTE_M_MIN_BLOCKS", "1");
1820 }
1821 let fused = run_with_strategy(16, GkrRoundStrategy::FoldEval)?;
1822 let precompute = run_with_strategy(16, GkrRoundStrategy::PrecomputeM)?;
1823 assert_proofs_equal(&fused, &precompute);
1824 unsafe {
1825 std::env::remove_var("SWIRL_CUDA_GKR_PRECOMPUTE_M_MIN_N");
1826 std::env::remove_var("SWIRL_CUDA_GKR_PRECOMPUTE_M_MIN_BLOCKS");
1827 }
1828 Ok(())
1829 }
1830}