@@ -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 } ,
0 commit comments