openvm_recursion_circuit/whir/folding/
trace.rs

1use core::{borrow::BorrowMut, convert::TryInto};
2
3use openvm_stark_sdk::config::baby_bear_poseidon2::{EF, F};
4use p3_field::{BasedVectorSpace, PrimeCharacteristicRing};
5use p3_matrix::dense::RowMajorMatrix;
6
7use super::WhirFoldingCols;
8use crate::{
9    tracegen::{RowMajorChip, StandardTracegenCtx},
10    whir::WhirBlobCpu,
11};
12
13#[repr(C)]
14#[derive(Clone, Copy, Debug, Default)]
15pub struct FoldRecord {
16    pub whir_round: u32,
17    pub query_idx: u32,
18    pub coset_idx: u32,
19    pub height: u32,
20    pub coset_size: u32,
21    pub coset_shift: F,
22    pub twiddle: F,
23    pub z_final: F,
24    pub value: EF,
25    pub left_value: EF,
26    pub right_value: EF,
27    pub y_final: EF,
28    pub alpha: EF,
29}
30
31impl FoldRecord {
32    #[allow(clippy::too_many_arguments)]
33    pub fn new(
34        whir_round: usize,
35        query_idx: usize,
36        twiddle: F,
37        coset_shift: F,
38        coset_size: usize,
39        coset_idx: usize,
40        height: usize,
41        left_value: EF,
42        right_value: EF,
43        value: EF,
44        alpha: EF,
45    ) -> Self {
46        debug_assert!(height > 0, "folding record height must be > 0");
47        Self {
48            whir_round: whir_round.try_into().unwrap(),
49            query_idx: query_idx.try_into().unwrap(),
50            coset_idx: coset_idx.try_into().unwrap(),
51            height: height.try_into().unwrap(),
52            coset_size: coset_size.try_into().unwrap(),
53            coset_shift,
54            twiddle,
55            value,
56            left_value,
57            right_value,
58            z_final: F::ZERO,
59            y_final: EF::ZERO,
60            alpha,
61        }
62    }
63
64    pub fn set_final_values(&mut self, z_final: F, y_final: EF) {
65        self.z_final = z_final;
66        self.y_final = y_final;
67    }
68}
69
70pub(crate) struct FoldingTraceGenerator;
71
72impl RowMajorChip<F> for FoldingTraceGenerator {
73    type Ctx<'a> = (StandardTracegenCtx<'a>, &'a WhirBlobCpu);
74
75    #[tracing::instrument(level = "trace", skip_all)]
76    fn generate_trace(
77        &self,
78        ctx: &Self::Ctx<'_>,
79        required_height: Option<usize>,
80    ) -> Option<RowMajorMatrix<F>> {
81        let fold_records = &ctx.1.fold_records;
82        let num_rows_per_proof = fold_records.layout().items_per_proof();
83        let num_valid_rows = fold_records.len();
84        let height = if let Some(h) = required_height {
85            if h < num_valid_rows {
86                return None;
87            }
88            h
89        } else {
90            num_valid_rows.next_power_of_two()
91        };
92        let width = WhirFoldingCols::<F>::width();
93
94        let mut trace = vec![F::ZERO; height * width];
95
96        for (row_idx, row) in trace.chunks_mut(width).take(num_valid_rows).enumerate() {
97            let proof_idx = row_idx / num_rows_per_proof;
98            let i = row_idx % num_rows_per_proof;
99            let record = fold_records[(proof_idx, i)];
100            let height = record.height as usize;
101
102            let cols: &mut WhirFoldingCols<F> = row.borrow_mut();
103            cols.is_valid = F::ONE;
104            cols.proof_idx = F::from_usize(proof_idx);
105            cols.is_root = F::from_bool(record.coset_size == 1);
106            cols.alpha
107                .copy_from_slice(record.alpha.as_basis_coefficients_slice());
108            cols.height = F::from_usize(height);
109            cols.whir_round = F::from_u32(record.whir_round);
110            cols.query_idx = F::from_u32(record.query_idx);
111            cols.coset_idx = F::from_u32(record.coset_idx);
112            cols.left_value
113                .copy_from_slice(record.left_value.as_basis_coefficients_slice());
114            cols.right_value
115                .copy_from_slice(record.right_value.as_basis_coefficients_slice());
116            cols.value
117                .copy_from_slice(record.value.as_basis_coefficients_slice());
118            cols.twiddle = record.twiddle;
119            cols.coset_shift = record.coset_shift;
120            cols.coset_size = F::from_u32(record.coset_size);
121            cols.z_final = record.z_final;
122            cols.y_final
123                .copy_from_slice(record.y_final.as_basis_coefficients_slice());
124        }
125
126        Some(RowMajorMatrix::new(trace, width))
127    }
128}