Skip to content

Commit fba57ba

Browse files
kunxian-xiaclaude
andcommitted
feat: add jagged vs direct PCS comparison bench and use p3 RowMajorMatrix
- Add comparison benchmark (comparison.rs) measuring commit, batch_open, batch_verify, and proof size for jagged PCS vs direct inner PCS - Refactor jagged_commit to accept p3::matrix::dense::RowMajorMatrix instead of witness::RowMajorMatrix, supporting non-power-of-two matrix heights (internally padded to next power of two before bit-reversal) - Update jagged_pcs bench to use BabyBearExt4, parallel make_rmm, jittered non-power-of-two heights, and eq-table-based column evaluation - Cap reshape_log_height to 25 to fit BabyBear two-adicity constraint Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
1 parent c530e7f commit fba57ba

4 files changed

Lines changed: 392 additions & 50 deletions

File tree

crates/mpcs/Cargo.toml

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -62,3 +62,7 @@ name = "jagged_sumcheck"
6262
[[bench]]
6363
harness = false
6464
name = "jagged_pcs"
65+
66+
[[bench]]
67+
harness = false
68+
name = "comparison"

crates/mpcs/benches/comparison.rs

Lines changed: 329 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,329 @@
1+
use std::time::Duration;
2+
3+
use criterion::*;
4+
use ff_ext::{BabyBearExt4, FromUniformBytes};
5+
use mpcs::{
6+
Basefold, BasefoldRSParams, PolynomialCommitmentScheme, SecurityLevel, jagged_batch_open,
7+
jagged_batch_verify, jagged_commit,
8+
};
9+
use multilinear_extensions::{util::ceil_log2, virtual_poly::build_eq_x_r_vec_sequential};
10+
use p3::{
11+
field::FieldAlgebra,
12+
babybear::BabyBear,
13+
matrix::{Matrix, dense::RowMajorMatrix},
14+
maybe_rayon::prelude::*,
15+
};
16+
use rand::{Rng, thread_rng};
17+
use transcript::{BasicTranscript, Transcript};
18+
use witness::{InstancePaddingStrategy, RowMajorMatrix as WitnessRowMajorMatrix};
19+
20+
type E = BabyBearExt4;
21+
type F = BabyBear;
22+
type Pcs = Basefold<E, BasefoldRSParams>;
23+
24+
const NUM_SAMPLES: usize = 10;
25+
const NUM_MATRICES: usize = 35;
26+
const NUM_COLS: usize = 32;
27+
28+
fn make_rmm(num_rows: usize, num_cols: usize) -> RowMajorMatrix<F> {
29+
let values: Vec<F> = (0..num_rows * num_cols)
30+
.into_par_iter()
31+
.map(|i| F::from_canonical_u32(((i as u64 * 13 + 7) % (1 << 30)) as u32))
32+
.collect();
33+
RowMajorMatrix::new(values, num_cols)
34+
}
35+
36+
fn sample_heights(rng: &mut impl Rng, num_matrices: usize) -> Vec<usize> {
37+
(0..num_matrices)
38+
.map(|_| {
39+
let log = rng.gen_range(16u32..=22);
40+
let base = 1usize << log;
41+
let lo = base - base / 4;
42+
let hi = base + base / 4;
43+
rng.gen_range(lo..=hi)
44+
})
45+
.collect()
46+
}
47+
48+
fn eval_all_columns_at_point(rmm: &RowMajorMatrix<F>, point: &[E]) -> Vec<E> {
49+
let w = rmm.width();
50+
let eq = build_eq_x_r_vec_sequential(point);
51+
let mut col_evals = vec![E::ZERO; w];
52+
for (eq_r, row) in eq.iter().zip(rmm.rows()) {
53+
for (col_eval, val) in col_evals.iter_mut().zip(row) {
54+
*col_eval += *eq_r * val;
55+
}
56+
}
57+
col_evals
58+
}
59+
60+
fn bench_comparison(c: &mut Criterion) {
61+
let mut group = c.benchmark_group("jagged_vs_direct");
62+
group.sample_size(NUM_SAMPLES);
63+
64+
let mut rng = thread_rng();
65+
let heights = sample_heights(&mut rng, NUM_MATRICES);
66+
let log_heights: Vec<usize> = heights.iter().map(|h| ceil_log2(*h)).collect();
67+
let max_s = *log_heights.iter().max().unwrap();
68+
69+
println!("Matrix heights: {:?}", heights);
70+
71+
let rmms: Vec<_> = heights.iter().map(|&h| make_rmm(h, NUM_COLS)).collect();
72+
let total_evals: usize = rmms.iter().map(|rmm| rmm.height() * rmm.width()).sum();
73+
let num_giga_vars = ceil_log2(total_evals);
74+
75+
println!(
76+
"num_matrices={NUM_MATRICES}, num_cols={NUM_COLS}, \
77+
total_evals={total_evals}, num_giga_vars={num_giga_vars}, max_s={max_s}"
78+
);
79+
80+
let point: Vec<E> = (0..max_s).map(|_| E::random(&mut rng)).collect();
81+
82+
let evals: Vec<E> = rmms
83+
.iter()
84+
.zip(log_heights.iter())
85+
.flat_map(|(rmm, &s_i)| eval_all_columns_at_point(rmm, &point[(max_s - s_i)..]))
86+
.collect();
87+
88+
// Per-matrix points and evals (used by direct batch_open).
89+
let per_matrix_point_evals: Vec<(Vec<E>, Vec<E>)> = log_heights
90+
.iter()
91+
.enumerate()
92+
.map(|(i, &s_i)| {
93+
let matrix_point = point[(max_s - s_i)..].to_vec();
94+
let matrix_evals = evals[i * NUM_COLS..(i + 1) * NUM_COLS].to_vec();
95+
(matrix_point, matrix_evals)
96+
})
97+
.collect();
98+
99+
// ======================== Jagged PCS ========================
100+
// BabyBear two-adicity is 27; RS rate_log=1 needs level+1 ≤ 27, so max poly_size is 2^25.
101+
let reshape_log_height = num_giga_vars.saturating_sub(4).min(25);
102+
let jagged_poly_size = 1usize << reshape_log_height;
103+
let jagged_param = Pcs::setup(jagged_poly_size, SecurityLevel::Conjecture100bits).unwrap();
104+
let (jagged_pp, jagged_vp) = Pcs::trim(jagged_param, jagged_poly_size).unwrap();
105+
106+
group.bench_function("jagged/commit", |b| {
107+
b.iter_custom(|iters| {
108+
let mut time = Duration::new(0, 0);
109+
for _ in 0..iters {
110+
let rmms_clone = rmms.clone();
111+
let instant = std::time::Instant::now();
112+
let _ =
113+
jagged_commit::<E, Pcs>(&jagged_pp, rmms_clone, reshape_log_height).unwrap();
114+
time += instant.elapsed();
115+
}
116+
time
117+
})
118+
});
119+
120+
let t0 = std::time::Instant::now();
121+
let jagged_comm =
122+
jagged_commit::<E, Pcs>(&jagged_pp, rmms.clone(), reshape_log_height).unwrap();
123+
println!("jagged_commit: {:?}", t0.elapsed());
124+
let jagged_pure_comm = jagged_comm.to_commitment();
125+
126+
group.bench_function("jagged/batch_open", |b| {
127+
b.iter_batched(
128+
|| {
129+
let mut t = BasicTranscript::<E>::new(b"bench");
130+
Pcs::write_commitment(&jagged_pure_comm.inner, &mut t).unwrap();
131+
t
132+
},
133+
|mut t| {
134+
jagged_batch_open::<E, Pcs>(&jagged_pp, &jagged_comm, &point, &evals, &mut t)
135+
.unwrap();
136+
},
137+
BatchSize::SmallInput,
138+
);
139+
});
140+
141+
let jagged_proof = {
142+
let mut t = BasicTranscript::<E>::new(b"bench");
143+
Pcs::write_commitment(&jagged_pure_comm.inner, &mut t).unwrap();
144+
let t0 = std::time::Instant::now();
145+
let proof =
146+
jagged_batch_open::<E, Pcs>(&jagged_pp, &jagged_comm, &point, &evals, &mut t).unwrap();
147+
println!("jagged_batch_open: {:?}", t0.elapsed());
148+
proof
149+
};
150+
let jagged_proof_size = bincode::serialize(&jagged_proof)
151+
.map(|v| v.len())
152+
.unwrap_or(0);
153+
154+
{
155+
let mut t = BasicTranscript::<E>::new(b"bench");
156+
Pcs::write_commitment(&jagged_pure_comm.inner, &mut t).unwrap();
157+
let t0 = std::time::Instant::now();
158+
jagged_batch_verify::<E, Pcs>(
159+
&jagged_vp,
160+
&jagged_pure_comm,
161+
&point,
162+
&evals,
163+
&jagged_proof,
164+
&mut t,
165+
)
166+
.unwrap();
167+
println!("jagged_batch_verify: {:?}", t0.elapsed());
168+
}
169+
170+
group.bench_function("jagged/batch_verify", |b| {
171+
b.iter_batched(
172+
|| {
173+
let mut t = BasicTranscript::<E>::new(b"bench");
174+
Pcs::write_commitment(&jagged_pure_comm.inner, &mut t).unwrap();
175+
t
176+
},
177+
|mut t| {
178+
jagged_batch_verify::<E, Pcs>(
179+
&jagged_vp,
180+
&jagged_pure_comm,
181+
&point,
182+
&evals,
183+
&jagged_proof,
184+
&mut t,
185+
)
186+
.unwrap();
187+
},
188+
BatchSize::SmallInput,
189+
);
190+
});
191+
192+
// ======================== Direct Inner PCS ========================
193+
// Pcs::batch_commit expects witness::RowMajorMatrix, so convert.
194+
let to_witness = |rmms: &[RowMajorMatrix<F>]| -> Vec<WitnessRowMajorMatrix<F>> {
195+
rmms.iter()
196+
.map(|rmm| {
197+
WitnessRowMajorMatrix::new_by_values(
198+
rmm.values.clone(),
199+
rmm.width(),
200+
InstancePaddingStrategy::Default,
201+
)
202+
})
203+
.collect()
204+
};
205+
206+
let direct_poly_size = 1usize << max_s;
207+
let direct_param = Pcs::setup(direct_poly_size, SecurityLevel::Conjecture100bits).unwrap();
208+
let (direct_pp, direct_vp) = Pcs::trim(direct_param, direct_poly_size).unwrap();
209+
210+
group.bench_function("direct/commit", |b| {
211+
b.iter_custom(|iters| {
212+
let mut time = Duration::new(0, 0);
213+
for _ in 0..iters {
214+
let w_rmms = to_witness(&rmms);
215+
let instant = std::time::Instant::now();
216+
let _ = Pcs::batch_commit(&direct_pp, w_rmms).unwrap();
217+
time += instant.elapsed();
218+
}
219+
time
220+
})
221+
});
222+
223+
let t0 = std::time::Instant::now();
224+
let direct_comm = Pcs::batch_commit(&direct_pp, to_witness(&rmms)).unwrap();
225+
println!("direct_commit: {:?}", t0.elapsed());
226+
let direct_pure_comm = Pcs::get_pure_commitment(&direct_comm);
227+
228+
let make_direct_transcript = |comm: &<Pcs as PolynomialCommitmentScheme<E>>::Commitment| {
229+
let mut t = BasicTranscript::<E>::new(b"bench");
230+
Pcs::write_commitment(comm, &mut t).unwrap();
231+
for (_, matrix_evals) in &per_matrix_point_evals {
232+
t.append_field_element_exts(matrix_evals);
233+
}
234+
t
235+
};
236+
237+
group.bench_function("direct/batch_open", |b| {
238+
b.iter_batched(
239+
|| make_direct_transcript(&direct_pure_comm),
240+
|mut t| {
241+
Pcs::batch_open(
242+
&direct_pp,
243+
vec![(&direct_comm, per_matrix_point_evals.clone())],
244+
&mut t,
245+
)
246+
.unwrap();
247+
},
248+
BatchSize::SmallInput,
249+
);
250+
});
251+
252+
let direct_proof = {
253+
let mut t = make_direct_transcript(&direct_pure_comm);
254+
let t0 = std::time::Instant::now();
255+
let proof = Pcs::batch_open(
256+
&direct_pp,
257+
vec![(&direct_comm, per_matrix_point_evals.clone())],
258+
&mut t,
259+
)
260+
.unwrap();
261+
println!("direct_batch_open: {:?}", t0.elapsed());
262+
proof
263+
};
264+
let direct_proof_size = bincode::serialize(&direct_proof)
265+
.map(|v| v.len())
266+
.unwrap_or(0);
267+
268+
let direct_verify_rounds: Vec<_> = log_heights
269+
.iter()
270+
.enumerate()
271+
.map(|(i, &s_i)| {
272+
let matrix_point = point[(max_s - s_i)..].to_vec();
273+
let matrix_evals = evals[i * NUM_COLS..(i + 1) * NUM_COLS].to_vec();
274+
(s_i, (matrix_point, matrix_evals))
275+
})
276+
.collect();
277+
278+
{
279+
let mut t = make_direct_transcript(&direct_pure_comm);
280+
let t0 = std::time::Instant::now();
281+
Pcs::batch_verify(
282+
&direct_vp,
283+
vec![(direct_pure_comm.clone(), direct_verify_rounds.clone())],
284+
&direct_proof,
285+
&mut t,
286+
)
287+
.unwrap();
288+
println!("direct_batch_verify: {:?}", t0.elapsed());
289+
}
290+
291+
group.bench_function("direct/batch_verify", |b| {
292+
b.iter_batched(
293+
|| make_direct_transcript(&direct_pure_comm),
294+
|mut t| {
295+
Pcs::batch_verify(
296+
&direct_vp,
297+
vec![(direct_pure_comm.clone(), direct_verify_rounds.clone())],
298+
&direct_proof,
299+
&mut t,
300+
)
301+
.unwrap();
302+
},
303+
BatchSize::SmallInput,
304+
);
305+
});
306+
307+
group.finish();
308+
309+
println!("\n=== Proof Size Comparison ===");
310+
println!(
311+
"Jagged PCS: {jagged_proof_size:>10} bytes ({:.1} KB)",
312+
jagged_proof_size as f64 / 1024.0
313+
);
314+
println!(
315+
"Direct PCS: {direct_proof_size:>10} bytes ({:.1} KB)",
316+
direct_proof_size as f64 / 1024.0
317+
);
318+
println!(
319+
"Ratio (direct / jagged): {:.2}x",
320+
direct_proof_size as f64 / jagged_proof_size as f64
321+
);
322+
}
323+
324+
criterion_group! {
325+
name = benches;
326+
config = Criterion::default().warm_up_time(Duration::from_millis(3000));
327+
targets = bench_comparison,
328+
}
329+
criterion_main!(benches);

0 commit comments

Comments
 (0)