diff --git a/Cargo.lock b/Cargo.lock index f32f14ce9..44837b0fe 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -7049,7 +7049,7 @@ dependencies = [ [[package]] name = "whir" version = "0.1.0" -source = "git+https://github.com/WizardOfMenlo/whir/?rev=0aeaa7f337c743d9ddfcb9d909628d6491e3355c#0aeaa7f337c743d9ddfcb9d909628d6491e3355c" +source = "git+https://github.com/worldfnd/whir.git?rev=33fbecf6ef883054ce8d93ad1bf1ef6f20b47e4f#33fbecf6ef883054ce8d93ad1bf1ef6f20b47e4f" dependencies = [ "ark-ff", "ark-serialize", @@ -7071,6 +7071,7 @@ dependencies = [ "sha3", "spongefish", "static_assertions", + "thiserror 2.0.18", "tracing", "zerocopy", ] diff --git a/Cargo.toml b/Cargo.toml index 67aee4a7f..4f2b6afc9 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -212,4 +212,4 @@ spongefish = { git = "https://github.com/arkworks-rs/spongefish", features = [ "sha2", ], rev = "fcc277f8a857fdeeadd7cca92ab08de63b1ff1a1" } spongefish-pow = { git = "https://github.com/arkworks-rs/spongefish", rev = "fcc277f8a857fdeeadd7cca92ab08de63b1ff1a1" } -whir = { git = "https://github.com/WizardOfMenlo/whir/", rev = "0aeaa7f337c743d9ddfcb9d909628d6491e3355c", features = ["tracing", "rs_in_order"] } +whir = { git = "https://github.com/worldfnd/whir.git", rev = "33fbecf6ef883054ce8d93ad1bf1ef6f20b47e4f", features = ["tracing", "rs_in_order"] } diff --git a/provekit/backend/bn254/benches/rs_bench.rs b/provekit/backend/bn254/benches/rs_bench.rs index b013fda53..de6465b1a 100644 --- a/provekit/backend/bn254/benches/rs_bench.rs +++ b/provekit/backend/bn254/benches/rs_bench.rs @@ -3,7 +3,10 @@ use { ark_ff::UniformRand, divan::{black_box, Bencher}, provekit_backend_bn254::RSFr, - whir::algebra::ntt::{NttEngine, ReedSolomon}, + whir::{ + algebra::ntt::{Messages, NttEngine, ReedSolomon}, + buffer::{Buffer, BufferOps}, + }, }; // (exp, expansion, coset_sz): matches whir's expand_from_coeff bench cases. @@ -19,21 +22,29 @@ const TEST_CASES: &[(usize, usize, usize)] = &[ (22, 4, 4), ]; -fn make_messages(exp: usize, coset_sz: usize) -> Vec> { +fn make_messages(exp: usize, coset_sz: usize) -> Vec> { let message_length = 1 << (exp - coset_sz); let num_messages = 1 << coset_sz; let mut rng = ark_std::rand::thread_rng(); (0..num_messages) - .map(|_| (0..message_length).map(|_| Fr::rand(&mut rng)).collect()) + .map(|_| { + Buffer::from( + (0..message_length) + .map(|_| Fr::rand(&mut rng)) + .collect::>(), + ) + }) .collect() } -fn make_mask(num_messages: usize) -> Vec { +fn make_mask(num_messages: usize) -> Buffer { let mask_length = 1 << 10; let mut rng = ark_std::rand::thread_rng(); - (0..num_messages * mask_length) - .map(|_| Fr::rand(&mut rng)) - .collect() + Buffer::from( + (0..num_messages * mask_length) + .map(|_| Fr::rand(&mut rng)) + .collect::>(), + ) } #[divan::bench(args = TEST_CASES)] @@ -43,9 +54,11 @@ fn rs_fr(bencher: Bencher, case: &(usize, usize, usize)) { bencher .with_inputs(|| make_messages(exp, coset_sz)) .bench_values(|coeffs| { - let refs: Vec<&[Fr]> = coeffs.iter().map(Vec::as_slice).collect(); - let codeword_length = refs[0].len() * expansion; - black_box(RSFr.interleaved_encode(&refs, &mask, codeword_length)) + let vector_refs: Vec<&Buffer> = coeffs.iter().collect(); + let message_length = coeffs[0].len(); + let rs_messages = Messages::new(&vector_refs, message_length, 1); + let codeword_length = message_length * expansion; + black_box(RSFr.interleaved_encode(rs_messages, &mask, codeword_length)) }); } @@ -57,9 +70,11 @@ fn whir_ntt_engine(bencher: Bencher, case: &(usize, usize, usize)) { bencher .with_inputs(|| make_messages(exp, coset_sz)) .bench_values(|coeffs| { - let refs: Vec<&[Fr]> = coeffs.iter().map(Vec::as_slice).collect(); - let codeword_length = refs[0].len() * expansion; - black_box(reference.interleaved_encode(&refs, &mask, codeword_length)) + let vector_refs: Vec<&Buffer> = coeffs.iter().collect(); + let message_length = coeffs[0].len(); + let rs_messages = Messages::new(&vector_refs, message_length, 1); + let codeword_length = message_length * expansion; + black_box(reference.interleaved_encode(rs_messages, &mask, codeword_length)) }); } diff --git a/provekit/backend/bn254/src/ntt.rs b/provekit/backend/bn254/src/ntt.rs index 2b4a64c1a..f5cd98917 100644 --- a/provekit/backend/bn254/src/ntt.rs +++ b/provekit/backend/bn254/src/ntt.rs @@ -3,7 +3,10 @@ use { ark_ff::{AdditiveGroup, FftField, Field}, ntt::ntt_nr, tracing::instrument, - whir::algebra::ntt::ReedSolomon, + whir::{ + algebra::ntt::{Messages, ReedSolomon}, + buffer::{Buffer, BufferOps}, + }, }; #[derive(Debug)] @@ -42,26 +45,36 @@ impl ReedSolomon for RSFr { } #[instrument(skip(self, messages, masks), fields( - num_messages = messages.len(), - message_len = messages.first().map(|c| c.len()), + num_messages = messages.vectors.len() * messages.interleaving_depth, + message_len = messages.message_length, codeword_length = codeword_length, - mask_len = masks.len().checked_div(messages.len()) - + mask_len = masks.len().checked_div(messages.vectors.len() * messages.interleaving_depth) ))] fn interleaved_encode( &self, - messages: &[&[Fr]], - masks: &[Fr], + messages: Messages<'_, Fr>, + masks: &Buffer, codeword_length: usize, - ) -> Vec { + ) -> Buffer { + let masks = masks.to_slice(); + let messages = messages + .vectors + .iter() + .flat_map(|message| { + message + .to_slice() + .chunks_exact(messages.message_length) + .take(messages.interleaving_depth) + }) + .collect::>(); if messages.is_empty() { - return vec![]; + return Buffer::from(vec![]); } let num_messages = messages.len(); let message_length = messages[0].len(); - for message in messages { + for message in &messages { assert_eq!(message_length, message.len()) } @@ -96,7 +109,7 @@ impl ReedSolomon for RSFr { ntt_nr(&mut result, codeword_length, num_cosets); - result + Buffer::from(result) } fn generator(&self, codeword_length: usize) -> Fr { @@ -136,7 +149,8 @@ mod tests { let mut data = messages_flat; data.resize(total, Fr::ZERO); - let messages: Vec<&[Fr]> = data.chunks(message_length).collect(); + let messages: Vec> = data.chunks(message_length).map(|c| Buffer::from(c)).collect(); + let messages_refs: Vec<&Buffer> = messages.iter().collect(); // Our masks are interleaved: num_messages x mask_length in row-major order // i.e. [m0_c0, m1_c0, m0_c1, m1_c1, ...] @@ -154,11 +168,16 @@ mod tests { } } + let messages = Messages::new(&messages_refs, message_length, 1); + let masks = Buffer::from(masks); + let masks_transposed = Buffer::from(masks_transposed); + let indices: Vec = (0..codeword_length).collect(); let reference = NttEngine::::new_from_fftfield(); - let our_codeword = RSFr.interleaved_encode(&messages, &masks, codeword_length); - let ref_codeword = reference.interleaved_encode(&messages, &masks_transposed, codeword_length); + let our_codeword = RSFr.interleaved_encode(messages.clone(), &masks, codeword_length); + let ref_codeword = + reference.interleaved_encode(messages, &masks_transposed, codeword_length); let our_points = RSFr.evaluation_points(message_length, codeword_length, &indices); let ref_points = reference.evaluation_points(message_length, codeword_length, &indices); @@ -166,12 +185,12 @@ mod tests { // Pair each evaluation point with its num_messages-wide slice, then sort // by point so that ordering differences between implementations don't matter. let mut our_rows: Vec<_> = our_points.iter().enumerate() - .map(|(i, pt)| (pt.into_bigint(), &our_codeword[i * num_messages..(i + 1) * num_messages])) + .map(|(i, pt)| (pt.into_bigint(), &our_codeword.to_slice()[i * num_messages..(i + 1) * num_messages])) .collect(); our_rows.sort_by_key(|(k, _)| *k); let mut ref_rows: Vec<_> = ref_points.iter().enumerate() - .map(|(i, pt)| (pt.into_bigint(), &ref_codeword[i * num_messages..(i + 1) * num_messages])) + .map(|(i, pt)| (pt.into_bigint(), &ref_codeword.to_slice()[i * num_messages..(i + 1) * num_messages])) .collect(); ref_rows.sort_by_key(|(k, _)| *k); diff --git a/provekit/common/src/prefix_covector.rs b/provekit/common/src/prefix_covector.rs index e45fb099c..e48717254 100644 --- a/provekit/common/src/prefix_covector.rs +++ b/provekit/common/src/prefix_covector.rs @@ -1,6 +1,9 @@ use { ark_ff::{Field, One, Zero}, - whir::algebra::{embedding::Embedding, linear_form::LinearForm, mixed_dot, multilinear_extend}, + whir::{ + algebra::{embedding::Embedding, linear_form::LinearForm, mixed_dot, multilinear_extend}, + buffer::{Buffer, BufferOps}, + }, }; /// A covector that stores only a power-of-two prefix, with the rest @@ -209,9 +212,10 @@ pub fn build_prefix_covectors( #[must_use] pub fn compute_alpha_evals( embedding: &M, - polynomial: &[M::Source], + polynomial: &Buffer, alphas: &[Vec; N], ) -> Vec { + let polynomial = polynomial.to_slice(); alphas .iter() .map(|w| mixed_dot(embedding, w, &polynomial[..w.len()])) @@ -226,8 +230,9 @@ pub fn compute_public_eval( embedding: &M, x: M::Target, num_public_inputs: usize, - polynomial: &[M::Source], + polynomial: &Buffer, ) -> M::Target { + let polynomial = polynomial.to_slice(); let n = num_public_inputs + 1; let mut eval = M::Target::zero(); let mut x_pow = M::Target::one(); @@ -331,8 +336,9 @@ pub fn compute_challenge_eval( embedding: &M, x: M::Target, challenge_offsets: &[usize], - polynomial: &[M::Source], + polynomial: &Buffer, ) -> M::Target { + let polynomial = polynomial.to_slice(); let mut eval = M::Target::zero(); let mut x_pow = M::Target::one(); for &offset in challenge_offsets { @@ -615,6 +621,7 @@ mod tests { poly[1] = fe(42); poly[5] = fe(99); poly[11] = fe(17); + let poly = Buffer::from(poly); let embedding = whir::algebra::embedding::Identity::::new(); let eval = compute_challenge_eval(&embedding, x, &offsets, &poly); diff --git a/provekit/common/src/whir_r1cs.rs b/provekit/common/src/whir_r1cs.rs index 5ff886d82..0f92ed3ca 100644 --- a/provekit/common/src/whir_r1cs.rs +++ b/provekit/common/src/whir_r1cs.rs @@ -199,7 +199,7 @@ impl WhirR1CSScheme

{ /// (≈118-bit algebraic hardness plus light per-round PoW). fn whir_protocol_params(hash_id: EngineId) -> ProtocolParameters { ProtocolParameters { - unique_decoding: false, + decoding_regime: whir::protocols::params::DecodingRegime::Johnson, security_level: 128, pow_bits: 10, initial_folding_factor: WHIR_INITIAL_FOLDING_FACTOR, diff --git a/provekit/prover/src/whir_r1cs.rs b/provekit/prover/src/whir_r1cs.rs index 7b1baf97f..296968f63 100644 --- a/provekit/prover/src/whir_r1cs.rs +++ b/provekit/prover/src/whir_r1cs.rs @@ -23,7 +23,6 @@ use { Base, Ext, FieldHash, PrefixCovector, ProofField, PublicInputs, WhirR1CSProof, WhirR1CSScheme, R1CS, }, - std::borrow::Cow, whir::{ algebra::{ dot, @@ -31,6 +30,7 @@ use { linear_form::LinearForm, mixed_dot, }, + buffer::{Buffer, BufferOps}, protocols::whir::Witness as WhirWitness, transcript::{Codec, DuplexSpongeInterface, ProverState, VerifierMessage}, }, @@ -42,14 +42,14 @@ pub struct BlindingState { pub polynomial: Vec<[Ext

; 4]>, /// `polynomial` flattened (length `4 * m_0`) and zero-padded to the /// blinding commitment domain — the vector actually committed. - pub vector: Vec>, + pub vector: Buffer>, /// WHIR witness for the ext blinding commitment. pub witness: WhirWitness, Identity>>, } pub struct WhirR1CSCommitment { pub witness: WhirWitness, P::Embedding>, - pub polynomial: Vec>, + pub polynomial: Buffer>, pub blinding: Option>, } @@ -145,6 +145,7 @@ where // Commit the base-field witness directly (non-hiding — openings leak // witness values; see the `whir_witness` field docs). + let padded_witness = Buffer::from(padded_witness); let witness_commitment = self.whir_witness.commit(merlin, &[&padded_witness]); // Commit the Spartan sumcheck blinding `g` separately, natively in the @@ -155,6 +156,7 @@ where let blind_len = self.blinding_domain_size(); let mut g_vector: Vec> = g.iter().flatten().copied().collect(); g_vector.resize(blind_len, >::zero()); + let g_vector = Buffer::from(g_vector); let blinding_witness = self.whir_blinding.commit(merlin, &[&g_vector]); Some(BlindingState { polynomial: g, @@ -396,10 +398,10 @@ where let final_claim = scheme.whir_witness.prove( &mut merlin, - vec![Cow::Borrowed(commitment.polynomial.as_slice())], - vec![Cow::Owned(commitment.witness)], + &[&commitment.polynomial], + vec![&commitment.witness], boxed_weights, - Cow::Borrowed(&evaluations), + Buffer::from(evaluations), ); spark_row.zip(spark_weights).map(|(row, spark_weights)| { @@ -500,10 +502,10 @@ where let final_claim = scheme.whir_witness.prove( &mut merlin, - vec![Cow::Borrowed(p1.as_slice())], - vec![Cow::Owned(w1)], + &[&p1], + vec![&w1], boxed_weights, - Cow::Borrowed(&evaluations), + Buffer::from(evaluations), ); let claimed = @@ -544,10 +546,10 @@ where let final_claim = scheme.whir_witness.prove( &mut merlin, - vec![Cow::Borrowed(p2.as_slice())], - vec![Cow::Owned(w2)], + &[&p2], + vec![&w2], boxed_weights, - Cow::Borrowed(&evaluations), + Buffer::from(evaluations), ); let claimed = @@ -594,10 +596,10 @@ where let blinding_covector = OffsetCovector::new(blinding_weights, 0, blind_domain); let _ = scheme.whir_blinding.prove( &mut merlin, - vec![Cow::Borrowed(blinding_state.vector.as_slice())], - vec![Cow::Owned(blinding_state.witness)], + &[&blinding_state.vector], + vec![&blinding_state.witness], vec![Box::new(blinding_covector) as Box>>], - Cow::Borrowed(&[blinding_eval]), + Buffer::from(vec![blinding_eval]), ); } @@ -727,13 +729,14 @@ pub fn run_zk_sumcheck_prover( merlin: &mut ProverState, m_0: usize, blinding_polynomial: &[[M::Target; 4]], - blinding_vector: &[M::Target], + blinding_vector: &Buffer, ) -> (Vec, M::Target) where M: Embedding, M::Target: Codec, S: DuplexSpongeInterface, { + let blinding_vector = blinding_vector.to_slice(); let [mut a, mut b, mut c] = mles; let r: Vec = merlin.verifier_message_vec(m_0); let mut eq = calculate_evaluations_over_boolean_hypercube_for_eq(&r, 1 << r.len()); @@ -887,9 +890,10 @@ fn combined_round_message>( fn create_weights_and_evaluations( embedding: &M, m: usize, - polynomial: &[M::Source], + polynomial: &Buffer, alphas: [Vec; N], ) -> (Vec>, Vec) { + let polynomial = polynomial.to_slice(); let domain_size = 1usize << m; let mut weights = Vec::with_capacity(N); @@ -909,8 +913,9 @@ fn create_weights_and_evaluations( fn compute_evaluations( embedding: &M, weights: &[PrefixCovector], - polynomial: &[M::Source], + polynomial: &Buffer, ) -> Vec { + let polynomial = polynomial.to_slice(); weights .iter() .map(|w| mixed_dot(embedding, w.vector(), &polynomial[..w.vector().len()])) @@ -920,10 +925,11 @@ fn compute_evaluations( fn compute_public_weight_evaluation( embedding: &M, weights: &mut Vec>, - polynomial: &[M::Source], + polynomial: &Buffer, public_weights: PrefixCovector, ) -> M::Target { let n = public_weights.vector().len(); + let polynomial = polynomial.to_slice(); let eval = mixed_dot(embedding, public_weights.vector(), &polynomial[..n]); weights.insert(0, public_weights); eval diff --git a/provekit/r1cs-compiler/src/noir_proof_scheme.rs b/provekit/r1cs-compiler/src/noir_proof_scheme.rs index a77cef538..6c802dedb 100644 --- a/provekit/r1cs-compiler/src/noir_proof_scheme.rs +++ b/provekit/r1cs-compiler/src/noir_proof_scheme.rs @@ -91,6 +91,8 @@ impl NoirCompiler { program: ProgramArtifact, hash_config: provekit_common::HashConfig, ) -> Result { + provekit_backend_bn254::register(); + info!("Program noir version: {}", program.noir_version); info!("Program entry point: fn main{};", PrintAbi(&program.abi)); ensure!( diff --git a/provekit/r1cs-compiler/src/whir_r1cs.rs b/provekit/r1cs-compiler/src/whir_r1cs.rs index 9e7ea54ba..549f128dc 100644 --- a/provekit/r1cs-compiler/src/whir_r1cs.rs +++ b/provekit/r1cs-compiler/src/whir_r1cs.rs @@ -34,6 +34,8 @@ impl MavrosSchemeBuilder for WhirR1CSScheme { has_public_inputs: bool, hash_config: HashConfig, ) -> Self { + provekit_backend_bn254::register(); + let num_witnesses = r1cs.witness_layout.size(); let num_constraints = r1cs.constraints.len(); let a_num_entries: usize = r1cs.constraints.iter().map(|c| c.a.len()).sum(); @@ -73,6 +75,8 @@ mod tests { expected_m: usize, expected_m_0: usize, ) { + provekit_backend_bn254::register(); + let from_dimensions = WhirR1CSScheme::::new_from_dimensions( num_witnesses, num_constraints, @@ -107,11 +111,14 @@ mod tests { /// Assert both WHIR commitments reach 128-bit security for field `P`. fn assert_configs_secure(size: usize) { + provekit_backend_bn254::register(); + provekit_backend_goldilocks::register(); + let field = std::any::type_name::

(); let witness = WhirR1CSScheme::

::new_witness_config_for_size(size, whir::hash::SHA2); let blinding = WhirR1CSScheme::

::new_blinding_config_for_size(size, whir::hash::SHA2); - let sec_witness = witness.security_level(witness.initial_committer.num_vectors, 1); - let sec_blinding = blinding.security_level(blinding.initial_committer.num_vectors, 1); + let sec_witness = witness.security_level(witness.initial_committer.num_vectors(), 1); + let sec_blinding = blinding.security_level(blinding.initial_committer.num_vectors(), 1); assert!( sec_witness >= 128.0, "Witness commitment security {sec_witness:.2} < 128 bits at size {size} for {field}" @@ -136,6 +143,8 @@ mod tests { #[test] fn mavros_dimensions_use_largest_commitment_not_total_witnesses() { + provekit_backend_bn254::register(); + let scheme = WhirR1CSScheme::::new_from_dimensions( 600_000, 8, diff --git a/provekit/spark/src/memory.rs b/provekit/spark/src/memory.rs index 50326477a..2b6299859 100644 --- a/provekit/spark/src/memory.rs +++ b/provekit/spark/src/memory.rs @@ -9,11 +9,11 @@ use { ark_std::One, provekit_backend_bn254::{FieldElement, TranscriptSponge, WhirConfig}, rayon::prelude::*, - std::borrow::Cow, tracing::instrument, whir::{ algebra::{linear_form::MultilinearExtension, multilinear_extend}, - protocols::irs_commit::Commitment, + buffer::{Buffer, BufferOps}, + protocols::whir::Commitment, transcript::{ProverState, VerifierState}, }, }; @@ -71,7 +71,7 @@ pub fn prove_axis_init_final_product( produce_whir_proof( merlin, evaluation_randomness, - &[config.final_timestamp], + &[&Buffer::from(config.final_timestamp)], config.whir_config, final_ts_witness, )?; @@ -138,7 +138,7 @@ pub fn verify_axis( pub fn produce_whir_proof( merlin: &mut ProverState, evaluation_point: &[FieldElement], - vectors: &[&[FieldElement]], + vectors: &[&Buffer], config: &WhirConfig, witness: &WhirWitness, ) -> Result<()> { @@ -146,18 +146,18 @@ pub fn produce_whir_proof( let evaluations: Vec = vectors .iter() - .map(|v| multilinear_extend(v, evaluation_point)) + .map(|v| multilinear_extend(v.to_slice(), evaluation_point)) .collect(); _ = config.prove( merlin, - vectors.iter().map(|v| Cow::Borrowed(*v)).collect(), - vec![Cow::Borrowed(witness)], + vectors, + vec![witness], vec![Box::new(lf) as Box< dyn whir::algebra::linear_form::LinearForm, >], - Cow::Borrowed(&evaluations), + Buffer::from(evaluations), ); Ok(()) diff --git a/provekit/spark/src/prover.rs b/provekit/spark/src/prover.rs index c1317bca5..7fb14bf53 100644 --- a/provekit/spark/src/prover.rs +++ b/provekit/spark/src/prover.rs @@ -26,10 +26,10 @@ use { HashConfig, WhirR1CSProof, }, rayon::{join, prelude::*}, - std::borrow::Cow, tracing::instrument, whir::{ algebra::{linear_form::MultilinearExtension, multilinear_extend}, + buffer::Buffer, engines::EngineId, parameters::ProtocolParameters, transcript::{DomainSeparator, ProverState, VerifierMessage}, @@ -46,10 +46,12 @@ pub fn new_whir_config_for_size( batch_size: usize, hash_id: EngineId, ) -> WhirConfig { + provekit_backend_bn254::register(); + let nv = log_size.max(4); let whir_params = ProtocolParameters { - unique_decoding: false, + decoding_regime: whir::protocols::params::DecodingRegime::Johnson, initial_folding_factor: 3, security_level: 128, pow_bits: 10, @@ -457,21 +459,21 @@ fn prove_combined_rs_ws_product( ]; let _ = whir_configs.num_terms_5batched.prove( merlin, - vec![ - Cow::Borrowed(&matrix.coo.val), - Cow::Borrowed(&row_field), - Cow::Borrowed(&read_row_field), - Cow::Borrowed(&col_field), - Cow::Borrowed(&read_col_field), + &[ + &Buffer::from(matrix.coo.val.as_slice()), + &Buffer::from(row_field), + &Buffer::from(read_row_field), + &Buffer::from(col_field), + &Buffer::from(read_col_field), ], - vec![Cow::Borrowed(vals_rs_ws_witness)], + vec![vals_rs_ws_witness], vec![ Box::new(fold_lf_for_vals_rs_ws) as Box>, Box::new(eval_lf_for_vals_rs_ws) as Box>, ], - Cow::Borrowed(&vals_rs_ws_evaluations), + Buffer::from(Vec::from(vals_rs_ws_evaluations)), ); let (row_value_eval, col_value_eval) = tracing::info_span!("multilinear_extend_e_values") @@ -494,13 +496,16 @@ fn prove_combined_rs_ws_product( ]; let _ = whir_configs.num_terms_2batched.prove( merlin, - vec![Cow::Borrowed(&e_values.e_rx), Cow::Borrowed(&e_values.e_ry)], - vec![Cow::Borrowed(e_values_witness)], + &[ + &Buffer::from(e_values.e_rx.as_slice()), + &Buffer::from(e_values.e_ry.as_slice()), + ], + vec![e_values_witness], vec![ Box::new(fold_lf) as Box>, Box::new(eval_lf) as Box>, ], - Cow::Borrowed(&evaluations), + Buffer::from(Vec::from(evaluations)), ); Ok(()) @@ -512,9 +517,10 @@ fn commit_e_values( whir_configs: &SparkWhirConfigs, e_values: &EValuesForMatrix, ) -> WhirWitness { - whir_configs - .num_terms_2batched - .commit(merlin, &[&e_values.e_rx, &e_values.e_ry]) + whir_configs.num_terms_2batched.commit(merlin, &[ + &Buffer::from(e_values.e_rx.as_slice()), + &Buffer::from(e_values.e_ry.as_slice()), + ]) } pub fn run_parallel_sumchecks( diff --git a/provekit/spark/src/serde_whir_witness.rs b/provekit/spark/src/serde_whir_witness.rs index 4f11481b3..6dcd57400 100644 --- a/provekit/spark/src/serde_whir_witness.rs +++ b/provekit/spark/src/serde_whir_witness.rs @@ -3,17 +3,21 @@ use { provekit_backend_bn254::FieldElement, provekit_common::utils::serde_ark_vec, serde::{ser::SerializeStruct, Deserialize, Deserializer, Serialize, Serializer}, - whir::protocols::{ - irs_commit::{Evaluations, Witness}, - matrix_commit, + whir::{ + buffer::{Buffer, BufferOps}, + protocols::{ + irs_commit::{self, Evaluations}, + matrix_commit, + whir::Witness, + }, }, }; pub fn serialize(w: &WhirWitness, s: S) -> Result { let mut st = s.serialize_struct("WhirWitness", 4)?; - st.serialize_field("masks", &ArkVecRef(&w.masks))?; - st.serialize_field("matrix", &ArkVecRef(&w.matrix))?; - st.serialize_field("matrix_witness", &w.matrix_witness)?; + st.serialize_field("masks", &ArkVecRef(w.irs.masks.to_slice()))?; + st.serialize_field("matrix", &ArkVecRef(w.irs.matrix.to_slice()))?; + st.serialize_field("matrix_witness", &w.irs.matrix_witness)?; st.serialize_field("out_of_domain", &EvaluationsRef(&w.out_of_domain))?; st.end() } @@ -21,21 +25,23 @@ pub fn serialize(w: &WhirWitness, s: S) -> Result>(d: D) -> Result { let m = WitnessMirror::deserialize(d)?; Ok(Witness { - masks: m.masks, - matrix: m.matrix, - matrix_witness: m.matrix_witness, - out_of_domain: Evaluations { + irs: irs_commit::Witness { + masks: Buffer::from(m.masks), + matrix: Buffer::from(m.matrix), + matrix_witness: m.matrix_witness, + }, + out_of_domain: Evaluations { points: m.out_of_domain.points, - matrix: m.out_of_domain.matrix, + matrix: Buffer::from(m.out_of_domain.matrix), }, }) } -struct ArkVecRef<'a>(&'a Vec); +struct ArkVecRef<'a>(&'a [FieldElement]); impl Serialize for ArkVecRef<'_> { fn serialize(&self, s: S) -> Result { - serde_ark_vec::serialize(self.0, s) + serde_ark_vec::serialize(&Vec::from(self.0), s) } } @@ -45,7 +51,7 @@ impl Serialize for EvaluationsRef<'_> { fn serialize(&self, s: S) -> Result { let mut st = s.serialize_struct("Evaluations", 2)?; st.serialize_field("points", &ArkVecRef(&self.0.points))?; - st.serialize_field("matrix", &ArkVecRef(&self.0.matrix))?; + st.serialize_field("matrix", &ArkVecRef(self.0.matrix.to_slice()))?; st.end() } } diff --git a/provekit/spark/src/setup.rs b/provekit/spark/src/setup.rs index e9124a90b..6a2521bec 100644 --- a/provekit/spark/src/setup.rs +++ b/provekit/spark/src/setup.rs @@ -8,7 +8,8 @@ use { provekit_common::{HashConfig, WhirR1CSProof}, tracing::instrument, whir::{ - protocols::irs_commit::Commitment, + buffer::Buffer, + protocols::whir::Commitment, transcript::{codecs::Empty, DomainSeparator, Proof, ProverState, VerifierState}, }, }; @@ -40,26 +41,22 @@ pub fn preprocess_spark( .whir_configs .num_terms_5batched .commit(&mut merlin, &[ - &matrix.coo.val, - &row_field, - &read_row_field, - &col_field, - &read_col_field, + &Buffer::from(matrix.coo.val.as_slice()), + &Buffer::from(row_field), + &Buffer::from(read_row_field), + &Buffer::from(col_field), + &Buffer::from(read_col_field), ]); - drop((row_field, col_field, read_row_field, read_col_field)); let final_row_field = matrix.timestamps.final_row_field(); let final_row_ts_witness = scheme .whir_configs .row - .commit(&mut merlin, &[&final_row_field]); - drop(final_row_field); + .commit(&mut merlin, &[&Buffer::from(final_row_field)]); let final_col_field = matrix.timestamps.final_col_field(); let final_col_ts_witness = scheme .whir_configs .col - .commit(&mut merlin, &[&final_col_field]); - drop(final_col_field); - + .commit(&mut merlin, &[&Buffer::from(final_col_field)]); let proof = merlin.proof(); let setup = SparkSetup { whir_configs: scheme.whir_configs, diff --git a/provekit/spark/src/types.rs b/provekit/spark/src/types.rs index 364611440..545de4a96 100644 --- a/provekit/spark/src/types.rs +++ b/provekit/spark/src/types.rs @@ -13,10 +13,10 @@ use { HashConfig, WhirR1CSProof, }, serde::{Deserialize, Serialize}, - whir::protocols::irs_commit, + whir::{algebra::embedding::Identity, protocols::whir::Witness}, }; -pub type WhirWitness = irs_commit::Witness; +pub type WhirWitness = Witness>; #[derive(Serialize, Deserialize)] #[serde(transparent)] diff --git a/tooling/provekit-gnark/src/gnark_config.rs b/tooling/provekit-gnark/src/gnark_config.rs index 4b08c866b..ed33b324d 100644 --- a/tooling/provekit-gnark/src/gnark_config.rs +++ b/tooling/provekit-gnark/src/gnark_config.rs @@ -75,29 +75,29 @@ impl WHIRConfigGnark { pub fn new(whir_params: &WhirConfig) -> Self { let n_rounds = whir_params.n_rounds(); let n_vars = whir_params.initial_num_variables(); - let message_length = whir_params.initial_committer.vector_size - / whir_params.initial_committer.interleaving_depth; + let message_length = whir_params.initial_committer.vector_size() + / whir_params.initial_committer.interleaving_depth(); let rate = - (whir_params.initial_committer.codeword_length / message_length).ilog2() as usize; + (whir_params.initial_committer.codeword_length() / message_length).ilog2() as usize; // Folding factor: initial round uses initial_sumcheck.num_rounds, // subsequent rounds use round_configs[i].sumcheck.num_rounds let mut folding_factor = Vec::with_capacity(n_rounds + 1); - folding_factor.push(whir_params.initial_sumcheck.num_rounds); + folding_factor.push(whir_params.initial_sumcheck.num_rounds()); for rc in &whir_params.round_configs { - folding_factor.push(rc.sumcheck.num_rounds); + folding_factor.push(rc.sumcheck.num_rounds()); } let ood_samples: Vec = whir_params .round_configs .iter() - .map(|rc| rc.irs_committer.out_domain_samples) + .map(|rc| rc.out_domain_samples) .collect(); let num_queries: Vec = whir_params .round_configs .iter() - .map(|rc| rc.irs_committer.in_domain_samples) + .map(|rc| rc.irs_committer.in_domain_samples()) .collect(); let pow_bits: Vec = whir_params @@ -115,9 +115,9 @@ impl WHIRConfigGnark { let sumcheck_pow_thresholds: Vec = whir_params .round_configs .iter() - .map(|rc| rc.sumcheck.round_pow.threshold) + .map(|rc| rc.sumcheck.round_pow().threshold) .collect(); - let initial_sumcheck_pow_threshold = whir_params.initial_sumcheck.round_pow.threshold; + let initial_sumcheck_pow_threshold = whir_params.initial_sumcheck.round_pow().threshold; let initial_skip_pow_threshold = whir_params.initial_skip_pow.threshold; // If there are no folding rounds, fall back to the initial commitment's @@ -125,14 +125,14 @@ impl WHIRConfigGnark { let final_queries = whir_params .round_configs .last() - .map_or(whir_params.initial_committer.in_domain_samples, |rc| { - rc.irs_committer.in_domain_samples + .map_or(whir_params.initial_committer.in_domain_samples(), |rc| { + rc.irs_committer.in_domain_samples() }); let final_pow_bits = f64::from(whir::protocols::proof_of_work::difficulty( whir_params.final_pow.threshold, )) as i32; let final_folding_pow_bits = f64::from(whir::protocols::proof_of_work::difficulty( - whir_params.final_sumcheck.round_pow.threshold, + whir_params.final_sumcheck.round_pow().threshold, )) as i32; // Reconstruct the starting domain to get its generator @@ -140,8 +140,8 @@ impl WHIRConfigGnark { .expect("Should have found an appropriate domain"); let domain_generator = format!("{}", domain.group_gen()); - let batch_size = whir_params.initial_committer.num_vectors; - let initial_in_domain_samples = whir_params.initial_committer.in_domain_samples; + let batch_size = whir_params.initial_committer.num_vectors(); + let initial_in_domain_samples = whir_params.initial_committer.in_domain_samples(); WHIRConfigGnark { n_rounds, @@ -159,7 +159,7 @@ impl WHIRConfigGnark { final_pow_bits, final_pow_threshold: whir_params.final_pow.threshold, final_folding_pow_bits, - final_folding_pow_threshold: whir_params.final_sumcheck.round_pow.threshold, + final_folding_pow_threshold: whir_params.final_sumcheck.round_pow().threshold, domain_generator, batch_size, initial_in_domain_samples,