diff --git a/Cargo.lock b/Cargo.lock index dc22ef9..f59299b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -23,7 +23,7 @@ checksum = "9555578bc9e57714c812a1f84e4fc5b4d21fcb063490c624de019f7464c91268" [[package]] name = "circuits" version = "0.1.0" -source = "git+https://github.com/sublinearlabs/sl-core.git#fc79d9c6b30791b97868e48f8934bfed68e4b011" +source = "git+https://github.com/sublinearlabs/sl-core.git#f6ca2101de0e3734c389a4021e6012cdc052d5eb" dependencies = [ "p3-field", "p3-goldilocks", @@ -33,9 +33,9 @@ dependencies = [ [[package]] name = "crunchy" -version = "0.2.3" +version = "0.2.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "43da5946c66ffcc7745f48db692ffbb10a83bfe0afd96235c5c2a4fb23994929" +checksum = "460fbee9c2c2f33933d720630a6a0bac33ba7053db5344fac858d4b8952d77d5" [[package]] name = "either" @@ -43,6 +43,15 @@ version = "1.15.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "48c757948c5ede0e46177b7add2e67155f70e33c07fea8284df6576da70b3719" +[[package]] +name = "fields" +version = "0.1.0" +source = "git+https://github.com/sublinearlabs/sl-core.git#f6ca2101de0e3734c389a4021e6012cdc052d5eb" +dependencies = [ + "p3-field", + "p3-mersenne-31", +] + [[package]] name = "gcd" version = "2.3.0" @@ -312,8 +321,9 @@ checksum = "3b3cff922bd51709b605d9ead9aa71031d81447142d828eb4a6eba76fe619f9b" [[package]] name = "poly" version = "0.1.0" -source = "git+https://github.com/sublinearlabs/sl-core.git#fc79d9c6b30791b97868e48f8934bfed68e4b011" +source = "git+https://github.com/sublinearlabs/sl-core.git#f6ca2101de0e3734c389a4021e6012cdc052d5eb" dependencies = [ + "fields", "p3-challenger", "p3-field", "p3-goldilocks", @@ -407,7 +417,7 @@ checksum = "fe895eb47f22e2ddd4dabc02bce419d2e643c8e3b585c78158b349195bc24d82" [[package]] name = "sum_check" version = "0.1.0" -source = "git+https://github.com/sublinearlabs/sl-core.git#fc79d9c6b30791b97868e48f8934bfed68e4b011" +source = "git+https://github.com/sublinearlabs/sl-core.git#f6ca2101de0e3734c389a4021e6012cdc052d5eb" dependencies = [ "anyhow", "p3-challenger", @@ -420,9 +430,9 @@ dependencies = [ [[package]] name = "syn" -version = "2.0.103" +version = "2.0.104" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e4307e30089d6fd6aff212f2da3a1f9e32f3223b1f010fb09b7c95f90f3ca1e8" +checksum = "17b6f705963418cdb9927482fa304bc562ece2fdd4f616084c50b7023b435a40" dependencies = [ "proc-macro2", "quote", @@ -472,8 +482,9 @@ dependencies = [ [[package]] name = "transcript" version = "0.1.0" -source = "git+https://github.com/sublinearlabs/sl-core.git#fc79d9c6b30791b97868e48f8934bfed68e4b011" +source = "git+https://github.com/sublinearlabs/sl-core.git#f6ca2101de0e3734c389a4021e6012cdc052d5eb" dependencies = [ + "fields", "p3-challenger", "p3-field", "p3-keccak", diff --git a/README.md b/README.md index 8de00e1..64f7139 100644 --- a/README.md +++ b/README.md @@ -117,10 +117,11 @@ cargo test test_libra_protocol ## Performance The Libra protocol provides: -- **Prover time**: O(C log C) where C is circuit size -- **Verifier time**: O(log C) -- **Proof size**: O(log C) +- **Prover time**: O(C) +- **Verifier time**: O(d log C) +- **Proof size**: O(d log C) - **Memory usage**: Linear in circuit size +where C is circuit size and d is the circuit depth ## Contributing diff --git a/src/prover.rs b/src/prover.rs index dba07db..773ba95 100644 --- a/src/prover.rs +++ b/src/prover.rs @@ -26,123 +26,72 @@ pub fn prove>( let mut wcs = vec![]; // Get the output vector - let mut output_evals: Vec> = output.layers[circuit.layers.len()] + let output_evals: Vec> = output.layers[circuit.layers.len()] .iter() .map(|val| Fields::::Base(*val)) .collect(); - if output_evals.len() == 1 { - output_evals.push(Fields::Base(F::zero())); - } - // Build the output polynomial let output_mle = - MultilinearPoly::new_from_vec((output_evals.len() as f64).log2() as usize, output_evals); + MultilinearPoly::new_extend_to_power_of_two(output_evals, Fields::Base(F::zero())); // Adds the output to the transcript transcript.observe_base_element(&output.layers[circuit.layers.len()]); - // Gets the addi and muli for the output layer - let (add_i, mul_i) = - LibraGKRLayeredCircuitTr::::add_and_mul_mle(circuit, circuit.layers.len() - 1); - - // Gets w_i+1 - let mut w_i_plus_one_poly = MultilinearPoly::new_from_vec( - (output.layers[circuit.layers.len() - 1].len() as f64).log2() as usize, - output.layers[circuit.layers.len() - 1] - .iter() - .map(|val| Fields::Base(*val)) - .collect::>>(), - ); - // Sample random challenge for the first round let g = transcript .sample_n_challenges(output_mle.num_vars()) - .iter() - .map(|val| Fields::Extension(*val)) + .into_iter() + .map(Fields::Extension) .collect::>>(); - let mut igz = generate_eq(&g); - - // Prepares parameters for phase one of Libra - let (mut mul_ahg, mut add_b_ahg, mut add_c_ahg) = - prepare_phase_one_params(&igz, &add_i, &mul_i, &w_i_plus_one_poly.evaluations); - - let mut claimed_sum = output_mle.evaluate(&g); - - // Proves the sumcheck relation using Libra algorithms - let (mut sumcheck_proof, mut wb, mut wc) = prove_libra_sumcheck( - ProveLibraInput { - claimed_sum: &claimed_sum, - igz: &igz, - mul_ahg: &mul_ahg, - add_b_ahg: &add_b_ahg, - add_c_ahg: &add_c_ahg, - add_i: &add_i, - mul_i: &mul_i, - w_i_plus_one_poly: &w_i_plus_one_poly, - }, - &mut transcript, - ); - - let (mut rb, mut rc) = ( - sumcheck_proof.challenges[..sumcheck_proof.challenges.len() / 2].to_vec(), - sumcheck_proof.challenges[sumcheck_proof.challenges.len() / 2..].to_vec(), - ); - - // Add messages to the transcript - transcript.observe_ext_element(&[wb.to_extension_field()]); - transcript.observe_ext_element(&[wc.to_extension_field()]); - transcript.observe_ext_element(&[sumcheck_proof.claimed_sum.to_extension_field()]); - transcript.observe_ext_element(&sumcheck_proof.round_polynomials.iter().fold( - vec![], - |mut acc, val| { - acc.extend( - val.iter() - .map(|val| val.to_extension_field()) - .collect::>(), - ); - acc - }, - )); - - // Adds messages to the proof - sumcheck_proofs.push(sumcheck_proof); - wbs.push(wb); - wcs.push(wc); - - // Samples alpha and beta for folding - let mut alpha_n_beta = transcript - .sample_n_challenges(2) - .iter() - .map(|val| Fields::Extension(*val)) - .collect::>>(); + let mut challenges = vec![]; + let mut wb = Fields::Extension(E::zero()); + let mut wc = Fields::Extension(E::zero()); + let mut alpha_n_beta = vec![]; + + for i in (1..=circuit.layers.len()).rev() { + let claimed_sum = if i == circuit.layers.len() { + output_mle.evaluate(&g) + } else { + // Sample alpha and beta + alpha_n_beta = transcript + .sample_n_challenges(2) + .into_iter() + .map(Fields::Extension) + .collect::>>(); + (alpha_n_beta[0] * wb) + (alpha_n_beta[1] * wc) + }; - for i in (1..circuit.layers.len()).rev() { // Get addi and muli let (add_i, mul_i) = LibraGKRLayeredCircuitTr::::add_and_mul_mle(circuit, i - 1); - claimed_sum = alpha_n_beta[0] * w_i_plus_one_poly.evaluate(&rb) - + alpha_n_beta[1] * w_i_plus_one_poly.evaluate(&rc); + let w_i_plus_one_eval = output.layers[i - 1] + .iter() + .map(|val| Fields::::Base(*val)) + .collect(); // Gets w_i+1 - w_i_plus_one_poly = MultilinearPoly::new_from_vec( - (output.layers[i - 1].len() as f64).log2() as usize, - output.layers[i - 1] - .iter() - .map(|val| Fields::::Base(*val)) - .collect(), - ); + let w_i_plus_one_poly = + MultilinearPoly::new_extend_to_power_of_two(w_i_plus_one_eval, Fields::Base(F::zero())); // Fold Igz for rb and rc using alpha and beta - igz = igz_n_to_1_fold(&[&rb, &rc], &alpha_n_beta); + let igz = if i == circuit.layers.len() { + generate_eq(&g) + } else { + let (rb, rc) = ( + challenges[..challenges.len() / 2].to_vec(), + challenges[challenges.len() / 2..].to_vec(), + ); + igz_n_to_1_fold(&[&rb, &rc], &alpha_n_beta) + }; // Gets new addi and muli based on rb, rc, alpha and beta - (mul_ahg, add_b_ahg, add_c_ahg) = + let (mul_ahg, add_b_ahg, add_c_ahg) = prepare_phase_one_params(&igz, &add_i, &mul_i, &w_i_plus_one_poly.evaluations); // Proves sumcheck relation using Libra algorithms - (sumcheck_proof, wb, wc) = prove_libra_sumcheck( + let (sumcheck_proof, wb_eval, wc_eval) = prove_libra_sumcheck( ProveLibraInput { claimed_sum: &claimed_sum, igz: &igz, @@ -156,23 +105,18 @@ pub fn prove>( &mut transcript, ); - (rb, rc) = ( - sumcheck_proof.challenges[..sumcheck_proof.challenges.len() / 2].to_vec(), - sumcheck_proof.challenges[sumcheck_proof.challenges.len() / 2..].to_vec(), - ); + wb = wb_eval; + wc = wc_eval; + challenges = sumcheck_proof.challenges.to_vec(); // Adds the messages to the transcript - transcript.observe_ext_element(&[wb.to_extension_field()]); - transcript.observe_ext_element(&[wc.to_extension_field()]); - transcript.observe_ext_element(&[sumcheck_proof.claimed_sum.to_extension_field()]); - transcript.observe_ext_element(&sumcheck_proof.round_polynomials.iter().fold( + transcript.observe(&[wb]); + transcript.observe(&[wc]); + transcript.observe(&[sumcheck_proof.claimed_sum]); + transcript.observe(&sumcheck_proof.round_polynomials.iter().fold( vec![], |mut acc, val| { - acc.extend( - val.iter() - .map(|val| val.to_extension_field()) - .collect::>(), - ); + acc.extend(val); acc }, )); @@ -181,13 +125,6 @@ pub fn prove>( sumcheck_proofs.push(sumcheck_proof); wbs.push(wb); wcs.push(wc); - - // Sample alpha and beta - alpha_n_beta = transcript - .sample_n_challenges(2) - .iter() - .map(|val| Fields::Extension(*val)) - .collect::>>(); } LibraProof::new( diff --git a/src/tests.rs b/src/tests.rs index 909f8b6..7018cec 100644 --- a/src/tests.rs +++ b/src/tests.rs @@ -48,13 +48,12 @@ fn test_libra_protocol() { let input = [1, 2, 3, 2, 1, 2, 4, 1] .into_iter() - .map(Mersenne31::from_canonical_usize) - .collect::>(); + .map(F::from_canonical_usize) + .collect::>(); let output = circuit.excecute(&input); - let proof: LibraProof> = - prove(&circuit, output); + let proof: LibraProof = prove(&circuit, output); let verify = verify(&circuit, proof, input); @@ -339,30 +338,14 @@ fn test_precompute() { let precomputed = generate_eq(&challenges); let expected = [ - Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_usize(0), - )), - Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_usize(0), - )), - Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_usize(0), - )), - Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_usize(0), - )), - Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_usize(8), - )), - -Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_usize(10), - )), - -Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_usize(12), - )), - Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_usize(15), - )), + Fields::Extension(E::from_base(F::from_canonical_usize(0))), + Fields::Extension(E::from_base(F::from_canonical_usize(0))), + Fields::Extension(E::from_base(F::from_canonical_usize(0))), + Fields::Extension(E::from_base(F::from_canonical_usize(0))), + Fields::Extension(E::from_base(F::from_canonical_usize(8))), + -Fields::Extension(E::from_base(F::from_canonical_usize(10))), + -Fields::Extension(E::from_base(F::from_canonical_usize(12))), + Fields::Extension(E::from_base(F::from_canonical_usize(15))), ]; assert_eq!(precomputed, expected); @@ -393,44 +376,19 @@ fn test_build_ahg() { // f(out, left, right) in the sparse form let f1 = vec![(0, 0, 1), (1, 2, 3), (2, 4, 5), (3, 6, 7)]; - let f3 = vec![ - Fields::from_u32(3), - Fields::from_u32(4), - Fields::from_u32(5), - Fields::from_u32(6), - Fields::from_u32(7), - Fields::from_u32(8), - Fields::from_u32(9), - Fields::from_u32(10), - ]; + let f3 = Fields::::from_u32_vec(vec![3, 4, 5, 6, 7, 8, 9, 10]); let ahg = initialize_phase_one(&igz, &f1, &f3); let expected = [ - Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_u32(32), - )), - Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_u32(0), - )), - -Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_u32(60), - )), - Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_u32(0), - )), - -Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_u32(96), - )), - Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_u32(0), - )), - Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_u32(150), - )), - Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_u32(0), - )), + Fields::Extension(E::from_base(F::from_canonical_u32(32))), + Fields::Extension(E::from_base(F::from_canonical_u32(0))), + -Fields::Extension(E::from_base(F::from_canonical_u32(60))), + Fields::Extension(E::from_base(F::from_canonical_u32(0))), + -Fields::Extension(E::from_base(F::from_canonical_u32(96))), + Fields::Extension(E::from_base(F::from_canonical_u32(0))), + Fields::Extension(E::from_base(F::from_canonical_u32(150))), + Fields::Extension(E::from_base(F::from_canonical_u32(0))), ]; assert_eq!(ahg, expected); @@ -485,30 +443,14 @@ fn test_build_af1() { let af1 = initialize_phase_two(&igz, &iux, &f1); let expected = [ - Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_usize(0), - )), - -Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_usize(288), - )), - Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_usize(0), - )), - -Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_usize(420), - )), - Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_usize(0), - )), - -Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_usize(576), - )), - Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_usize(0), - )), - -Fields::Extension(BinomialExtensionField::from_base( - Mersenne31::from_canonical_usize(840), - )), + Fields::Extension(E::from_base(F::from_canonical_usize(0))), + -Fields::Extension(E::from_base(F::from_canonical_usize(288))), + Fields::Extension(E::from_base(F::from_canonical_usize(0))), + -Fields::Extension(E::from_base(F::from_canonical_usize(420))), + Fields::Extension(E::from_base(F::from_canonical_usize(0))), + -Fields::Extension(E::from_base(F::from_canonical_usize(576))), + Fields::Extension(E::from_base(F::from_canonical_usize(0))), + -Fields::Extension(E::from_base(F::from_canonical_usize(840))), ]; assert_eq!(af1, expected); diff --git a/src/utils.rs b/src/utils.rs index c4fcc4e..27dfa2d 100644 --- a/src/utils.rs +++ b/src/utils.rs @@ -1,6 +1,6 @@ use circuits::layered_circuit::utils::get_gate_properties; use p3_field::{ExtensionField, Field}; -use poly::{Fields, MultilinearExtension, mle::MultilinearPoly, utils::generate_eq, vpoly::VPoly}; +use poly::{Fields, mle::MultilinearPoly, utils::generate_eq, vpoly::VPoly}; use std::rc::Rc; use sum_check::primitives::SumCheckProof; @@ -123,13 +123,13 @@ pub fn build_phase_one_libra_sumcheck_poly>( add_c_ahg: &[Fields], w_i_plus_one_poly: &MultilinearPoly, ) -> VPoly { - let n_vars = w_i_plus_one_poly.num_vars(); + let padding = Fields::Extension(E::zero()); VPoly::new( vec![ - MultilinearPoly::new_from_vec(n_vars, mul_ahg.to_vec()), - MultilinearPoly::new_from_vec(n_vars, add_b_ahg.to_vec()), - MultilinearPoly::new_from_vec(n_vars, add_c_ahg.to_vec()), + MultilinearPoly::new_extend_to_power_of_two(mul_ahg.to_vec(), padding), + MultilinearPoly::new_extend_to_power_of_two(add_b_ahg.to_vec(), padding), + MultilinearPoly::new_extend_to_power_of_two(add_c_ahg.to_vec(), padding), w_i_plus_one_poly.clone(), ], 2, @@ -159,13 +159,20 @@ pub fn build_phase_two_libra_sumcheck_poly>( wb: &Fields, w_i_plus_one_poly: &MultilinearPoly, ) -> VPoly { - let n_vars = w_i_plus_one_poly.num_vars(); - VPoly::new( vec![ - MultilinearPoly::new_from_vec(n_vars, mul_af1.to_vec()), - MultilinearPoly::new_from_vec(n_vars, add_af1.to_vec()), - MultilinearPoly::new_from_vec(n_vars, vec![*wb; w_i_plus_one_poly.evaluations.len()]), + MultilinearPoly::new_extend_to_power_of_two( + mul_af1.to_vec(), + Fields::Extension(E::zero()), + ), + MultilinearPoly::new_extend_to_power_of_two( + add_af1.to_vec(), + Fields::Extension(E::zero()), + ), + MultilinearPoly::new_extend_to_power_of_two( + vec![*wb; w_i_plus_one_poly.evaluations.len()], + Fields::Extension(E::zero()), + ), w_i_plus_one_poly.clone(), ], 2, diff --git a/src/verifier.rs b/src/verifier.rs index 48ccd50..b4b503a 100644 --- a/src/verifier.rs +++ b/src/verifier.rs @@ -26,98 +26,31 @@ pub fn verify>( let mut transcript = Transcript::::init(); // Adds output to the transcript - transcript.observe_base_element( - &proofs - .circuit_output - .iter() - .map(|val| val.to_base_field().unwrap()) - .collect::>(), - ); + transcript.observe(&proofs.circuit_output); // Gets output vector - let mut output: Vec> = proofs.circuit_output; - - if output.len() == 1 { - output.push(Fields::Base(F::zero())); - } + let output: Vec> = proofs.circuit_output; // Build output polynomial - let output_mle = MultilinearPoly::new_from_vec((output.len() as f64).log2() as usize, output); + let output_mle = MultilinearPoly::new_extend_to_power_of_two(output, Fields::Base(F::zero())); // Samples challenge for round one let g = transcript .sample_n_challenges(output_mle.num_vars()) - .iter() - .map(|val| Fields::Extension(*val)) + .into_iter() + .map(Fields::Extension) .collect::>>(); // Gets claimed sum by evaluating output polynomial at random challenge let mut claimed_sum = output_mle.evaluate(&g); - // Asserts claimed sum is equal to the proof claimed sum - assert_eq!(claimed_sum, proofs.sumcheck_proofs[0].claimed_sum); - - // Verify prover sumcheck - let (sumcheck_claimed_sum, rb_n_rc) = SumCheck::>::verify_partial( - &proofs.sumcheck_proofs[0], - &mut transcript, - ); - - // Get the correct number of vars - let n_vars = compute_num_vars(circuit.layers.len() - 1, circuit.layers.len() - 1); - - // Gets addi and muli - let (add_i, mul_i) = - LibraGKRLayeredCircuitTr::::add_and_mul_mle(circuit, circuit.layers.len() - 1); - - // Gets challenges for use in oracle check - let g_bc = [g.clone(), rb_n_rc.clone()].concat(); - - let mut rb = rb_n_rc[..rb_n_rc.len() / 2].to_vec(); + let mut alpha_n_beta = vec![]; - let mut rc = rb_n_rc[(rb_n_rc.len() / 2)..].to_vec(); + let mut rb = vec![]; - // Add messages to the transcript - transcript.observe_ext_element(&[proofs.wbs[0].to_extension_field()]); - transcript.observe_ext_element(&[proofs.wcs[0].to_extension_field()]); - transcript.observe_ext_element(&[proofs.sumcheck_proofs[0].claimed_sum.to_extension_field()]); - transcript.observe_ext_element(&proofs.sumcheck_proofs[0].round_polynomials.iter().fold( - vec![], - |mut acc, val| { - acc.extend( - val.iter() - .map(|val| val.to_extension_field()) - .collect::>(), - ); - acc - }, - )); - - // Calculate the expected claimed sum given wb and wc and challenges - let expected_sum = eval_layer_mle_given_wb_n_wc( - &add_i, - &mul_i, - &g_bc, - &proofs.wbs[0], - &proofs.wcs[0], - circuit.layers.len() - 1, - n_vars, - ); + let mut rc = vec![]; - // Performs oracle check on the first round - assert_eq!(sumcheck_claimed_sum, expected_sum.to_extension_field()); - - // Get alpha and beta - let mut alpha_n_beta = transcript - .sample_n_challenges(2) - .iter() - .map(|val| Fields::Extension(*val)) - .collect::>>(); - - // Get claimed sum for the next round by calculating: (alpha * wb) + (beta * wc) - claimed_sum = (alpha_n_beta[0] * proofs.wbs[0]) + (alpha_n_beta[1] * proofs.wcs[0]); - - for i in 1..(proofs.sumcheck_proofs.len() - 1) { + for i in 0..proofs.sumcheck_proofs.len() { // Assert claimed sum equals prover claimed sum assert_eq!(claimed_sum, proofs.sumcheck_proofs[i].claimed_sum); @@ -141,20 +74,59 @@ pub fn verify>( ); // Calculates expected claimed sum given wb and wc values - let expected_claimed_sum = eval_new_addi_n_muli_at_rb_bc_n_rc_bc( - EvalNewAddNMulInput { - add_i: &add_i, - mul_i: &mul_i, - alpha_n_beta: &alpha_n_beta, - rb: &rb, - rc: &rc, - bc: &rb_n_rc, - wb: &proofs.wbs[i], - wc: &proofs.wcs[i], - }, - proofs.sumcheck_proofs.len() - i - 1, - n_vars, - ); + let expected_claimed_sum = if i == 0 { + let g_bc = [g.clone(), rb_n_rc.clone()].concat(); + + eval_layer_mle_given_wb_n_wc( + &add_i, + &mul_i, + &g_bc, + &proofs.wbs[0], + &proofs.wcs[0], + circuit.layers.len() - 1, + n_vars, + ) + } else if i == proofs.sumcheck_proofs.len() - 1 { + let input_poly: MultilinearPoly = MultilinearPoly::new_extend_to_power_of_two( + input.iter().map(|val| Fields::Base(*val)).collect(), + Fields::Base(F::zero()), + ); + // // Calculate wb and wc on the input polynomial + let wb = input_poly.evaluate(&rb_n_rc[..rb_n_rc.len() / 2]); + + let wc = input_poly.evaluate(&rb_n_rc[(rb_n_rc.len() / 2)..]); + + // Get the expected claimed sum + eval_new_addi_n_muli_at_rb_bc_n_rc_bc( + EvalNewAddNMulInput { + add_i: &add_i, + mul_i: &mul_i, + alpha_n_beta: &alpha_n_beta, + rb: &rb, + rc: &rc, + bc: &rb_n_rc, + wb: &wb, + wc: &wc, + }, + proofs.sumcheck_proofs.len() - 1, + n_vars, + ) + } else { + eval_new_addi_n_muli_at_rb_bc_n_rc_bc( + EvalNewAddNMulInput { + add_i: &add_i, + mul_i: &mul_i, + alpha_n_beta: &alpha_n_beta, + rb: &rb, + rc: &rc, + bc: &rb_n_rc, + wb: &proofs.wbs[i], + wc: &proofs.wcs[i], + }, + proofs.sumcheck_proofs.len() - i - 1, + n_vars, + ) + }; // Do oracle check assert_eq!( @@ -163,18 +135,13 @@ pub fn verify>( ); // Add messages to the transcript - transcript.observe_ext_element(&[proofs.wbs[i].to_extension_field()]); - transcript.observe_ext_element(&[proofs.wcs[i].to_extension_field()]); - transcript - .observe_ext_element(&[proofs.sumcheck_proofs[i].claimed_sum.to_extension_field()]); - transcript.observe_ext_element(&proofs.sumcheck_proofs[i].round_polynomials.iter().fold( + transcript.observe(&[proofs.wbs[i]]); + transcript.observe(&[proofs.wcs[i]]); + transcript.observe(&[proofs.sumcheck_proofs[i].claimed_sum]); + transcript.observe(&proofs.sumcheck_proofs[i].round_polynomials.iter().fold( vec![], |mut acc, val| { - acc.extend( - val.iter() - .map(|val| val.to_extension_field()) - .collect::>(), - ); + acc.extend(val); acc }, )); @@ -182,81 +149,19 @@ pub fn verify>( // Sample alpha and beta alpha_n_beta = transcript .sample_n_challenges(2) - .iter() - .map(|val| Fields::Extension(*val)) + .into_iter() + .map(Fields::Extension) .collect::>>(); + let (wb, wc) = (proofs.wbs[i], proofs.wcs[i]); + // Get claimed sum for the next round - claimed_sum = (alpha_n_beta[0] * proofs.wbs[i]) + (alpha_n_beta[1] * proofs.wcs[i]); + claimed_sum = (alpha_n_beta[0] * wb) + (alpha_n_beta[1] * wc); rb = rb_n_rc[..rb_n_rc.len() / 2].to_vec(); rc = rb_n_rc[(rb_n_rc.len() / 2)..].to_vec(); } - // Verify sumcheck proof for the final round - let (sumcheck_claimed_sum, rb_n_rc) = SumCheck::>::verify_partial( - &proofs.sumcheck_proofs[proofs.sumcheck_proofs.len() - 1], - &mut transcript, - ); - - // Get addi and muli - let (add_i, mul_i) = LibraGKRLayeredCircuitTr::::add_and_mul_mle(circuit, 0); - - // Get the number of vars - let n_vars = compute_num_vars(0, circuit.layers.len() - 1); - - // Calculate expected claimed sum given wb and wc - let expected_claimed_sum = eval_new_addi_n_muli_at_rb_bc_n_rc_bc( - EvalNewAddNMulInput { - add_i: &add_i, - mul_i: &mul_i, - alpha_n_beta: &alpha_n_beta, - rb: &rb, - rc: &rc, - bc: &rb_n_rc, - wb: &proofs.wbs[proofs.sumcheck_proofs.len() - 1], - wc: &proofs.wcs[proofs.sumcheck_proofs.len() - 1], - }, - proofs.sumcheck_proofs.len() - 1, - n_vars, - ); - - // Do oracle check - assert_eq!( - sumcheck_claimed_sum, - expected_claimed_sum.to_extension_field() - ); - - // Get the input polynomial - let input_poly: MultilinearPoly = MultilinearPoly::new_from_vec( - (input.len() as f64).log2() as usize, - input.iter().map(|val| Fields::Base(*val)).collect(), - ); - - // Calculate wb and wc on the input polynomial - let wb = input_poly.evaluate(&rb_n_rc[..rb_n_rc.len() / 2]); - - let wc = input_poly.evaluate(&rb_n_rc[(rb_n_rc.len() / 2)..]); - - // Get the expected claimed sum - let oracle_query = eval_new_addi_n_muli_at_rb_bc_n_rc_bc( - EvalNewAddNMulInput { - add_i: &add_i, - mul_i: &mul_i, - alpha_n_beta: &alpha_n_beta, - rb: &rb, - rc: &rc, - bc: &rb_n_rc, - wb: &wb, - wc: &wc, - }, - proofs.sumcheck_proofs.len() - 1, - n_vars, - ); - - // Do the oracle check - assert_eq!(expected_claimed_sum, oracle_query); - Ok(true) }