Skip to main content

openvm_cuda_backend/
monomial.rs

1//! Monomial expansion for GPU zerocheck evaluation.
2//!
3//! This module expands symbolic constraint DAGs into monomials for efficient
4//! GPU evaluation via the monomial kernel.
5
6use 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/// Packed variable following the CUDA monomial layout.
17#[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/// A (constraint_idx, coefficient) pair in F[lambda].
71#[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/// Term mapping a monomial to its interaction context.
79/// For numerator: sum_i(coefficient_i * eq_3bs[interaction_idx_i])
80/// For denominator: sum_i(coefficient_i * beta_pows[field_idx_i] * eq_3bs[interaction_idx_i])
81#[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, // For denom: index into message for beta_pows. For numer: unused.
87}
88
89#[derive(Clone)]
90struct ExpandedMonomial<F> {
91    pub variables: Vec<PackedVar>,
92    pub coefficient: F,
93}
94
95/// Expanded monomials serialized for GPU upload.
96#[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// ============================================================================
289// INTERACTION (LOGUP) MONOMIAL EXPANSION
290// ============================================================================
291
292/// Expanded interaction monomials for GPU upload.
293#[derive(Clone, Debug, PartialEq, Eq)]
294pub struct ExpandedInteractionMonomials<F> {
295    /// Numerator monomial headers
296    pub numer_headers: Vec<MonomialHeader>,
297    /// Numerator monomial variables
298    pub numer_variables: Vec<PackedVar>,
299    /// Numerator interaction terms: (interaction_idx, 0, coefficient)
300    pub numer_terms: Vec<InteractionMonomialTerm<F>>,
301    /// Denominator monomial headers
302    pub denom_headers: Vec<MonomialHeader>,
303    /// Denominator monomial variables
304    pub denom_variables: Vec<PackedVar>,
305    /// Denominator interaction terms: (interaction_idx, field_idx, coefficient)
306    pub denom_terms: Vec<InteractionMonomialTerm<F>>,
307    /// Number of interactions
308    pub num_interactions: u32,
309    /// Maximum message length across all interactions
310    pub max_fields_len: usize,
311}
312
313impl<F: Field> ExpandedInteractionMonomials<F> {
314    /// Create expanded interaction monomials from symbolic constraints.
315    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        // Build numerator monomials: expand count expressions
337        // Group by variables -> (interaction_idx, coefficient)
338        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        // Build denominator monomials: expand each message field
353        // Each message field contributes: coeff * beta_pows[field_idx] * eq_3bs[interaction_idx]
354        // Group by variables -> (interaction_idx, field_idx, coefficient)
355        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, // unused for numerator
378                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
405/// Expand a `SymbolicExpression` into a list of monomials.
406fn 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}