openvm_recursion_circuit/whir/folding/
trace.rs1use 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}