Skip to main content

uniform_runner/
uniform_runner.rs

1//! Uniform-shape proving benchmark.
2//!
3//! Creates N identical synthetic AIRs sized via CLI flags, with:
4//! - Zero witness (all trace cells are 0)
5//! - Boolean constraints: `x * (x - 1) = 0`
6//! - Self-canceling bus interactions: `bus_send([x])` + `bus_receive([x])`
7//!
8//! Usage (CPU):
9//!   cargo run -p openvm-benchmark-synthetic --release --bin uniform_runner -- \
10//!     --num-airs 4 --cols-per-air 100 --constraints-per-col 2 --log-rows-per-air 20
11//!
12//! Usage (GPU):
13//!   cargo run -p openvm-benchmark-synthetic --release --features cuda --bin uniform_runner -- \
14//!     --num-airs 4 --cols-per-air 100 --constraints-per-col 2 --log-rows-per-air 20
15//!
16//! To also write a metrics.json file:
17//!   METRICS_OUTPUT=metrics.json cargo run -p openvm-benchmark-synthetic --release \
18//!     --bin uniform_runner -- ...
19
20use 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    /// Number of AIRs
46    #[arg(long, default_value_t = 1)]
47    num_airs: usize,
48
49    /// Number of columns per AIR
50    #[arg(long, default_value_t = 20)]
51    cols_per_air: usize,
52
53    /// Boolean constraints per column (can be fractional, e.g. 0.5 means every
54    /// other column gets a constraint)
55    #[arg(long, default_value_t = 1.0)]
56    constraints_per_col: f64,
57
58    /// Send/receive bus interaction pairs per column (can be fractional, e.g.
59    /// 0.25 means a pair on every 8th column). Each pair creates one send and
60    /// one receive with the same message, so they cancel out.
61    #[arg(long, default_value_t = 0.25)]
62    interactions_per_col: f64,
63
64    /// Log2 of the number of rows per AIR.
65    #[arg(long, default_value_t = 18)]
66    log_rows_per_air: usize,
67
68    /// Log2 of stacked height for the PCS. Default is 24
69    /// (MAX_APP_LOG_STACKED_HEIGHT), matching default_app_config() in the
70    /// openvm cli crate.
71    #[arg(long, default_value_t = 24)]
72    log_stacked_height: usize,
73}
74
75// ---------------------------------------------------------------------------
76// BenchmarkAir
77// ---------------------------------------------------------------------------
78
79#[derive(Clone, Debug)]
80pub struct BenchmarkAir {
81    pub num_columns: usize,
82    /// Total boolean constraints in this AIR (spread across columns round-robin).
83    pub total_constraints: usize,
84    /// Total send+receive pairs in this AIR (spread across columns round-robin).
85    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        // Boolean constraints spread round-robin across columns
102        for i in 0..self.total_constraints {
103            builder.assert_bool(local[i % self.num_columns]);
104        }
105
106        // Self-canceling bus interaction pairs spread round-robin across columns
107        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// ---------------------------------------------------------------------------
116// Backend-specific proving
117// ---------------------------------------------------------------------------
118
119#[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
176// ---------------------------------------------------------------------------
177// main
178// ---------------------------------------------------------------------------
179
180fn main() {
181    let args = Args::parse();
182
183    // run_with_metric_collection sets up tracing + metrics recording.
184    // If METRICS_OUTPUT is set, writes a metrics.json on completion.
185    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    // Compute actual integer counts per AIR from fractional rates
196    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    // Create AIRs
240    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    // Keygen (always CPU)
253    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    // Generate zero traces on CPU
260    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    // Prove (backend-specific)
270    println!("Proving...");
271    let start = Instant::now();
272    let proof = prove(&params, &pk, cpu_traces);
273    let prove_time = start.elapsed();
274    println!("Done proving, time: {prove_time:?}");
275
276    // Verify (always CPU)
277    println!("Verifying...");
278    keygen_engine.verify(&vk, &proof).unwrap();
279    println!("Verifies!");
280}