Skip to main content

openvm_cuda_backend/logup_zerocheck/
fractional.rs

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    /// Peak work-buffer bytes (excluding the `S_frac * real_len` layer/input buffer).
63    ///
64    /// This must stay in sync with `max_work_size` in `fractional_sumcheck_gpu` and the
65    /// precompute-M EF auxiliary allocations. If the sumcheck implementation changes, this
66    /// method must be updated as well — it is the source of truth for batching budgets that
67    /// depend on the fractional-GKR peak. Also update the conservative interaction memory
68    /// estimate in `openvm_stark_backend::memory_metering`.
69    ///
70    /// ## Formula (FoldEval path, dominates for large inputs)
71    ///
72    /// ```text
73    /// work_buffer = logical_len / 4   Frac entries → S_frac * L/4 bytes
74    /// ```
75    ///
76    /// For the precompute-M path (`max_work_size` in `fractional_sumcheck_gpu`):
77    ///
78    /// ```text
79    /// work_buffer = max(L >> (1 + GKR_WINDOW_SIZE), 2^22)   Frac entries
80    /// + S_ef * (2^(2w) + max_window(m_partial) + 2^(w+1))  EF bytes
81    /// ```
82    ///
83    /// This method returns the maximum over both paths to give a conservative budget.
84    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        // Conservative estimate for M_precompute_EF: m_total + m_partial bound + eq_prefix/suffix.
93        // In practice the max_window term is ceil(2^(rem_n - w) / tail_tile) * 2^(2w) EF elems;
94        // use 2 * 2^(2w) as a safe floor (tail_tile >= 1, rem_n bounded by total_rounds).
95        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/// Describes which buffer operation to use for the next fused compute+fold round.
102#[derive(Debug, Clone, Copy)]
103enum BufferTarget {
104    /// Out-of-place: layer → work_buffer
105    LayerToWork,
106    /// Out-of-place: work_buffer → layer
107    WorkToLayer,
108    /// In-place on layer buffer
109    InPlaceLayer,
110    /// In-place on work_buffer
111    InPlaceWork,
112}
113
114/// Encapsulates ping-pong buffer scheduling state for GKR inner rounds.
115///
116/// This struct manages the decision of whether to use in-place or out-of-place
117/// (ping-pong) kernel variants based on buffer capacities and current data location.
118struct BufferScheduler {
119    /// True if data currently resides in work_buffer, false if in layer.
120    data_in_work_buffer: bool,
121    /// Maximum capacity of work_buffer in elements.
122    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    /// Creates a new scheduler with data initially in layer buffer.
138    fn new(work_buffer_cap: usize) -> Self {
139        Self {
140            data_in_work_buffer: false,
141            work_buffer_cap,
142        }
143    }
144
145    /// Returns true if we can use ping-pong (out-of-place) for the given post-fold size.
146    fn can_pingpong(&self, post_fold_size: usize) -> bool {
147        post_fold_size <= self.work_buffer_cap
148    }
149
150    /// Determines the next buffer target for a fused compute+fold operation.
151    ///
152    /// For last outer round: uses ping-pong when possible for __restrict__ optimization.
153    /// For non-last outer rounds: preserves layer for tree revert operations.
154    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                // Ping-pong to other buffer
160                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                // In-place on layer (data must be in layer for early rounds of last outer round)
169                debug_assert!(
170                    !self.data_in_work_buffer,
171                    "in-place path requires data in layer"
172                );
173                BufferTarget::InPlaceLayer
174            }
175        } else {
176            // Non-last outer round: preserve layer for tree revert
177            if self.data_in_work_buffer {
178                // Already in work_buffer, stay there in-place
179                BufferTarget::InPlaceWork
180            } else {
181                // Data in layer, move to work_buffer (out-of-place)
182                self.data_in_work_buffer = true;
183                BufferTarget::LayerToWork
184            }
185        }
186    }
187
188    /// Returns the buffer target for the final fold (no fused compute).
189    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            // Non-last outer round: fold into work_buffer to preserve layer
198            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    // w <= 5 so CPU builds are trivial; avoid GPU kernel/alloc overhead for tiny tables.
321    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
334/// Get low/high eq pointers for the tail portion of the eq buffer, skipping `drop_count` layers.
335/// See `docs/cuda-backend/gkr-prover.md` § "Eq buffer sqrt decomposition".
336fn 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    // The first GKR fold drops bit 1 of the logical input index under the bit-reversed
407    // layout. The compact threshold is therefore the max prefix index after removing
408    // that bit from any real source position.
409    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/// Observes s_evals in transcript, updates accumulators, and returns the sampled challenge.
441#[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/// Fused revert + compute round: reverts the tree layer and computes s'_0(1) and s'_0(2).
475///
476/// See `docs/cuda-backend/gkr-prover.md` § "Sumcheck round strategies" for context.
477#[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/// Fused compute round: computes s' polynomial AND folds the pq_buffer for next round.
528///
529/// This kernel fuses the fold operation (using `r_prev` from the previous round) into the current
530/// round's compute, eliminating one kernel launch and reducing memory traffic.
531#[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/// In-place variant of [`do_fused_sumcheck_round`]. Reads and writes to the same buffer.
588#[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/// GKR fractional sumcheck prover. See `docs/cuda-backend/gkr-prover.md` (repo root) for the
647/// protocol and implementation details.
648#[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    // total_rounds = l_skip + n_logup
695    let total_rounds = log2_strict_usize(total_leaves);
696    assert!(total_rounds > 0, "n_logup > 0 when there are interactions");
697    // Build segment tree.
698    // - We only maintain the current layer
699    // - Input layer uses separate F and EF buffers to save memory
700    // - First tree layer converts (F, EF) to FracExt (EF, EF)
701
702    // We store it in bit-reversal order for coalesced memory accesses.
703    // For large N (> 1024), fuse bitrev + tree layers 0 and 1 into a single kernel pass,
704    // eliminating ~1.5N global memory reads. For small N, fall back to separate operations.
705    let virtual_input = real_len < total_leaves;
706    let start_layer_i = if total_leaves > 1024 {
707        unsafe {
708            // SAFETY: Frac<EF> has exact same memory layout and alignment as (EF, EF).
709            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 // layers 0+1 already done
714    } else {
715        // Fallback: separate bitrev + layer 0.
716        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 // layer 0 done; continue from layer 1
751    };
752
753    // Remaining layers: fuse consecutive pairs via two-layer kernel (~33% traffic savings).
754    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    // Remaining single layer (if odd number of layers left).
764    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    // Work buffer to avoid revert operations on layer. Only needed for non-last rounds.
842    // For the last round (round == total_rounds - 1), we fold in-place on layer.
843    // When precompute-M is active, the first multifold folds w+1 variables at once
844    // (pending r_prev + w window challenges). The largest case is the last outer round:
845    // input pq_size = total_leaves, output pq_size = total_leaves >> (1 + w).
846    // Fold-eval fallback rounds (rem_n < min_n) need at most 2^min_n elements.
847    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    // Note: `stream` local variable is used for FFI calls throughout the loop.
881
882    for round in 1..total_rounds {
883        let gkr_round_span = debug_span!("GKR", round).entered();
884
885        // Note: frac_build_tree_layer revert is now fused into do_sumcheck_round_and_revert below.
886
887        debug_assert_eq!(xi_prev.len(), round);
888        // eq_buffer stores eq(xi_prev[j..], x) for x in H_{xi_prev.len()-j} for
889        // j=1,...,xi_prev.len()-1.
890        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        // In round `j`, contains `s_{j-1}(r_{j-1})`. Starts with the sumcheck's sum claim.
917        let (numer_claim, denom_claim) =
918            reduce_to_single_evaluation(claims_per_layer.last().unwrap(), /* mu */ xi_prev[0]);
919        let mut prev_s_eval = numer_claim + lambda * denom_claim;
920        let mut eq_r_acc = EF::ONE;
921
922        // Round 0: compute + revert fused. The pq_buffer fold will be fused into next round's
923        // compute. This fuses frac_build_tree_layer(revert=true) with the first inner round
924        // compute.
925        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        // Fused rounds 1..(round-1): compute + fold using prev_r.
944        let mut prev_r = r0;
945        let active: &mut DeviceBuffer<Frac<EF>>;
946
947        match backend {
948            GkrRoundStrategy::FoldEval => {
949                // Existing fused path.
950                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                // Final fold after last r (no next compute to fuse with).
1066                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                // First window reads from `layer` with pending_fold=true
1165                // (M-build folds prev_r inline). Multifold writes to
1166                // `active_pq` (work_buffer for non-last rounds, layer for
1167                // last). After the first window, subsequent windows
1168                // read/write active_pq with pending_fold=false.
1169                let mut pending_fold = true;
1170                let layer_read_ptr = layer.as_ptr();
1171                // Virtual compact reads can recover spilled real entries from compact source slots
1172                // that are not owned by the output thread. Dense in-place multifold is safe, but
1173                // virtual last-round multifold must write to the existing work buffer to avoid
1174                // clobbering a source slot before another block reads it.
1175                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                // w+1 to accommodate r_prev prepended to window challenges
1184                // on the first iteration (inline fold on last outer round).
1185                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                    // Reuse tmp_block_sums for eq_r_window upload.
1218                    // Safe: M eval rounds finish before multifold needs eq_r_window,
1219                    // and tmp_block_sums is not needed until next fold-eval round.
1220                    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                    // When pending_fold is true, the buffer has rem_n+1 variables
1230                    // and the kernel folds inline. The effective tail dimension
1231                    // for tiling is the same either way (rem_n - w).
1232                    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; // save before eval loop overwrites prev_r
1251                    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, // inline fold only on first iteration
1277                            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                    // Compute eq_r_window for the multifold.
1334                    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                    // First tail round is standalone compute,
1399                    // subsequent rounds are fold+compute.
1400                    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
1517/// Reduces claims to a single evaluation point using linear interpolation.
1518fn 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
1527/// Reconstructs the full s(1,2,3) evaluations from s'(1,2) evaluations returned by GPU.
1528///
1529/// Goal: compute s({1,2,3}) from s(X) = eq(xi_j, X) * s'(X).
1530/// Reconstruct the full round polynomial s_t from GPU-computed s'_t(1), s'_t(2).
1531///
1532/// See `docs/cuda-backend/gkr-prover.md` § "Sumcheck round implementation" for the derivation.
1533fn 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    // sp_evals holds evaluations of degree 2 poly `eq(xi_{j+1..}, r_{j+1..}) * s'(X)` at {0,1,2}
1544    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    // We use that s_j(0) + s_j(1) = s_{j-1}(r_{j-1})
1549    // s_j(X) = eq(xi_j, X) * sp_j(X)
1550    // s_j(0) = (1 - xi_j) * sp_j(0)
1551    // s_j(1) = xi_j * sp_j(1)
1552    // So: (1 - xi_j) * sp_j(0) + xi_j * sp_j(1) = prev_s_eval
1553    // xi_j is randomly sampled so 1 - xi_j should be invertible
1554    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        // evaluate s at X = i + 1 (skip 0 evaluation)
1561        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
1573/// Generate random fractional leaves on device for benchmarking.
1574pub 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    /// Run fractional sumcheck with a given round strategy and return the proof + final randomness.
1623    fn run_with_strategy(
1624        n: usize,
1625        strategy: GkrRoundStrategy,
1626    ) -> Result<(super::FracSumcheckProof<SC>, Vec<EF>), FractionalSumcheckError> {
1627        // SAFETY: test sets process env; run with --test-threads=1.
1628        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        // SAFETY: test sets process env; run with --test-threads=1.
1719        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    /// Compares precompute-M against FoldEval at n=24,25,26.
1798    /// n=24: top layer rem_n=22 (one window).
1799    /// n=25: top layer rem_n=23 (one window, more tail).
1800    /// n=26: top layer rem_n=24 (one window, even more tail).
1801    #[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    /// Compares precompute-M against FoldEval with lowered thresholds to force multi-window
1813    /// iteration at small n. At n=16, round=15: stop=8, window 1 at base=1 (rem_n=14),
1814    /// window 2 at base=4 (rem_n=11), then fold-eval tail for the remaining round.
1815    #[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}