Skip to content
Merged
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
4 changes: 3 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 = "8804e80e8e890d01bb585f2bd5e5b564ac0fd80d", features = ["tracing", "rs_in_order"] }
54 changes: 41 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::{NttEngine, PolynomialSegment, Polynomials, 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,17 @@ 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 mask_refs = [&mask];
let segments = [
PolynomialSegment::from_rows(&vector_refs, 1),
PolynomialSegment::from_rows(&mask_refs, vector_refs.len()),
];
let codeword_length = message_length * expansion;
black_box(
RSFr.interleaved_encode(Polynomials::from_segments(&segments), codeword_length),
)
});
}

Expand All @@ -57,9 +76,18 @@ 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 mask_refs = [&mask];
let segments = [
PolynomialSegment::from_rows(&vector_refs, 1),
PolynomialSegment::from_rows(&mask_refs, vector_refs.len()),
];
let codeword_length = message_length * expansion;
black_box(
reference
.interleaved_encode(Polynomials::from_segments(&segments), codeword_length),
)
});
}

Expand Down
4 changes: 4 additions & 0 deletions provekit/backend/bn254/src/field.rs
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,10 @@ impl ProofField for Bn254Field {
}

impl FieldHash for Bn254Field {
fn register() {
crate::register();
}

fn hash_public_inputs(config: HashConfig, inputs: &[Base<Self>]) -> Ext<Self> {
crate::field_hash::hash_field_elements(config, inputs)
}
Expand Down
114 changes: 59 additions & 55 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::{Polynomials, ReedSolomon},
buffer::{Buffer, BufferOps},
},
};

#[derive(Debug)]
Expand Down Expand Up @@ -41,62 +44,56 @@ impl ReedSolomon<Fr> for RSFr {
.collect()
}

