Perform inversions in lagrange_interpolate as part of a batch.

This commit is contained in:
Sean Bowe 2020-10-15 14:08:13 -06:00
parent 5c563eca12
commit 63d7de3bc2
No known key found for this signature in database
GPG key ID: 95684257D8F8B031

View file

@ -431,7 +431,9 @@ fn log2_floor(num: usize) -> u32 {
pow pow
} }
/// Returns coefficients of an n - 1 degree polynomial given a set of n points and their evaluations /// Returns coefficients of an n - 1 degree polynomial given a set of n points
/// and their evaluations. This function will panic if two values in `points`
/// are the same.
pub fn lagrange_interpolate<F: Field>(points: &[F], evals: &[F]) -> Vec<F> { pub fn lagrange_interpolate<F: Field>(points: &[F], evals: &[F]) -> Vec<F> {
assert_eq!(points.len(), evals.len()); assert_eq!(points.len(), evals.len());
if points.len() == 1 { if points.len() == 1 {
@ -439,14 +441,33 @@ pub fn lagrange_interpolate<F: Field>(points: &[F], evals: &[F]) -> Vec<F> {
return vec![evals[0]]; return vec![evals[0]];
} else { } else {
let mut interpolation_polys = vec![]; let mut interpolation_polys = vec![];
let mut denoms = Vec::with_capacity(points.len());
for (j, x_j) in points.iter().enumerate() { for (j, x_j) in points.iter().enumerate() {
let mut denom = Vec::with_capacity(points.len() - 1);
for x_k in points
.iter()
.enumerate()
.filter(|&(k, _)| k != j)
.map(|a| a.1)
{
denom.push(*x_j - x_k);
}
denoms.push(denom);
}
// Compute (x_j - x_k)^(-1) for each j != i
denoms.iter_mut().flat_map(|v| v.iter_mut()).batch_invert();
for (j, denoms) in denoms.into_iter().enumerate() {
let mut tmp: Vec<F> = Vec::with_capacity(points.len()); let mut tmp: Vec<F> = Vec::with_capacity(points.len());
let mut product = Vec::with_capacity(points.len() - 1); let mut product = Vec::with_capacity(points.len() - 1);
tmp.push(F::one()); tmp.push(F::one());
for (k, x_k) in points.iter().enumerate() { for (x_k, denom) in points
if k != j { .iter()
// Compute (x_j - x_k)^(-1) .enumerate()
let denom = (*x_j - x_k).invert().unwrap(); .filter(|&(k, _)| k != j)
.map(|a| a.1)
.zip(denoms.into_iter())
{
product.resize(tmp.len() + 1, F::zero()); product.resize(tmp.len() + 1, F::zero());
for ((a, b), product) in tmp for ((a, b), product) in tmp
.iter() .iter()
@ -458,7 +479,6 @@ pub fn lagrange_interpolate<F: Field>(points: &[F], evals: &[F]) -> Vec<F> {
} }
std::mem::swap(&mut tmp, &mut product); std::mem::swap(&mut tmp, &mut product);
} }
}
assert_eq!(tmp.len(), points.len()); assert_eq!(tmp.len(), points.len());
assert_eq!(product.len(), points.len() - 1); assert_eq!(product.len(), points.len() - 1);
interpolation_polys.push(tmp); interpolation_polys.push(tmp);