Skip to content

Commit e8f8f5c

Browse files
authored
Feat: jagged pcs (#40)
1 parent d63b38e commit e8f8f5c

12 files changed

Lines changed: 3762 additions & 4 deletions

File tree

Cargo.lock

Lines changed: 13 additions & 1 deletion
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

crates/mpcs/Cargo.toml

Lines changed: 14 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -20,7 +20,6 @@ num-integer = "0.1"
2020
p3.workspace = true
2121
rand.workspace = true
2222
rand_chacha.workspace = true
23-
rayon = { workspace = true, optional = true }
2423
serde.workspace = true
2524
sumcheck.workspace = true
2625
tracing.workspace = true
@@ -30,10 +29,11 @@ witness.workspace = true
3029

3130
[dev-dependencies]
3231
criterion.workspace = true
32+
tracing-forest = "0.1"
3333

3434
[features]
3535
nightly-features = ["ff_ext/nightly-features"]
36-
parallel = ["p3/parallel", "dep:rayon"]
36+
parallel = ["p3/parallel"]
3737
whir = ["dep:whir"]
3838
print-trace = ["whir/print-trace"]
3939
sanity-check = []
@@ -54,3 +54,15 @@ name = "interpolate"
5454
harness = false
5555
name = "whir"
5656
required-features = ["whir"]
57+
58+
[[bench]]
59+
harness = false
60+
name = "jagged_sumcheck"
61+
62+
[[bench]]
63+
harness = false
64+
name = "jagged_pcs"
65+
66+
[[bench]]
67+
harness = false
68+
name = "comparison"

crates/mpcs/benches/comparison.rs

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

0 commit comments

Comments
 (0)