openvm_recursion_circuit/batch_constraint/eq_airs/eq_ns/
trace.rs1use 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}