#[instrument(skip(self, messages, masks), fields(
num_messages = messages.len(),
message_len = messages.first().map(|c| c.len()),
#[instrument(skip(self, polynomials), fields(
num_polynomials = polynomials.len(),
polynomial_len = polynomials.polynomial_length(),
codeword_length = codeword_length,
mask_len = masks.len().checked_div(messages.len())

))]
fn interleaved_encode(
&self,
messages: &[&[Fr]],
masks: &[Fr],
polynomials: Polynomials<'_, Fr>,
codeword_length: usize,
) -> Vec<Fr> {
if messages.is_empty() {
return vec![];
}

let num_messages = messages.len();

let message_length = messages[0].len();
for message in messages {
assert_eq!(message_length, message.len())
) -> Buffer<Fr> {
let num_polynomials = polynomials.len();
if num_polynomials == 0 {
return Buffer::from(vec![]);
}

let total_size = num_messages * codeword_length;
let polynomial_length = polynomials.polynomial_length();
assert!(polynomial_length <= codeword_length);
let total_size = num_polynomials * codeword_length;

let mut result = vec![Fr::ZERO; total_size];

(0..message_length).for_each(|column| {
let base = column * num_messages;
for row in 0..num_messages {
result[base + row] = messages[row][column];
let mut column_offset = 0;
for segment in polynomials.segments() {
let row_width = segment.row_width();
let rows_per_buffer = segment.rows_per_buffer();
for polynomial in 0..num_polynomials {
let buffer = segment.buffer(polynomial / rows_per_buffer).to_slice();
let row = polynomial % rows_per_buffer;
let start = row * row_width;
for column in 0..row_width {
result[(column_offset + column) * num_polynomials + polynomial] =
buffer[start + column];
}
}
});

result[message_length * num_messages..message_length * num_messages + masks.len()]
.copy_from_slice(masks);

let mask_length = masks.len() / num_messages;

let masked_message_length = message_length + mask_length;
column_offset += row_width;
}

let mut coset_size = self.next_order(masked_message_length).unwrap();
while !codeword_length.is_multiple_of(coset_size) {
let mut coset_size = self.next_order(polynomial_length).unwrap();
while codeword_length % coset_size != 0 {
coset_size = self.next_order(coset_size + 1).unwrap();
}
let num_cosets = codeword_length / coset_size;

let chunk_size = coset_size * num_messages;
let chunk_size = coset_size * num_polynomials;
for k in 1..num_cosets {
result.copy_within(0..chunk_size, k * chunk_size);
}

ntt_nr(&mut result, codeword_length, num_cosets);

result
Buffer::from(result)
}

fn generator(&self, codeword_length: usize) -> Fr {
Expand Down Expand Up @@ -136,42 +133,49 @@ 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(Buffer::from)
.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, ...]
let mask_total = num_messages * mask_length;
let mut masks = masks_flat;
masks.resize(mask_total, Fr::ZERO);

// Whir expects masks per-message (column-major from our perspective):
// [m0_c0, m0_c1, ..., m1_c0, m1_c1, ...]
// Transpose the num_messages x mask_length matrix.
let mut masks_transposed = vec![Fr::ZERO; mask_total];
for row in 0..num_messages {
for col in 0..mask_length {
masks_transposed[row * mask_length + col] = masks[col * num_messages + row];
}
}
let masks = Buffer::from(masks);
let mask_refs = [&masks];
let segments = [
whir::algebra::ntt::PolynomialSegment::from_rows(&messages_refs, 1),
whir::algebra::ntt::PolynomialSegment::from_rows(&mask_refs, num_messages),
];

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_points = RSFr.evaluation_points(message_length, codeword_length, &indices);
let ref_points = reference.evaluation_points(message_length, codeword_length, &indices);
let our_codeword = RSFr.interleaved_encode(
Polynomials::from_segments(&segments),
codeword_length,
);
let ref_codeword = reference.interleaved_encode(
Polynomials::from_segments(&segments),
codeword_length,
);

let our_points =
RSFr.evaluation_points(masked_message_length, codeword_length, &indices);
let ref_points =
reference.evaluation_points(masked_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
4 changes: 4 additions & 0 deletions provekit/backend/goldilocks/src/field.rs
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,10 @@ impl ProofField for GoldilocksEfField {
macro_rules! impl_goldilocks_field_hash {
($field:ty) => {
impl FieldHash for $field {
fn register() {
crate::register();
}

fn hash_public_inputs(config: HashConfig, inputs: &[Base<Self>]) -> Ext<Self> {
hash_field_elements(config, inputs)
}
Expand Down
5 changes: 5 additions & 0 deletions provekit/common/src/field.rs
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,11 @@ pub type Ext<P> = <<P as ProofField>::Embedding as Embedding>::Target;
/// Hash and byte-bridge glue, kept out of [`ProofField`]'s algebra surface and
/// composed as a supertrait so [`Ext<Self>`] is nameable.
pub trait FieldHash: ProofField {
/// Register this field's engines in WHIR's global registries.
///
/// Implementations must be idempotent.
fn register();

/// Instance-binding hash of base-field public inputs to an extension
/// transcript element.
fn hash_public_inputs(config: crate::HashConfig, inputs: &[Base<Self>]) -> Ext<Self>;
Expand Down
5 changes: 4 additions & 1 deletion provekit/common/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -29,5 +29,8 @@ pub use {
public_inputs::{PublicInputs, PublicInputsHash},
r1cs::R1CS,
sparse_matrix::{HydratedSparseMatrix, SparseMatrix},
whir_r1cs::{ProvekitProof, R1csHash, WhirR1CSProof, WhirR1CSScheme, MIN_WHIR_NUM_VARIABLES},
whir_r1cs::{
whir_protocol_params, ProvekitProof, R1csHash, WhirR1CSProof, WhirR1CSScheme,
MIN_WHIR_NUM_VARIABLES,
},
};
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/utils/serde_ark_vec.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ use {
std::{fmt, marker::PhantomData},
};

pub fn serialize<T, S>(vec: &Vec<T>, serializer: S) -> Result<S::Ok, S::Error>
pub fn serialize<T, S>(vec: &[T], serializer: S) -> Result<S::Ok, S::Error>
where
T: CanonicalSerialize,
S: Serializer,
Expand Down
Loading
Loading