1use 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 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
36pub struct InteractionEvalRules {
38 pub(crate) inner: EvalRules,
39 pub(crate) d_pair_idxs: DeviceBuffer<u32>,
44 pub(crate) max_fields_len: usize,
45}
46
47pub struct ConstraintOnlyRules<const BUFFER_VARS: bool> {
49 pub(crate) inner: EvalRules,
50}
51
52pub struct EvalRules {
53 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 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 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 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 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 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 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 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 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}