Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

2 changes: 1 addition & 1 deletion Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"] }
41 changes: 28 additions & 13 deletions provekit/backend/bn254/benches/rs_bench.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -19,21 +22,29 @@ const TEST_CASES: &[(usize, usize, usize)] = &[
(22, 4, 4),
];

fn make_messages(exp: usize, coset_sz: usize) -> Vec<Vec<Fr>> {
fn make_messages(exp: usize, coset_sz: usize) -> Vec<Buffer<Fr>> {
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::<Vec<_>>(),
)
})
.collect()
}

fn make_mask(num_messages: usize) -> Vec<Fr> {
fn make_mask(num_messages: usize) -> Buffer<Fr> {
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::<Vec<_>>(),
)
}

#[divan::bench(args = TEST_CASES)]
Expand All @@ -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<Fr>> = 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))
});
}

Expand All @@ -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<Fr>> = 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))
});
}

Expand Down
51 changes: 35 additions & 16 deletions provekit/backend/bn254/src/ntt.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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)]
Expand Down Expand Up @@ -42,26 +45,36 @@ impl ReedSolomon<Fr> 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<Fr>,
codeword_length: usize,
) -> Vec<Fr> {
) -> Buffer<Fr> {
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::<Vec<_>>();
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())
}

Expand Down Expand Up @@ -96,7 +109,7 @@ impl ReedSolomon<Fr> for RSFr {

ntt_nr(&mut result, codeword_length, num_cosets);

result
Buffer::from(result)
}

fn generator(&self, codeword_length: usize) -> Fr {
Expand Down Expand Up @@ -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<Buffer<Fr>> = data.chunks(message_length).map(|c| Buffer::from(c)).collect();
let messages_refs: Vec<&Buffer<Fr>> = 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, ...]
Expand All @@ -154,24 +168,29 @@ 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<usize> = (0..codeword_length).collect();

let reference = NttEngine::<Fr>::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);

// 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);

Expand Down
15 changes: 11 additions & 4 deletions provekit/common/src/prefix_covector.rs
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -209,9 +212,10 @@ pub fn build_prefix_covectors<const N: usize, F: Field>(
#[must_use]
pub fn compute_alpha_evals<const N: usize, M: Embedding>(
embedding: &M,
polynomial: &[M::Source],
polynomial: &Buffer<M::Source>,
alphas: &[Vec<M::Target>; N],
) -> Vec<M::Target> {
let polynomial = polynomial.to_slice();
alphas
.iter()
.map(|w| mixed_dot(embedding, w, &polynomial[..w.len()]))
Expand All @@ -226,8 +230,9 @@ pub fn compute_public_eval<M: Embedding>(
embedding: &M,
x: M::Target,
num_public_inputs: usize,
polynomial: &[M::Source],
polynomial: &Buffer<M::Source>,
) -> 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();
Expand Down Expand Up @@ -331,8 +336,9 @@ pub fn compute_challenge_eval<M: Embedding>(
embedding: &M,
x: M::Target,
challenge_offsets: &[usize],
polynomial: &[M::Source],
polynomial: &Buffer<M::Source>,
) -> 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 {
Expand Down Expand Up @@ -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::<FieldElement>::new();
let eval = compute_challenge_eval(&embedding, x, &offsets, &poly);
Expand Down
2 changes: 1 addition & 1 deletion provekit/common/src/whir_r1cs.rs
Original file line number Diff line number Diff line change
Expand Up @@ -199,7 +199,7 @@ impl<P: FieldHash> WhirR1CSScheme<P> {
/// (≈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,
Expand Down
Loading