1use std::sync::Arc;
7
8use openvm_stark_backend::air_builders::symbolic::{
9 symbolic_expression::SymbolicExpression,
10 symbolic_variable::{Entry, SymbolicVariable},
11 SymbolicConstraints, SymbolicExpressionDag, SymbolicExpressionNode,
12};
13use p3_field::Field;
14use rustc_hash::FxHashMap;
15
16#[repr(C)]
18#[derive(Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Debug)]
19pub struct PackedVar(pub u32);
20
21impl PackedVar {
22 pub fn new(entry_type: u8, part_index: u8, offset: u8, col_index: u16) -> Self {
23 assert!(offset < 16, "PackedVar offset must fit in 4 bits");
24 Self(
25 (entry_type as u32)
26 | ((part_index as u32) << 4)
27 | ((offset as u32) << 12)
28 | ((col_index as u32) << 16),
29 )
30 }
31
32 pub fn from_symbolic_var<F: Field>(var: &SymbolicVariable<F>) -> Self {
33 assert!(
34 var.index <= u16::MAX as usize,
35 "symbolic column index exceeds PackedVar capacity"
36 );
37 let (entry_type, part_index, offset) = match var.entry {
38 Entry::Main { part_index, offset } => (1, part_index as u8, offset as u8),
39 Entry::Preprocessed { offset } => (0, 0, offset as u8),
40 Entry::Public => (3, 0, 0),
41 Entry::Challenge => {
42 panic!("unsupported symbolic entry in zerocheck monomial extraction")
43 }
44 };
45 Self::new(entry_type, part_index, offset, var.index as u16)
46 }
47
48 pub fn is_first() -> Self {
49 Self::new(8, 0, 0, 0)
50 }
51
52 pub fn is_last() -> Self {
53 Self::new(9, 0, 0, 0)
54 }
55
56 pub fn is_transition() -> Self {
57 Self::new(10, 0, 0, 0)
58 }
59}
60
61#[repr(C)]
62#[derive(Clone, Copy, Debug, PartialEq, Eq)]
63pub struct MonomialHeader {
64 pub var_offset: u32,
65 pub term_offset: u32,
66 pub num_vars: u16,
67 pub num_terms: u16,
68}
69
70#[repr(C)]
72#[derive(Clone, Copy, Debug, PartialEq, Eq)]
73pub struct LambdaTerm<F> {
74 pub constraint_idx: u32,
75 pub coefficient: F,
76}
77
78#[repr(C)]
82#[derive(Clone, Copy, Debug, PartialEq, Eq)]
83pub struct InteractionMonomialTerm<F> {
84 pub coefficient: F,
85 pub interaction_idx: u16,
86 pub field_idx: u16, }
88
89#[derive(Clone)]
90struct ExpandedMonomial<F> {
91 pub variables: Vec<PackedVar>,
92 pub coefficient: F,
93}
94
95#[derive(Clone, Debug, PartialEq, Eq)]
97pub struct ExpandedMonomials<F> {
98 pub headers: Vec<MonomialHeader>,
99 pub variables: Vec<PackedVar>,
100 pub lambda_terms: Vec<LambdaTerm<F>>,
101}
102
103impl<F: Field> ExpandedMonomials<F> {
104 pub fn from_dag(dag: &SymbolicExpressionDag<F>) -> Self {
105 let mut cache: FxHashMap<usize, Arc<[ExpandedMonomial<F>]>> = FxHashMap::default();
106
107 let mut monomial_map: FxHashMap<Vec<PackedVar>, Vec<(u32, F)>> = FxHashMap::default();
108
109 for (constraint_idx, &dag_idx) in dag.constraint_idx.iter().enumerate() {
110 let expanded = expand_node_cached(dag, dag_idx, &mut cache);
111 for mono in expanded.iter() {
112 if mono.coefficient == F::ZERO {
113 continue;
114 }
115 monomial_map
116 .entry(mono.variables.clone())
117 .or_default()
118 .push((constraint_idx as u32, mono.coefficient));
119 }
120 }
121
122 let mut monomials: Vec<_> = monomial_map.into_iter().collect();
123 monomials.sort_by(|(vars_a, _), (vars_b, _)| {
124 vars_a
125 .len()
126 .cmp(&vars_b.len())
127 .then_with(|| vars_a.cmp(vars_b))
128 });
129
130 let (headers, all_vars, all_lambda_terms) = serialize_monomials(
131 monomials,
132 |terms| terms.sort_by_key(|(constraint_idx, _)| *constraint_idx),
133 |(idx, coeff)| LambdaTerm {
134 constraint_idx: idx,
135 coefficient: coeff,
136 },
137 );
138
139 Self {
140 headers,
141 variables: all_vars,
142 lambda_terms: all_lambda_terms,
143 }
144 }
145}
146
147fn expand_node_cached<F: Field>(
148 dag: &SymbolicExpressionDag<F>,
149 idx: usize,
150 cache: &mut FxHashMap<usize, Arc<[ExpandedMonomial<F>]>>,
151) -> Arc<[ExpandedMonomial<F>]> {
152 if let Some(cached) = cache.get(&idx) {
153 return Arc::clone(cached);
154 }
155
156 let result_vec = match &dag.nodes[idx] {
157 SymbolicExpressionNode::Constant(c) => expand_leaf(vec![], *c),
158 SymbolicExpressionNode::Variable(v) => {
159 expand_leaf(vec![PackedVar::from_symbolic_var(v)], F::ONE)
160 }
161 SymbolicExpressionNode::IsFirstRow => expand_leaf(vec![PackedVar::is_first()], F::ONE),
162 SymbolicExpressionNode::IsLastRow => expand_leaf(vec![PackedVar::is_last()], F::ONE),
163 SymbolicExpressionNode::IsTransition => {
164 expand_leaf(vec![PackedVar::is_transition()], F::ONE)
165 }
166 SymbolicExpressionNode::Add {
167 left_idx,
168 right_idx,
169 ..
170 } => {
171 let left = expand_node_cached(dag, *left_idx, cache);
172 let right = expand_node_cached(dag, *right_idx, cache);
173 let mut result = Vec::with_capacity(left.len() + right.len());
174 result.extend(left.iter().cloned());
175 result.extend(right.iter().cloned());
176 combine_like_terms(result)
177 }
178 SymbolicExpressionNode::Sub {
179 left_idx,
180 right_idx,
181 ..
182 } => {
183 let left = expand_node_cached(dag, *left_idx, cache);
184 let right = expand_node_cached(dag, *right_idx, cache);
185 let mut result = Vec::with_capacity(left.len() + right.len());
186 result.extend(left.iter().cloned());
187 for mut mono in right.iter().cloned() {
188 mono.coefficient = -mono.coefficient;
189 result.push(mono);
190 }
191 combine_like_terms(result)
192 }
193 SymbolicExpressionNode::Mul {
194 left_idx,
195 right_idx,
196 ..
197 } => {
198 let left = expand_node_cached(dag, *left_idx, cache);
199 let right = expand_node_cached(dag, *right_idx, cache);
200 let mut result = Vec::with_capacity(left.len() * right.len());
201 for l in left.iter() {
202 for r in right.iter() {
203 let mut vars = l.variables.clone();
204 vars.extend(&r.variables);
205 vars.sort();
206 result.push(ExpandedMonomial {
207 variables: vars,
208 coefficient: l.coefficient * r.coefficient,
209 });
210 }
211 }
212 combine_like_terms(result)
213 }
214 SymbolicExpressionNode::Neg { idx, .. } => expand_node_cached(dag, *idx, cache)
215 .iter()
216 .cloned()
217 .map(|mut mono| {
218 mono.coefficient = -mono.coefficient;
219 mono
220 })
221 .collect(),
222 };
223
224 let result = Arc::from(result_vec);
225 cache.insert(idx, Arc::clone(&result));
226 result
227}
228
229fn expand_leaf<F: Field>(variables: Vec<PackedVar>, coefficient: F) -> Vec<ExpandedMonomial<F>> {
230 vec![ExpandedMonomial {
231 variables,
232 coefficient,
233 }]
234}
235
236fn serialize_monomials<TermIn, TermOut, FSort, FMap>(
237 monomials: impl IntoIterator<Item = (Vec<PackedVar>, Vec<TermIn>)>,
238 mut sort_terms: FSort,
239 mut map_term: FMap,
240) -> (Vec<MonomialHeader>, Vec<PackedVar>, Vec<TermOut>)
241where
242 FSort: FnMut(&mut Vec<TermIn>),
243 FMap: FnMut(TermIn) -> TermOut,
244{
245 let iter = monomials.into_iter();
246 let (min, _) = iter.size_hint();
247 let mut headers = Vec::with_capacity(min);
248 let mut all_vars = Vec::new();
249 let mut all_terms = Vec::new();
250
251 for (vars, mut terms) in iter {
252 sort_terms(&mut terms);
253 assert!(
254 vars.len() <= u16::MAX as usize,
255 "monomial has too many variables for PackedVar header"
256 );
257 assert!(
258 terms.len() <= u16::MAX as usize,
259 "monomial has too many terms for PackedVar header"
260 );
261 headers.push(MonomialHeader {
262 var_offset: all_vars.len() as u32,
263 num_vars: vars.len() as u16,
264 term_offset: all_terms.len() as u32,
265 num_terms: terms.len() as u16,
266 });
267 all_vars.extend(vars);
268 all_terms.extend(terms.into_iter().map(&mut map_term));
269 }
270
271 (headers, all_vars, all_terms)
272}
273
274fn combine_like_terms<F: Field>(monomials: Vec<ExpandedMonomial<F>>) -> Vec<ExpandedMonomial<F>> {
275 let mut map: FxHashMap<Vec<PackedVar>, F> = FxHashMap::default();
276 for mono in monomials {
277 *map.entry(mono.variables).or_insert(F::ZERO) += mono.coefficient;
278 }
279 map.into_iter()
280 .filter(|(_, coeff)| *coeff != F::ZERO)
281 .map(|(variables, coefficient)| ExpandedMonomial {
282 variables,
283 coefficient,
284 })
285 .collect()
286}
287
288#[derive(Clone, Debug, PartialEq, Eq)]
294pub struct ExpandedInteractionMonomials<F> {
295 pub numer_headers: Vec<MonomialHeader>,
297 pub numer_variables: Vec<PackedVar>,
299 pub numer_terms: Vec<InteractionMonomialTerm<F>>,
301 pub denom_headers: Vec<MonomialHeader>,
303 pub denom_variables: Vec<PackedVar>,
305 pub denom_terms: Vec<InteractionMonomialTerm<F>>,
307 pub num_interactions: u32,
309 pub max_fields_len: usize,
311}
312
313impl<F: Field> ExpandedInteractionMonomials<F> {
314 pub fn from_symbolic_constraints(symbolic: &SymbolicConstraints<F>) -> Self {
316 let interactions = &symbolic.interactions;
317 if interactions.is_empty() {
318 return Self {
319 numer_headers: Vec::new(),
320 numer_variables: Vec::new(),
321 numer_terms: Vec::new(),
322 denom_headers: Vec::new(),
323 denom_variables: Vec::new(),
324 denom_terms: Vec::new(),
325 num_interactions: 0,
326 max_fields_len: 0,
327 };
328 }
329
330 let max_fields_len = interactions
331 .iter()
332 .map(|i| i.message.len())
333 .max()
334 .unwrap_or(0);
335
336 let mut numer_map: FxHashMap<Vec<PackedVar>, Vec<(u16, F)>> = FxHashMap::default();
339 for (interaction_idx, interaction) in interactions.iter().enumerate() {
340 let monomials = expand_symbolic_expression(&interaction.count);
341 for mono in monomials {
342 if mono.coefficient == F::ZERO {
343 continue;
344 }
345 numer_map
346 .entry(mono.variables)
347 .or_default()
348 .push((interaction_idx as u16, mono.coefficient));
349 }
350 }
351
352 let mut denom_map: FxHashMap<Vec<PackedVar>, Vec<(u16, u16, F)>> = FxHashMap::default();
356 for (interaction_idx, interaction) in interactions.iter().enumerate() {
357 for (field_idx, field_expr) in interaction.message.iter().enumerate() {
358 let monomials = expand_symbolic_expression(field_expr);
359 for mono in monomials {
360 if mono.coefficient == F::ZERO {
361 continue;
362 }
363 denom_map.entry(mono.variables).or_default().push((
364 interaction_idx as u16,
365 field_idx as u16,
366 mono.coefficient,
367 ));
368 }
369 }
370 }
371
372 let (numer_headers, numer_variables, numer_terms) = serialize_monomials(
373 numer_map,
374 |_| {},
375 |(idx, coeff)| InteractionMonomialTerm {
376 interaction_idx: idx,
377 field_idx: 0, coefficient: coeff,
379 },
380 );
381
382 let (denom_headers, denom_variables, denom_terms) = serialize_monomials(
383 denom_map,
384 |_| {},
385 |(int_idx, field_idx, coeff)| InteractionMonomialTerm {
386 interaction_idx: int_idx,
387 field_idx,
388 coefficient: coeff,
389 },
390 );
391
392 Self {
393 numer_headers,
394 numer_variables,
395 numer_terms,
396 denom_headers,
397 denom_variables,
398 denom_terms,
399 num_interactions: interactions.len() as u32,
400 max_fields_len,
401 }
402 }
403}
404
405fn expand_symbolic_expression<F: Field>(expr: &SymbolicExpression<F>) -> Vec<ExpandedMonomial<F>> {
407 match expr {
408 SymbolicExpression::Constant(c) => expand_leaf(vec![], *c),
409 SymbolicExpression::Variable(v) => {
410 expand_leaf(vec![PackedVar::from_symbolic_var(v)], F::ONE)
411 }
412 SymbolicExpression::IsFirstRow => expand_leaf(vec![PackedVar::is_first()], F::ONE),
413 SymbolicExpression::IsLastRow => expand_leaf(vec![PackedVar::is_last()], F::ONE),
414 SymbolicExpression::IsTransition => expand_leaf(vec![PackedVar::is_transition()], F::ONE),
415 SymbolicExpression::Add { x, y, .. } => {
416 let mut result = expand_symbolic_expression(x);
417 result.extend(expand_symbolic_expression(y));
418 combine_like_terms(result)
419 }
420 SymbolicExpression::Sub { x, y, .. } => {
421 let mut result = expand_symbolic_expression(x);
422 for mut mono in expand_symbolic_expression(y) {
423 mono.coefficient = -mono.coefficient;
424 result.push(mono);
425 }
426 combine_like_terms(result)
427 }
428 SymbolicExpression::Mul { x, y, .. } => {
429 let left_monomials = expand_symbolic_expression(x);
430 let right_monomials = expand_symbolic_expression(y);
431 let mut result = Vec::with_capacity(left_monomials.len() * right_monomials.len());
432 for l in &left_monomials {
433 for r in &right_monomials {
434 let mut vars = l.variables.clone();
435 vars.extend(&r.variables);
436 vars.sort();
437 result.push(ExpandedMonomial {
438 variables: vars,
439 coefficient: l.coefficient * r.coefficient,
440 });
441 }
442 }
443 combine_like_terms(result)
444 }
445 SymbolicExpression::Neg { x, .. } => expand_symbolic_expression(x)
446 .into_iter()
447 .map(|mut mono| {
448 mono.coefficient = -mono.coefficient;
449 mono
450 })
451 .collect(),
452 }
453}