Clean up prover implementation

This commit is contained in:
Sean Bowe 2020-08-27 14:03:43 -06:00
parent 154568c387
commit b453b845b8
No known key found for this signature in database
GPG key ID: 95684257D8F8B031
2 changed files with 76 additions and 73 deletions

View file

@ -115,7 +115,7 @@ impl<C: CurveAffine> Proof<C> {
&|index| advice_cosets[index].clone(), &|index| advice_cosets[index].clone(),
&|mut a, b| { &|mut a, b| {
parallelize(&mut a, |a, start| { parallelize(&mut a, |a, start| {
for (a, b) in a.into_iter().zip(b[start..].iter()) { for (a, b) in a.iter_mut().zip(b[start..].iter()) {
*a += b; *a += b;
} }
}); });
@ -123,7 +123,7 @@ impl<C: CurveAffine> Proof<C> {
}, },
&|mut a, b| { &|mut a, b| {
parallelize(&mut a, |a, start| { parallelize(&mut a, |a, start| {
for (a, b) in a.into_iter().zip(b[start..].iter()) { for (a, b) in a.iter_mut().zip(b[start..].iter()) {
*a *= b; *a *= b;
} }
}); });
@ -188,7 +188,9 @@ impl<C: CurveAffine> Proof<C> {
let fixed_evals: Vec<_> = meta let fixed_evals: Vec<_> = meta
.fixed_queries .fixed_queries
.iter() .iter()
.map(|&(wire, at)| eval_polynomial(&srs.fixed_polys[wire.0], domain.rotate_omega(x_3, at))) .map(|&(wire, at)| {
eval_polynomial(&srs.fixed_polys[wire.0], domain.rotate_omega(x_3, at))
})
.collect(); .collect();
let h_evals: Vec<_> = h_pieces let h_evals: Vec<_> = h_pieces
@ -227,78 +229,63 @@ impl<C: CurveAffine> Proof<C> {
let mut q_blinds = vec![C::Scalar::zero(); meta.rotations.len()]; let mut q_blinds = vec![C::Scalar::zero(); meta.rotations.len()];
let mut q_evals: Vec<_> = vec![C::Scalar::zero(); meta.rotations.len()]; let mut q_evals: Vec<_> = vec![C::Scalar::zero(); meta.rotations.len()];
{ {
let mut accumulate = |point_index: usize, new_poly: &Vec<_>, blind, eval| {
q_polys[point_index]
.as_mut()
.map(|poly| {
parallelize(poly, |q, start| {
for (q, a) in q.iter_mut().zip(new_poly[start..].iter()) {
*q *= &x_4;
*q += a;
}
});
})
.or_else(|| {
q_polys[point_index] = Some(new_poly.clone());
Some(())
});
q_blinds[point_index] *= &x_4;
q_blinds[point_index] += &blind;
q_evals[point_index] *= &x_4;
q_evals[point_index] += &eval;
};
for (query_index, &(wire, ref at)) in meta.advice_queries.iter().enumerate() { for (query_index, &(wire, ref at)) in meta.advice_queries.iter().enumerate() {
let point_index = (*meta.rotations.get(at).unwrap()).0; let point_index = (*meta.rotations.get(at).unwrap()).0;
if q_polys[point_index].is_none() { accumulate(
q_polys[point_index] = Some(advice_polys[wire.0].clone()); point_index,
q_blinds[point_index] = advice_blinds[wire.0]; &advice_polys[wire.0],
q_evals[point_index] = advice_evals[query_index]; advice_blinds[wire.0],
} else { advice_evals[query_index],
parallelize(q_polys[point_index].as_mut().unwrap(), |q, start| { );
for (q, a) in q.iter_mut().zip(advice_polys[wire.0][start..].iter()) {
*q *= &x_4;
*q += a;
}
});
q_blinds[point_index] *= &x_4;
q_blinds[point_index] += &advice_blinds[wire.0];
q_evals[point_index] *= &x_4;
q_evals[point_index] += &advice_evals[query_index];
}
} }
for (query_index, &(wire, ref at)) in meta.fixed_queries.iter().enumerate() { for (query_index, &(wire, ref at)) in meta.fixed_queries.iter().enumerate() {
let point_index = (*meta.rotations.get(at).unwrap()).0; let point_index = (*meta.rotations.get(at).unwrap()).0;
if q_polys[point_index].is_none() { accumulate(
q_polys[point_index] = Some(srs.fixed_polys[wire.0].clone()); point_index,
q_blinds[point_index] = C::Scalar::one(); &srs.fixed_polys[wire.0],
q_evals[point_index] = fixed_evals[query_index]; C::Scalar::one(),
} else { fixed_evals[query_index],
parallelize(q_polys[point_index].as_mut().unwrap(), |q, start| { );
for (q, a) in q.iter_mut().zip(srs.fixed_polys[wire.0][start..].iter()) {
*q *= &x_4;
*q += a;
}
});
q_blinds[point_index] *= &x_4;
q_blinds[point_index] += &C::Scalar::one();
q_evals[point_index] *= &x_4;
q_evals[point_index] += &fixed_evals[query_index];
}
} }
// We query the h(X) polynomial at x_3
let current_index = (*meta.rotations.get(&Rotation::default()).unwrap()).0;
for ((h_poly, h_blind), h_eval) in h_pieces for ((h_poly, h_blind), h_eval) in h_pieces
.into_iter() .into_iter()
.zip(h_blinds.iter()) .zip(h_blinds.iter())
.zip(h_evals.iter()) .zip(h_evals.iter())
{ {
// We query the h(X) polynomial at x_3 accumulate(current_index, &h_poly, *h_blind, *h_eval);
let point_index = (*meta.rotations.get(&Rotation::default()).unwrap()).0;
if q_polys[point_index].is_none() {
q_polys[point_index] = Some(h_poly);
q_blinds[point_index] = *h_blind;
q_evals[point_index] = *h_eval;
} else {
parallelize(q_polys[point_index].as_mut().unwrap(), |q, start| {
for (q, a) in q.iter_mut().zip(h_poly[start..].iter()) {
*q *= &x_4;
*q += a;
}
});
q_blinds[point_index] *= &x_4;
q_blinds[point_index] += h_blind;
q_evals[point_index] *= &x_4;
q_evals[point_index] += h_eval;
}
} }
} }
let x_5: C::Scalar = get_challenge_scalar(Challenge(transcript.squeeze().get_lower_128())); let x_5: C::Scalar = get_challenge_scalar(Challenge(transcript.squeeze().get_lower_128()));
let mut f_poly = None; let mut f_poly: Option<Vec<C::Scalar>> = None;
for (&row, &point_index) in meta.rotations.iter() { for (&row, &point_index) in meta.rotations.iter() {
let mut poly = q_polys[point_index.0].as_ref().unwrap().clone(); let mut poly = q_polys[point_index.0].as_ref().unwrap().clone();
let point = domain.rotate_omega(x_3, row); let point = domain.rotate_omega(x_3, row);
@ -306,16 +293,17 @@ impl<C: CurveAffine> Proof<C> {
let mut poly = kate_division(&poly, point); let mut poly = kate_division(&poly, point);
poly.push(C::Scalar::zero()); poly.push(C::Scalar::zero());
if f_poly.is_none() { f_poly = f_poly
f_poly = Some(poly); .map(|mut f_poly| {
} else { parallelize(&mut f_poly, |q, start| {
parallelize(f_poly.as_mut().unwrap(), |q, start| { for (q, a) in q.iter_mut().zip(poly[start..].iter()) {
for (q, a) in q.iter_mut().zip(poly[start..].iter()) { *q *= &x_5;
*q *= &x_5; *q += a;
*q += a; }
} });
}); f_poly
} })
.or_else(|| Some(poly));
} }
let mut f_poly = f_poly.unwrap(); let mut f_poly = f_poly.unwrap();
let mut f_blind = C::Scalar::random(); let mut f_blind = C::Scalar::random();

View file

@ -35,7 +35,12 @@ impl<C: CurveAffine> Proof<C> {
// transcript on the scalar field. // transcript on the scalar field.
let mut transcript_scalar = HScalar::init(C::Scalar::one()); let mut transcript_scalar = HScalar::init(C::Scalar::one());
for eval in self.advice_evals.iter().chain(self.fixed_evals.iter()).chain(self.h_evals.iter()) { for eval in self
.advice_evals
.iter()
.chain(self.fixed_evals.iter())
.chain(self.h_evals.iter())
{
transcript_scalar.absorb(*eval); transcript_scalar.absorb(*eval);
} }
@ -84,23 +89,33 @@ impl<C: CurveAffine> Proof<C> {
let mut q_evals: Vec<_> = vec![C::Scalar::zero(); srs.meta.rotations.len()]; let mut q_evals: Vec<_> = vec![C::Scalar::zero(); srs.meta.rotations.len()];
{ {
let mut accumulate = |point_index: usize, new_commitment, eval| { let mut accumulate = |point_index: usize, new_commitment, eval| {
q_commitments[point_index] = q_commitments[point_index].map(|mut commitment| { q_commitments[point_index] = q_commitments[point_index]
commitment *= x_4; .map(|mut commitment| {
commitment += new_commitment; commitment *= x_4;
commitment commitment += new_commitment;
}).or_else(|| Some(new_commitment.to_projective())); commitment
})
.or_else(|| Some(new_commitment.to_projective()));
q_evals[point_index] *= &x_4; q_evals[point_index] *= &x_4;
q_evals[point_index] += &eval; q_evals[point_index] += &eval;
}; };
for (query_index, &(wire, ref at)) in srs.meta.advice_queries.iter().enumerate() { for (query_index, &(wire, ref at)) in srs.meta.advice_queries.iter().enumerate() {
let point_index = (*srs.meta.rotations.get(at).unwrap()).0; let point_index = (*srs.meta.rotations.get(at).unwrap()).0;
accumulate(point_index, self.advice_commitments[wire.0], self.advice_evals[query_index]); accumulate(
point_index,
self.advice_commitments[wire.0],
self.advice_evals[query_index],
);
} }
for (query_index, &(wire, ref at)) in srs.meta.fixed_queries.iter().enumerate() { for (query_index, &(wire, ref at)) in srs.meta.fixed_queries.iter().enumerate() {
let point_index = (*srs.meta.rotations.get(at).unwrap()).0; let point_index = (*srs.meta.rotations.get(at).unwrap()).0;
accumulate(point_index, srs.fixed_commitments[wire.0], self.fixed_evals[query_index]); accumulate(
point_index,
srs.fixed_commitments[wire.0],
self.fixed_evals[query_index],
);
} }
let current_index = (*srs.meta.rotations.get(&Rotation::default()).unwrap()).0; let current_index = (*srs.meta.rotations.get(&Rotation::default()).unwrap()).0;