1use std::{sync::Arc, time::Instant};
21
22use clap::Parser;
23use openvm_stark_backend::{
24 interaction::InteractionBuilder,
25 keygen::types::MultiStarkProvingKey,
26 proof::Proof,
27 prover::{AirProvingContext, ColMajorMatrix, DeviceDataTransporter, ProvingContext},
28 AirRef, PartitionedBaseAir, StarkEngine,
29};
30use openvm_stark_sdk::config::{
31 app_params_with_100_bits_security,
32 baby_bear_poseidon2::{BabyBearPoseidon2Config, BabyBearPoseidon2RefEngine},
33};
34use p3_air::{Air, AirBuilder, BaseAir, BaseAirWithPublicValues};
35use p3_baby_bear::BabyBear;
36use p3_field::PrimeCharacteristicRing;
37use p3_matrix::Matrix;
38
39type F = BabyBear;
40type SC = BabyBearPoseidon2Config;
41
42#[derive(Parser)]
43#[command(about = "Synthetic proving benchmark")]
44struct Args {
45 #[arg(long, default_value_t = 1)]
47 num_airs: usize,
48
49 #[arg(long, default_value_t = 20)]
51 cols_per_air: usize,
52
53 #[arg(long, default_value_t = 1.0)]
56 constraints_per_col: f64,
57
58 #[arg(long, default_value_t = 0.25)]
62 interactions_per_col: f64,
63
64 #[arg(long, default_value_t = 18)]
66 log_rows_per_air: usize,
67
68 #[arg(long, default_value_t = 24)]
72 log_stacked_height: usize,
73}
74
75#[derive(Clone, Debug)]
80pub struct BenchmarkAir {
81 pub num_columns: usize,
82 pub total_constraints: usize,
84 pub total_interaction_pairs: usize,
86}
87
88impl<F> BaseAir<F> for BenchmarkAir {
89 fn width(&self) -> usize {
90 self.num_columns
91 }
92}
93impl<F> BaseAirWithPublicValues<F> for BenchmarkAir {}
94impl<F> PartitionedBaseAir<F> for BenchmarkAir {}
95
96impl<AB: AirBuilder + InteractionBuilder> Air<AB> for BenchmarkAir {
97 fn eval(&self, builder: &mut AB) {
98 let main = builder.main();
99 let local = main.row_slice(0).unwrap();
100
101 for i in 0..self.total_constraints {
103 builder.assert_bool(local[i % self.num_columns]);
104 }
105
106 for i in 0..self.total_interaction_pairs {
108 let field = vec![local[i % self.num_columns]];
109 builder.push_interaction(0, field.clone(), AB::Expr::ONE, 0);
110 builder.push_interaction(0, field, AB::Expr::NEG_ONE, 0);
111 }
112 }
113}
114
115#[cfg(not(feature = "cuda"))]
120fn prove(
121 params: &openvm_stark_backend::SystemParams,
122 pk: &MultiStarkProvingKey<SC>,
123 cpu_traces: Vec<ColMajorMatrix<F>>,
124) -> Proof<SC> {
125 let engine: BabyBearPoseidon2RefEngine = StarkEngine::new(params.clone());
126 let d_pk = engine.device().transport_pk_to_device(pk);
127 let ctx = ProvingContext::new(
128 cpu_traces
129 .into_iter()
130 .enumerate()
131 .map(|(i, trace)| (i, AirProvingContext::simple_no_pis(trace)))
132 .collect(),
133 );
134 engine.prove(&d_pk, ctx).unwrap()
135}
136
137#[cfg(feature = "cuda")]
138fn prove(
139 params: &openvm_stark_backend::SystemParams,
140 pk: &MultiStarkProvingKey<SC>,
141 cpu_traces: Vec<ColMajorMatrix<F>>,
142) -> Proof<SC> {
143 use openvm_cuda_backend::{prelude::SC as CudaSC, BabyBearPoseidon2GpuEngine, GpuBackend};
144
145 let engine = BabyBearPoseidon2GpuEngine::new(params.clone());
146 let device = engine.device();
147 let d_pk = <_ as DeviceDataTransporter<CudaSC, GpuBackend>>::transport_pk_to_device(device, pk);
148 let ctx = ProvingContext::new(
149 cpu_traces
150 .iter()
151 .enumerate()
152 .map(|(i, trace)| {
153 let d_trace =
154 <_ as DeviceDataTransporter<CudaSC, GpuBackend>>::transport_matrix_to_device(
155 device, trace,
156 );
157 (i, AirProvingContext::simple_no_pis(d_trace))
158 })
159 .collect(),
160 );
161 engine.prove(&d_pk, ctx).unwrap()
162}
163
164fn fmt_num(n: usize) -> String {
165 let s = n.to_string();
166 let mut result = String::with_capacity(s.len() + s.len() / 3);
167 for (i, c) in s.chars().enumerate() {
168 if i > 0 && (s.len() - i).is_multiple_of(3) {
169 result.push(',');
170 }
171 result.push(c);
172 }
173 result
174}
175
176fn main() {
181 let args = Args::parse();
182
183 openvm_stark_sdk::bench::run_with_metric_collection("METRICS_OUTPUT", || run(&args));
186}
187
188fn run(args: &Args) {
189 assert!(args.num_airs > 0);
190 assert!(args.cols_per_air > 0);
191
192 let trace_height = 1usize << args.log_rows_per_air;
193 let total_cells = args.num_airs * args.cols_per_air * trace_height;
194
195 let constraints_per_air =
197 (args.constraints_per_col * args.cols_per_air as f64).round() as usize;
198 let interaction_pairs_per_air =
199 (args.interactions_per_col * args.cols_per_air as f64 / 2.0).round() as usize;
200
201 let total_constraints = constraints_per_air * args.num_airs;
202 let total_bus_interactions = interaction_pairs_per_air * 2 * args.num_airs;
203 let constraint_instances = total_constraints * trace_height;
204 let bus_interaction_messages = total_bus_interactions * trace_height;
205
206 let backend_name = if cfg!(feature = "cuda") { "GPU" } else { "CPU" };
207
208 println!("=== Synthetic Proving Benchmark ({backend_name}) ===");
209 println!(" num_airs: {}", fmt_num(args.num_airs));
210 println!(" cols_per_air: {}", fmt_num(args.cols_per_air));
211 println!(" constraints_per_col: {}", args.constraints_per_col);
212 println!(" interactions_per_col: {}", args.interactions_per_col);
213 println!(
214 " trace_height: {} (2^{})",
215 fmt_num(trace_height),
216 args.log_rows_per_air
217 );
218 println!(
219 " trace_cells: {} (2^{} rows * {} AIRs * {} columns / AIR)",
220 fmt_num(total_cells),
221 args.log_rows_per_air,
222 args.num_airs,
223 args.cols_per_air
224 );
225 println!(" constraints: {}", fmt_num(total_constraints));
226 println!(
227 " bus_interactions: {}",
228 fmt_num(total_bus_interactions)
229 );
230 println!(
231 " constraint_instances: {}",
232 fmt_num(constraint_instances)
233 );
234 println!(
235 " bus_interaction_msgs: {}",
236 fmt_num(bus_interaction_messages)
237 );
238
239 let airs: Vec<AirRef<SC>> = (0..args.num_airs)
241 .map(|_| {
242 Arc::new(BenchmarkAir {
243 num_columns: args.cols_per_air,
244 total_constraints: constraints_per_air,
245 total_interaction_pairs: interaction_pairs_per_air,
246 }) as AirRef<SC>
247 })
248 .collect();
249
250 let params = app_params_with_100_bits_security(args.log_stacked_height);
251
252 let keygen_engine: BabyBearPoseidon2RefEngine = StarkEngine::new(params.clone());
254 println!("\nKeygen...");
255 let start = Instant::now();
256 let (pk, vk) = keygen_engine.keygen(&airs);
257 println!(" time: {:?}", start.elapsed());
258
259 let cpu_traces: Vec<ColMajorMatrix<F>> = (0..args.num_airs)
261 .map(|_| {
262 ColMajorMatrix::new(
263 vec![F::ZERO; trace_height * args.cols_per_air],
264 args.cols_per_air,
265 )
266 })
267 .collect();
268
269 println!("Proving...");
271 let start = Instant::now();
272 let proof = prove(¶ms, &pk, cpu_traces);
273 let prove_time = start.elapsed();
274 println!("Done proving, time: {prove_time:?}");
275
276 println!("Verifying...");
278 keygen_engine.verify(&vk, &proof).unwrap();
279 println!("Verifies!");
280}