Skip to main content

openvm_cuda_backend/
pkey.rs

1//! Defines the symbolic rule data to precompute and store in the GPU proving key
2use itertools::Itertools;
3use openvm_cuda_common::{
4    copy::MemCopyH2D, d_buffer::DeviceBuffer, error::MemCopyError, stream::GpuDeviceCtx,
5};
6use openvm_stark_backend::{
7    air_builders::symbolic::{
8        symbolic_expression::SymbolicExpression,
9        symbolic_variable::{Entry, SymbolicVariable},
10        SymbolicConstraints, SymbolicDagBuilder, SymbolicExpressionDag,
11    },
12    keygen::types::StarkProvingKey,
13    StarkProtocolConfig,
14};
15use p3_field::PrimeCharacteristicRing;
16
17use crate::{
18    logup_zerocheck::rules::{codec::Codec, SymbolicRulesGpu},
19    monomial::{
20        ExpandedInteractionMonomials, ExpandedMonomials, InteractionMonomialTerm, LambdaTerm,
21        MonomialHeader, PackedVar,
22    },
23    prelude::F,
24};
25
26pub struct AirDataGpu {
27    pub interaction_rules: InteractionEvalRules,
28    /// Whether to buffer vars depends on the performance and memory access patterns of the kernel.
29    /// This may be tuned.
30    pub zerocheck_round0: ConstraintOnlyRules<true>,
31    pub zerocheck_mle: ConstraintOnlyRules<false>,
32    pub zerocheck_monomials: Option<ZerocheckMonomials>,
33    pub interaction_monomials: Option<InteractionMonomials>,
34}
35
36/// Used for GKR input evaluation and logup MLE sumcheck rounds.
37pub struct InteractionEvalRules {
38    pub(crate) inner: EvalRules,
39    /// Constraints consist of all `(numer, denom)` pairs **topologically sorted**. We map the
40    /// constraint idx back to unsorted order. ```text
41    /// constraint_idx => 2 * interaction_idx + is_denom
42    /// ```
43    pub(crate) d_pair_idxs: DeviceBuffer<u32>,
44    pub(crate) max_fields_len: usize,
45}
46
47/// Constraints only, no interactions
48pub struct ConstraintOnlyRules<const BUFFER_VARS: bool> {
49    pub(crate) inner: EvalRules,
50}
51
52pub struct EvalRules {
53    /// Encoded rules
54    pub d_rules: DeviceBuffer<u128>,
55    pub d_used_nodes: DeviceBuffer<usize>,
56    pub buffer_size: u32,
57}
58
59pub struct ZerocheckMonomials {
60    pub d_headers: DeviceBuffer<MonomialHeader>,
61    pub d_variables: DeviceBuffer<PackedVar>,
62    pub d_lambda_terms: DeviceBuffer<LambdaTerm<F>>,
63    pub num_monomials: u32,
64}
65
66pub struct InteractionMonomials {
67    pub d_numer_headers: DeviceBuffer<MonomialHeader>,
68    pub d_numer_variables: DeviceBuffer<PackedVar>,
69    pub d_numer_terms: DeviceBuffer<InteractionMonomialTerm<F>>,
70    pub num_numer_monomials: u32,
71    pub d_denom_headers: DeviceBuffer<MonomialHeader>,
72    pub d_denom_variables: DeviceBuffer<PackedVar>,
73    pub d_denom_terms: DeviceBuffer<InteractionMonomialTerm<F>>,
74    pub num_denom_monomials: u32,
75    pub max_fields_len: usize,
76    pub num_interactions: u32,
77}
78
79fn to_device_or_empty<T>(
80    data: &[T],
81    device_ctx: &GpuDeviceCtx,
82) -> Result<DeviceBuffer<T>, MemCopyError> {
83    if data.is_empty() {
84        Ok(DeviceBuffer::new())
85    } else {
86        data.to_device_on(device_ctx)
87    }
88}
89
90impl AirDataGpu {
91    pub fn new<S: StarkProtocolConfig<F = F>>(
92        pk: &StarkProvingKey<S>,
93        device_ctx: &GpuDeviceCtx,
94    ) -> Result<Self, MemCopyError> {
95        let dag = &pk.vk.symbolic_constraints;
96        let symbolic_constraints = SymbolicConstraints::from(dag);
97        let interaction_rules = InteractionEvalRules::new(&symbolic_constraints, device_ctx)?;
98        let zerocheck_round0 = ConstraintOnlyRules::<true>::new(&dag.constraints, device_ctx)?;
99        let zerocheck_mle = ConstraintOnlyRules::<false>::new(&dag.constraints, device_ctx)?;
100
101        let zerocheck_monomials = if dag.constraints.num_constraints() > 0 {
102            let expanded = ExpandedMonomials::from_dag(&dag.constraints);
103            Some(ZerocheckMonomials::from_expanded(&expanded, device_ctx)?)
104        } else {
105            None
106        };
107        let interaction_monomials = if !symbolic_constraints.interactions.is_empty() {
108            let expanded =
109                ExpandedInteractionMonomials::from_symbolic_constraints(&symbolic_constraints);
110            Some(InteractionMonomials::from_expanded(&expanded, device_ctx)?)
111        } else {
112            None
113        };
114        Ok(Self {
115            interaction_rules,
116            zerocheck_round0,
117            zerocheck_mle,
118            zerocheck_monomials,
119            interaction_monomials,
120        })
121    }
122}
123
124impl ZerocheckMonomials {
125    pub fn from_expanded(
126        expanded: &ExpandedMonomials<F>,
127        device_ctx: &GpuDeviceCtx,
128    ) -> Result<Self, MemCopyError> {
129        // Validate bounds for all monomial headers to prevent out-of-bounds access in CUDA kernel
130        let num_variables = expanded.variables.len();
131        let num_lambda_terms = expanded.lambda_terms.len();
132        for (i, hdr) in expanded.headers.iter().enumerate() {
133            let var_end = hdr.var_offset as usize + hdr.num_vars as usize;
134            let term_end = hdr.term_offset as usize + hdr.num_terms as usize;
135            assert!(
136                var_end <= num_variables,
137                "Monomial {i}: var_offset ({}) + num_vars ({}) = {var_end} exceeds variables.len() ({num_variables})",
138                hdr.var_offset,
139                hdr.num_vars
140            );
141            assert!(
142                term_end <= num_lambda_terms,
143                "Monomial {i}: term_offset ({}) + num_terms ({}) = {term_end} exceeds lambda_terms.len() ({num_lambda_terms})",
144                hdr.term_offset,
145                hdr.num_terms
146            );
147        }
148
149        Ok(Self {
150            d_headers: expanded.headers.to_device_on(device_ctx)?,
151            d_variables: expanded.variables.to_device_on(device_ctx)?,
152            d_lambda_terms: expanded.lambda_terms.to_device_on(device_ctx)?,
153            num_monomials: expanded.headers.len() as u32,
154        })
155    }
156}
157
158impl InteractionMonomials {
159    pub fn from_expanded(
160        expanded: &ExpandedInteractionMonomials<F>,
161        device_ctx: &GpuDeviceCtx,
162    ) -> Result<Self, MemCopyError> {
163        // Validate numerator monomial headers
164        let num_numer_vars = expanded.numer_variables.len();
165        let num_numer_terms = expanded.numer_terms.len();
166        for (i, hdr) in expanded.numer_headers.iter().enumerate() {
167            let var_end = hdr.var_offset as usize + hdr.num_vars as usize;
168            let term_end = hdr.term_offset as usize + hdr.num_terms as usize;
169            assert!(
170                var_end <= num_numer_vars,
171                "Numer monomial {i}: var_offset + num_vars exceeds bounds"
172            );
173            assert!(
174                term_end <= num_numer_terms,
175                "Numer monomial {i}: term_offset + num_terms exceeds bounds"
176            );
177        }
178
179        // Validate denominator monomial headers
180        let num_denom_vars = expanded.denom_variables.len();
181        let num_denom_terms = expanded.denom_terms.len();
182        for (i, hdr) in expanded.denom_headers.iter().enumerate() {
183            let var_end = hdr.var_offset as usize + hdr.num_vars as usize;
184            let term_end = hdr.term_offset as usize + hdr.num_terms as usize;
185            assert!(
186                var_end <= num_denom_vars,
187                "Denom monomial {i}: var_offset + num_vars exceeds bounds"
188            );
189            assert!(
190                term_end <= num_denom_terms,
191                "Denom monomial {i}: term_offset + num_terms exceeds bounds"
192            );
193        }
194
195        Ok(Self {
196            d_numer_headers: to_device_or_empty(&expanded.numer_headers, device_ctx)?,
197            d_numer_variables: to_device_or_empty(&expanded.numer_variables, device_ctx)?,
198            d_numer_terms: to_device_or_empty(&expanded.numer_terms, device_ctx)?,
199            num_numer_monomials: expanded.numer_headers.len() as u32,
200            d_denom_headers: to_device_or_empty(&expanded.denom_headers, device_ctx)?,
201            d_denom_variables: to_device_or_empty(&expanded.denom_variables, device_ctx)?,
202            d_denom_terms: to_device_or_empty(&expanded.denom_terms, device_ctx)?,
203            num_denom_monomials: expanded.denom_headers.len() as u32,
204            max_fields_len: expanded.max_fields_len,
205            num_interactions: expanded.num_interactions,
206        })
207    }
208}
209
210impl InteractionEvalRules {
211    pub fn new(
212        symbolic_constraints: &SymbolicConstraints<F>,
213        device_ctx: &GpuDeviceCtx,
214    ) -> Result<Self, MemCopyError> {
215        let interactions = &symbolic_constraints.interactions;
216        let num_interactions = interactions.len();
217        if num_interactions == 0 {
218            return Ok(Self {
219                inner: EvalRules::dummy(),
220
221                max_fields_len: 0,
222                d_pair_idxs: DeviceBuffer::new(),
223            });
224        }
225        let max_fields_len = interactions
226            .iter()
227            .map(|interaction| interaction.message.len())
228            .max()
229            .unwrap_or(0);
230        // [alpha, beta^0, ..., beta^max_fields_len]
231        let symbolic_challenges: Vec<SymbolicExpression<F>> = (0..max_fields_len + 2)
232            .map(|index| SymbolicVariable::<F>::new(Entry::Challenge, index).into())
233            .collect();
234
235        let mut frac_pairs = Vec::with_capacity(num_interactions * 2);
236        for interaction in interactions.iter() {
237            let numer = interaction.count.clone();
238            let b = SymbolicExpression::from_u32(interaction.bus_index as u32 + 1);
239            let betas = symbolic_challenges[1..].to_vec();
240            let mut denom = SymbolicExpression::from_u32(0);
241            for (j, expr) in interaction.message.iter().enumerate() {
242                denom += betas[j].clone() * expr.clone();
243            }
244            denom += betas[interaction.message.len()].clone() * b;
245            frac_pairs.push(numer);
246            frac_pairs.push(denom);
247        }
248        // build DAG without sorting constraint idxs:
249        let (dag, pair_idxs) = {
250            let mut dag_builder = SymbolicDagBuilder::new();
251            let mut dag_pair_idxs: Vec<(usize, u32)> = frac_pairs
252                .iter()
253                .enumerate()
254                .map(|(pair_idx, expr)| {
255                    let dag_idx = dag_builder.add_expr(expr);
256                    (dag_idx, pair_idx.try_into().unwrap())
257                })
258                .collect_vec();
259            dag_pair_idxs.sort();
260            let (constraint_idx, pair_idxs): (Vec<_>, Vec<_>) = dag_pair_idxs.into_iter().unzip();
261            // NOTE: do not sort pair_idxs since we need to keep them in pairs
262            let dag = SymbolicExpressionDag {
263                nodes: dag_builder.nodes,
264                constraint_idx,
265            };
266            (dag, pair_idxs)
267        };
268        let rules = SymbolicRulesGpu::new(&dag, false);
269        // Build used_nodes with duplicates, preserving order from constraint_idx
270        let used_nodes = dag
271            .constraint_idx
272            .iter()
273            .map(|&dag_idx| rules.dag_idx_to_rule_idx[&dag_idx])
274            .collect_vec();
275        let encoded_rules = rules.rules.iter().map(|c| c.encode()).collect_vec();
276        let d_rules = encoded_rules.to_device_on(device_ctx)?;
277        let d_used_nodes = used_nodes.to_device_on(device_ctx)?;
278        let d_pair_idxs = pair_idxs.to_device_on(device_ctx)?;
279        assert_eq!(
280            used_nodes.len(),
281            2 * num_interactions,
282            "Rules come in (numer, denom) pairs"
283        );
284
285        let inner = EvalRules {
286            d_rules,
287            d_used_nodes,
288            buffer_size: rules
289                .buffer_size
290                .try_into()
291                .expect("buffer_size exceeds u32"),
292        };
293
294        Ok(Self {
295            inner,
296            d_pair_idxs,
297            max_fields_len,
298        })
299    }
300}
301
302impl<const BUFFER_VARS: bool> ConstraintOnlyRules<BUFFER_VARS> {
303    pub fn new(
304        dag: &SymbolicExpressionDag<F>,
305        device_ctx: &GpuDeviceCtx,
306    ) -> Result<Self, MemCopyError> {
307        if dag.num_constraints() == 0 {
308            return Ok(Self {
309                inner: EvalRules::dummy(),
310            });
311        }
312
313        let rules = SymbolicRulesGpu::new(dag, BUFFER_VARS);
314        // Build used_nodes with duplicates, preserving order from constraint_idx
315        let used_nodes = dag
316            .constraint_idx
317            .iter()
318            .map(|&dag_idx| rules.dag_idx_to_rule_idx[&dag_idx])
319            .collect_vec();
320
321        let encoded_rules = rules.rules.iter().map(|c| c.encode()).collect_vec();
322        let d_rules = encoded_rules.to_device_on(device_ctx)?;
323        let d_used_nodes = used_nodes.to_device_on(device_ctx)?;
324
325        let inner = EvalRules {
326            d_rules,
327            d_used_nodes,
328            buffer_size: rules
329                .buffer_size
330                .try_into()
331                .expect("buffer_size exceeds u32"),
332        };
333        Ok(Self { inner })
334    }
335}
336
337impl EvalRules {
338    pub fn dummy() -> Self {
339        Self {
340            d_rules: DeviceBuffer::new(),
341            d_used_nodes: DeviceBuffer::new(),
342            buffer_size: 0,
343        }
344    }
345}