Skip to content

Commit 3916c34

Browse files
authored
Expose jagged assist claim helper (#64)
* fix: align jagged q layout with committed rows * fix(mpcs): compact jagged q over occupied rows * rollback unnessesary comments * rollback unnessesary change * Fuse jagged assist claim computation * Fix jagged assist clippy warnings
1 parent 83cf5bb commit 3916c34

6 files changed

Lines changed: 97 additions & 19 deletions

File tree

crates/mpcs/src/jagged/assist.rs

Lines changed: 67 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -35,27 +35,64 @@ pub fn assist_sumcheck_prove<E: ExtensionField>(
3535
n_robp: usize,
3636
transcript: &mut impl Transcript<E>,
3737
) -> (IOPProof<E>, Vec<E>) {
38+
let (_, proof, challenges) = assist_sumcheck_prove_impl(
39+
z_row_padded,
40+
rho_padded,
41+
eq_col,
42+
cumulative_heights,
43+
n_robp,
44+
transcript,
45+
false,
46+
);
47+
(proof, challenges)
48+
}
49+
50+
/// Compute the assist claimed sum, append it to the transcript, and run the
51+
/// assist sumcheck prover using the same ROBP precompute.
52+
pub fn assist_sumcheck_prove_and_append_claim<E: ExtensionField>(
53+
z_row_padded: &[E],
54+
rho_padded: &[E],
55+
eq_col: &[E],
56+
cumulative_heights: &[usize],
57+
n_robp: usize,
58+
transcript: &mut impl Transcript<E>,
59+
) -> (E, IOPProof<E>, Vec<E>) {
60+
assist_sumcheck_prove_impl(
61+
z_row_padded,
62+
rho_padded,
63+
eq_col,
64+
cumulative_heights,
65+
n_robp,
66+
transcript,
67+
true,
68+
)
69+
}
70+
71+
fn assist_sumcheck_prove_impl<E: ExtensionField>(
72+
z_row_padded: &[E],
73+
rho_padded: &[E],
74+
eq_col: &[E],
75+
cumulative_heights: &[usize],
76+
n_robp: usize,
77+
transcript: &mut impl Transcript<E>,
78+
append_claim: bool,
79+
) -> (E, IOPProof<E>, Vec<E>) {
3880
let num_polys = cumulative_heights.len() - 1;
3981
let n_vars = 2 * n_robp;
4082
let max_degree: usize = 2;
4183

42-
// Write transcript header (must match verifier).
43-
transcript.append_message(&n_vars.to_le_bytes());
44-
transcript.append_message(&max_degree.to_le_bytes());
45-
4684
// Precompute per-step symbol matrices.
4785
let step_mats: Vec<[TransitionMatrix<E>; 4]> = (0..n_robp)
4886
.map(|i| symbol_transition_matrices(z_row_padded[i], rho_padded[i]))
4987
.collect();
5088

51-
// Extract Boolean bits in step-major layout: c_bits[i][y], d_bits[i][y].
52-
// c_bits[i][y] = bit_i(t_y), d_bits[i][y] = bit_i(t_{y+1})
53-
let mut c_bits = vec![vec![0usize; num_polys]; n_robp];
54-
let mut d_bits = vec![vec![0usize; num_polys]; n_robp];
55-
for i in 0..n_robp {
56-
for y in 0..num_polys {
57-
c_bits[i][y] = (cumulative_heights[y] >> i) & 1;
58-
d_bits[i][y] = (cumulative_heights[y + 1] >> i) & 1;
89+
// Extract Boolean symbol pairs in step-major layout:
90+
// cd_bits[i][y] = 2 * bit_i(t_y) + bit_i(t_{y+1}).
91+
let mut cd_bits = vec![vec![0u8; num_polys]; n_robp];
92+
for (i, cd_bits_i) in cd_bits.iter_mut().enumerate() {
93+
for (y, cd_bit) in cd_bits_i.iter_mut().enumerate() {
94+
*cd_bit = ((((cumulative_heights[y] >> i) & 1) << 1)
95+
| ((cumulative_heights[y + 1] >> i) & 1)) as u8;
5996
}
6097
}
6198

@@ -89,14 +126,27 @@ pub fn assist_sumcheck_prove<E: ExtensionField>(
89126
let dst = &mut left[i];
90127
let src = &right[0];
91128
dst.into_par_iter().enumerate().for_each(|(y, dst_y)| {
92-
let cd = c_bits[i][y] * 2 + d_bits[i][y];
129+
let cd = cd_bits[i][y] as usize;
93130
*dst_y = mat_vec_mul(&step_mats[i][cd], &src[y]);
94131
});
95132
}
96133

134+
let source = source_vec();
135+
let mut claimed_sum = E::ZERO;
136+
for (y, eq) in eq_col.iter().enumerate().take(num_polys) {
137+
claimed_sum += *eq * dot4(&source, &bwd[0][y]);
138+
}
139+
if append_claim {
140+
transcript.append_field_element_ext(&claimed_sum);
141+
}
142+
143+
// Write transcript header (must match verifier).
144+
transcript.append_message(&n_vars.to_le_bytes());
145+
transcript.append_message(&max_degree.to_le_bytes());
146+
97147
// Initialize weights and forward vector.
98148
let mut weights: Vec<E> = eq_col[..num_polys].to_vec();
99-
let mut fwd: StateVec<E> = source_vec();
149+
let mut fwd: StateVec<E> = source;
100150

101151
let mut challenges: Vec<E> = Vec::with_capacity(n_vars);
102152
let mut proof_messages: Vec<IOPProverMessage<E>> = Vec::with_capacity(n_vars);
@@ -138,7 +188,7 @@ pub fn assist_sumcheck_prove<E: ExtensionField>(
138188
.map(|chunk| {
139189
let mut local_bwd_sum = [[E::ZERO; ROBP_WIDTH]; 4];
140190
for &y in chunk {
141-
let cd = c_bits[i][y] * 2 + d_bits[i][y];
191+
let cd = cd_bits[i][y] as usize;
142192
let w = weights[y];
143193
for s in 0..ROBP_WIDTH {
144194
local_bwd_sum[cd][s] += w * bwd[i + 1][y][s];
@@ -220,7 +270,7 @@ pub fn assist_sumcheck_prove<E: ExtensionField>(
220270
let start = chunk_idx * batch_size;
221271
for (j, w) in w_chunk.iter_mut().enumerate() {
222272
let y = start + j;
223-
let cd = c_bits[i][y] * 2 + d_bits[i][y];
273+
let cd = cd_bits[i][y] as usize;
224274
*w *= eq_cd[cd];
225275
}
226276
});
@@ -235,6 +285,7 @@ pub fn assist_sumcheck_prove<E: ExtensionField>(
235285
}
236286

237287
(
288+
claimed_sum,
238289
IOPProof {
239290
proofs: proof_messages,
240291
},

crates/mpcs/src/jagged/mod.rs

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -112,7 +112,9 @@ pub mod evaluator;
112112
pub mod sumcheck;
113113
mod types;
114114

115-
pub use assist::{assist_sumcheck_prove, compute_q_at_assist_point};
115+
pub use assist::{
116+
assist_sumcheck_prove, assist_sumcheck_prove_and_append_claim, compute_q_at_assist_point,
117+
};
116118
pub use evaluator::{evaluate_g, evaluate_g_backward, evaluate_g_forward};
117119
pub use sumcheck::{JaggedSumcheckInput, QPrimeEvaluations, jagged_sumcheck_prove};
118120
pub use types::{JaggedBatchOpenProof, JaggedCommitment, JaggedCommitmentWithWitness, JaggedProof};

crates/mpcs/src/lib.rs

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -282,8 +282,8 @@ pub mod jagged;
282282
pub use jagged::{
283283
JAGGED_RESHAPE_GROUP_WIDTH, Jagged, JaggedBatchOpenProof, JaggedCommitment,
284284
JaggedCommitmentWithWitness, JaggedProof, JaggedSumcheckInput, assist_sumcheck_prove,
285-
evaluate_g, evaluate_g_backward, evaluate_g_forward, jagged_batch_open, jagged_batch_verify,
286-
jagged_commit, jagged_sumcheck_prove,
285+
assist_sumcheck_prove_and_append_claim, evaluate_g, evaluate_g_backward, evaluate_g_forward,
286+
jagged_batch_open, jagged_batch_verify, jagged_commit, jagged_sumcheck_prove,
287287
};
288288
#[cfg(feature = "whir")]
289289
extern crate whir as whir_external;

crates/sumcheck/src/frontload.rs

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -103,6 +103,7 @@ use crate::{
103103
pub struct FrontloadProverState<E: ExtensionField> {
104104
pub challenges: Vec<Challenge<E>>,
105105
pub final_evaluations: Vec<Vec<E>>,
106+
pub claimed_sum: E,
106107
}
107108

108109
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
@@ -181,6 +182,7 @@ pub fn prove_2phase<'a, E: ExtensionField>(
181182
}
182183
let mut proofs = Vec::with_capacity(global_num_vars);
183184
let mut challenge: Option<Challenge<E>> = None;
185+
let mut claimed_sum = E::ZERO;
184186

185187
for round in 0..local_num_vars {
186188
workers.par_iter_mut().for_each(|worker| {
@@ -208,6 +210,9 @@ pub fn prove_2phase<'a, E: ExtensionField>(
208210
},
209211
)
210212
};
213+
if round == 0 {
214+
claimed_sum = evaluations[0] + evaluations[1];
215+
}
211216
evaluations.remove(0);
212217
transcript.append_field_element_exts(&evaluations);
213218
proofs.push(IOPProverMessage { evaluations });
@@ -249,6 +254,7 @@ pub fn prove_2phase<'a, E: ExtensionField>(
249254
FrontloadProverState {
250255
challenges,
251256
final_evaluations,
257+
claimed_sum,
252258
},
253259
)
254260
}
@@ -268,6 +274,7 @@ fn prove_inner<'a, E: ExtensionField>(
268274

269275
let mut proof = Vec::with_capacity(num_vars);
270276
let mut challenge: Option<Challenge<E>> = None;
277+
let mut claimed_sum = E::ZERO;
271278

272279
for round in 0..num_vars {
273280
if let Some(challenge) = challenge.take() {
@@ -276,6 +283,9 @@ fn prove_inner<'a, E: ExtensionField>(
276283
}
277284

278285
let mut evaluations = state.round_evaluations(round);
286+
if round == 0 {
287+
claimed_sum = evaluations[0] + evaluations[1];
288+
}
279289
evaluations.remove(0);
280290
transcript.append_field_element_exts(&evaluations);
281291
proof.push(IOPProverMessage { evaluations });
@@ -293,6 +303,7 @@ fn prove_inner<'a, E: ExtensionField>(
293303
FrontloadProverState {
294304
challenges: state.challenges,
295305
final_evaluations,
306+
claimed_sum,
296307
},
297308
)
298309
}

crates/sumcheck/src/prover.rs

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -192,6 +192,7 @@ impl<'a, E: ExtensionField> IOPProverState<'a, E> {
192192
max_num_variables,
193193
poly_meta: vec![],
194194
final_evaluations: Some(state.final_evaluations),
195+
claimed_sum: state.claimed_sum,
195196
phase2_numvar: None,
196197
}
197198
}
@@ -449,6 +450,7 @@ impl<'a, E: ExtensionField> IOPProverState<'a, E> {
449450
poly: polynomial,
450451
poly_meta: poly_meta.unwrap_or_else(|| vec![PolyMeta::Normal; num_polys]),
451452
final_evaluations: None,
453+
claimed_sum: E::ZERO,
452454
phase2_numvar,
453455
}
454456
}
@@ -512,6 +514,9 @@ impl<'a, E: ExtensionField> IOPProverState<'a, E> {
512514
exit_span!(start);
513515

514516
assert!(uni_polys.len() > 1);
517+
if self.round == 1 {
518+
self.claimed_sum = uni_polys[0] + uni_polys[1];
519+
}
515520
// NOTE remove uni_polys.eval(0) from lagrange domain
516521
// as verifier can derive via claim - uni_polys.eval(1)
517522
uni_polys.remove(0);
@@ -843,6 +848,10 @@ impl<'a, E: ExtensionField> IOPProverState<'a, E> {
843848
.collect_vec()
844849
}
845850

851+
pub fn claimed_sum(&self) -> E {
852+
self.claimed_sum
853+
}
854+
846855
pub fn expected_numvars_at_round(&self) -> usize {
847856
// first round start from 1
848857
let num_vars = self.max_num_variables + 1 - self.round;
@@ -1021,6 +1030,7 @@ impl<'a, E: ExtensionField> IOPProverState<'a, E> {
10211030
poly: polynomial,
10221031
poly_meta,
10231032
final_evaluations: None,
1033+
claimed_sum: E::ZERO,
10241034
phase2_numvar: None,
10251035
};
10261036

@@ -1130,6 +1140,9 @@ impl<'a, E: ExtensionField> IOPProverState<'a, E> {
11301140
exit_span!(start);
11311141

11321142
assert!(uni_polys.len() > 1);
1143+
if self.round == 1 {
1144+
self.claimed_sum = uni_polys[0] + uni_polys[1];
1145+
}
11331146
// NOTE remove uni_polys.eval(0) from lagrange domain
11341147
// as verifier can derive via claim - uni_polys.eval(1)
11351148
uni_polys.remove(0);

crates/sumcheck/src/structs.rs

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -99,6 +99,7 @@ pub struct IOPProverState<'a, E: ExtensionField> {
9999
pub(crate) max_num_variables: usize,
100100
pub(crate) poly_meta: Vec<PolyMeta>,
101101
pub(crate) final_evaluations: Option<Vec<Vec<E>>>,
102+
pub(crate) claimed_sum: E,
102103
/// phase 1 and phase 2 sumcheck we share similar implementation
103104
/// thus this option variable only use for phase 1 sumcheck to mark how many variables belongs to phase 2
104105
pub(crate) phase2_numvar: Option<usize>,

0 commit comments

Comments
 (0)