openvm_recursion_circuit/batch_constraint/eq_airs/eq_ns/
trace.rs

1use std::borrow::BorrowMut;
2
3use openvm_stark_backend::keygen::types::MultiStarkVerifyingKey;
4use openvm_stark_sdk::config::baby_bear_poseidon2::{BabyBearPoseidon2Config, EF, F};
5use p3_field::{BasedVectorSpace, PrimeCharacteristicRing};
6use p3_matrix::dense::RowMajorMatrix;
7use p3_maybe_rayon::prelude::*;
8
9use crate::{
10    batch_constraint::{eq_airs::eq_ns::air::EqNsColumns, SelectorCount},
11    system::Preflight,
12    tracegen::RowMajorChip,
13    utils::MultiVecWithBounds,
14};
15
16#[derive(Clone, Copy)]
17#[repr(C)]
18pub struct EqNsRecord {
19    xi: EF,
20    r: EF,
21    eq_r_ones: EF,
22    eq_r_zeroes: EF,
23    r_prod: EF,
24    eq: EF,
25    eq_sharp: EF,
26    sel_first_count: usize,
27    sel_last_and_trans_count: usize,
28    n_logup: usize,
29    n_max: usize,
30}
31
32pub struct EqNsTraceGenerator;
33
34impl RowMajorChip<F> for EqNsTraceGenerator {
35    type Ctx<'a> = (
36        &'a MultiStarkVerifyingKey<BabyBearPoseidon2Config>,
37        &'a [&'a Preflight],
38        &'a MultiVecWithBounds<SelectorCount, 1>,
39    );
40
41    #[tracing::instrument(level = "trace", skip_all)]
42    fn generate_trace(
43        &self,
44        ctx: &Self::Ctx<'_>,
45        required_height: Option<usize>,
46    ) -> Option<RowMajorMatrix<F>> {
47        let (vk, preflights, selector_counts) = ctx;
48        let l_skip = vk.inner.params.l_skip;
49        let records = preflights
50            .iter()
51            .enumerate()
52            .map(|(pidx, preflight)| {
53                let selector_counts = &selector_counts[[pidx]];
54                let n_global = preflight.proof_shape.n_global();
55                let n_max = preflight.proof_shape.n_max;
56                let rs = &preflight.batch_constraint.sumcheck_rnd;
57                let xi = &preflight.batch_constraint.xi;
58                let mut res = Vec::with_capacity(n_global + 1);
59                let mut eq_r_ones = EF::ONE;
60                let mut eq_r_zeroes = EF::ONE;
61                for i in 0..n_max {
62                    let counts = selector_counts[i + l_skip];
63                    res.push(EqNsRecord {
64                        xi: xi[l_skip + i],
65                        r: rs[1 + i],
66                        r_prod: EF::ONE,
67                        eq: preflight.batch_constraint.eq_ns[i],
68                        eq_sharp: preflight.batch_constraint.eq_sharp_ns[i],
69                        eq_r_ones,
70                        eq_r_zeroes,
71                        n_logup: preflight.proof_shape.n_logup,
72                        n_max: preflight.proof_shape.n_max,
73                        sel_first_count: counts.first,
74                        sel_last_and_trans_count: counts.last + counts.transition,
75                    });
76                    eq_r_ones *= rs[1 + i];
77                    eq_r_zeroes *= EF::ONE - rs[1 + i];
78                }
79                let counts = selector_counts[l_skip + n_max];
80                for i in n_max..=n_global {
81                    let (sel_first_count, sel_last_and_trans_count) = if i == n_max {
82                        (counts.first, counts.last + counts.transition)
83                    } else {
84                        (0, 0)
85                    };
86                    let xi = if i == n_global {
87                        EF::ZERO
88                    } else {
89                        xi[l_skip + i]
90                    };
91                    res.push(EqNsRecord {
92                        xi,
93                        r: EF::ONE,
94                        r_prod: EF::ONE,
95                        eq: preflight.batch_constraint.eq_ns[n_max],
96                        eq_sharp: preflight.batch_constraint.eq_sharp_ns[n_max],
97                        eq_r_ones,
98                        eq_r_zeroes,
99                        n_logup: preflight.proof_shape.n_logup,
100                        n_max: preflight.proof_shape.n_max,
101                        sel_first_count,
102                        sel_last_and_trans_count,
103                    });
104                }
105                for i in (0..n_global).rev() {
106                    res[i].r_prod = res[i + 1].r_prod * res[i].r;
107                }
108                res
109            })
110            .collect::<Vec<_>>();
111
112        let width = EqNsColumns::<F>::width();
113        let total_height = records.iter().map(|rows| rows.len()).sum::<usize>();
114        let padded_height = if let Some(height) = required_height {
115            if height < total_height {
116                return None;
117            }
118            height
119        } else {
120            total_height.next_power_of_two()
121        };
122
123        let mut trace = vec![F::ZERO; padded_height * width];
124        let mut cur_height = 0;
125        for (pidx, rows) in records.iter().enumerate() {
126            trace[cur_height * width..(cur_height + rows.len()) * width]
127                .par_chunks_exact_mut(width)
128                .zip(rows.par_iter())
129                .enumerate()
130                .for_each(|(i, (chunk, record))| {
131                    let cols: &mut EqNsColumns<_> = chunk.borrow_mut();
132                    cols.is_valid = F::ONE;
133                    cols.is_first = F::from_bool(i == 0);
134                    cols.proof_idx = F::from_usize(pidx);
135                    cols.n = F::from_usize(i);
136                    cols.n_logup = F::from_usize(record.n_logup);
137                    cols.n_max = F::from_usize(record.n_max);
138                    cols.n_less_than_n_logup = F::from_bool(i < record.n_logup);
139                    cols.n_less_than_n_max = F::from_bool(i < record.n_max);
140                    cols.is_transition_and_n_less_than_n_max =
141                        F::from_bool(i + 1 < rows.len() && i < record.n_max);
142                    cols.xi_n
143                        .copy_from_slice(record.xi.as_basis_coefficients_slice());
144                    cols.r_n
145                        .copy_from_slice(record.r.as_basis_coefficients_slice());
146                    cols.r_product
147                        .copy_from_slice(record.r_prod.as_basis_coefficients_slice());
148                    cols.r_pref_product
149                        .copy_from_slice(record.eq_r_ones.as_basis_coefficients_slice());
150                    cols.one_minus_r_pref_prod
151                        .copy_from_slice(record.eq_r_zeroes.as_basis_coefficients_slice());
152                    cols.eq
153                        .copy_from_slice(record.eq.as_basis_coefficients_slice());
154                    cols.eq_sharp
155                        .copy_from_slice(record.eq_sharp.as_basis_coefficients_slice());
156                    cols.sel_first_count = F::from_usize(record.sel_first_count);
157                    cols.sel_last_and_trans_count = F::from_usize(record.sel_last_and_trans_count);
158                });
159            let mut num_n_lift_int = vec![0; preflights[pidx].proof_shape.n_logup + 1];
160            let mut num_n_lift_con = vec![0; preflights[pidx].proof_shape.n_max + 1];
161            for (air_idx, vdata) in preflights[pidx].proof_shape.sorted_trace_vdata.iter() {
162                let num_interactions = vk.inner.per_air[*air_idx].num_interactions();
163                let n_lift = vdata.log_height.saturating_sub(l_skip);
164                num_n_lift_con[n_lift] += 1;
165                if num_interactions > 0 {
166                    num_n_lift_int[n_lift] += num_interactions;
167                }
168            }
169            let mut xi_mult = 0;
170            for (chunk, cnt) in trace[cur_height * width..]
171                .chunks_mut(width)
172                .zip(num_n_lift_int.into_iter())
173            {
174                xi_mult += cnt;
175                let cols: &mut EqNsColumns<_> = chunk.borrow_mut();
176                cols.xi_mult = F::from_usize(xi_mult);
177            }
178            for (chunk, cnt) in trace[cur_height * width..]
179                .chunks_mut(width)
180                .zip(num_n_lift_con.into_iter())
181            {
182                let cols: &mut EqNsColumns<_> = chunk.borrow_mut();
183                cols.num_traces = F::from_usize(cnt);
184            }
185            cur_height += rows.len();
186        }
187
188        Some(RowMajorMatrix::new(trace, width))
189    }
190}