use crate::{math::Math, polys::eq::EqPolynomial};
use core::ops::Index;
use ff::PrimeField;
use rayon::prelude::*;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct MultilinearPolynomial<Scalar: PrimeField> {
pub(crate) Z: Vec<Scalar>, }
impl<Scalar: PrimeField> MultilinearPolynomial<Scalar> {
pub fn new(Z: Vec<Scalar>) -> Self {
MultilinearPolynomial { Z }
}
pub fn bind_poly_var_top(&mut self, r: &Scalar) {
assert!(
self.Z.len() >= 2,
"Vector Z must have at least two elements to bind the top variable."
);
let n = self.Z.len() / 2;
let (left, right) = self.Z.split_at_mut(n);
zip_with_for_each!((left.par_iter_mut(), right.par_iter()), |a, b| {
*a += *r * (*b - *a);
});
self.Z.truncate(n);
}
pub fn bind_with(poly: &[Scalar], L: &[Scalar], r_len: usize) -> Vec<Scalar> {
assert_eq!(
poly.len(),
L.len() * r_len,
"poly length ({}) must equal L.len() * r_len ({} * {}) = {}",
poly.len(),
L.len(),
r_len,
L.len() * r_len
);
(0..r_len)
.into_par_iter()
.map(|i| {
let mut acc = Scalar::ZERO;
for j in 0..L.len() {
acc += L[j] * poly[j * r_len + i];
}
acc
})
.collect()
}
}
impl<Scalar: PrimeField> Index<usize> for MultilinearPolynomial<Scalar> {
type Output = Scalar;
#[inline(always)]
fn index(&self, _index: usize) -> &Scalar {
&(self.Z[_index])
}
}
pub(crate) struct SparsePolynomial<Scalar: PrimeField> {
num_vars: usize,
Z: Vec<Scalar>,
}
impl<Scalar: PrimeField> SparsePolynomial<Scalar> {
pub fn new(num_vars: usize, Z: Vec<Scalar>) -> Self {
SparsePolynomial { num_vars, Z }
}
pub fn evaluate(&self, r: &[Scalar]) -> Scalar {
assert_eq!(self.num_vars, r.len());
let num_vars_z = self.Z.len().next_power_of_two().log_2();
let chis = EqPolynomial::evals_from_points(&r[self.num_vars - 1 - num_vars_z..]);
let eval_partial: Scalar = self
.Z
.iter()
.zip(chis.iter())
.map(|(z, chi)| *z * *chi)
.sum();
let common = (0..self.num_vars - 1 - num_vars_z)
.map(|i| Scalar::ONE - r[i])
.product::<Scalar>();
common * eval_partial
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::provider::pasta::pallas;
use rand_core::{CryptoRng, OsRng, RngCore};
pub fn evaluate<Scalar: PrimeField>(
poly: &MultilinearPolynomial<Scalar>,
r: &[Scalar],
) -> Scalar {
let chis = EqPolynomial::evals_from_points(r);
zip_with!(
(chis.into_par_iter(), poly.Z.par_iter()),
|chi_i, Z_i| chi_i * Z_i
)
.sum()
}
pub fn evaluate_with<Scalar: PrimeField>(Z: &[Scalar], r: &[Scalar]) -> Scalar {
zip_with!(
(
EqPolynomial::evals_from_points(r).into_par_iter(),
Z.par_iter()
),
|a, b| a * b
)
.sum()
}
fn test_multilinear_polynomial_with<F: PrimeField>() {
let TWO = F::from(2);
let Z = vec![
F::ZERO,
F::ZERO,
F::ZERO,
F::ONE,
F::ZERO,
F::ONE,
F::ZERO,
TWO,
];
let m_poly = MultilinearPolynomial::<F>::new(Z.clone());
let x = vec![F::ONE, F::ONE, F::ONE];
assert_eq!(evaluate(&m_poly, x.as_slice()), TWO);
let y = evaluate_with(Z.as_slice(), x.as_slice());
assert_eq!(y, TWO);
}
fn test_sparse_polynomial_with<F: PrimeField>() {
let mut Z = vec![F::ONE, F::ONE, F::from(2)];
let m_poly = SparsePolynomial::<F>::new(4, Z.clone());
Z.resize(16, F::ZERO); let m_poly_dense = MultilinearPolynomial::new(Z);
let x = vec![F::from(5), F::from(8), F::from(5), F::from(3)];
assert_eq!(
m_poly.evaluate(x.as_slice()),
evaluate(&m_poly_dense, x.as_slice())
);
}
#[test]
fn test_multilinear_polynomial() {
test_multilinear_polynomial_with::<pallas::Scalar>();
}
#[test]
fn test_sparse_polynomial() {
test_sparse_polynomial_with::<pallas::Scalar>();
}
fn test_evaluation_with<F: PrimeField>() {
let num_evals = 4;
let mut evals: Vec<F> = Vec::with_capacity(num_evals);
for _ in 0..num_evals {
evals.push(F::from(8));
}
let dense_poly: MultilinearPolynomial<F> = MultilinearPolynomial::new(evals.clone());
assert_eq!(
evaluate(&dense_poly, vec![F::from(3), F::from(4)].as_slice()),
F::from(8)
);
}
#[test]
fn test_evaluation() {
test_evaluation_with::<pallas::Scalar>();
}
fn random<R: RngCore + CryptoRng, Scalar: PrimeField>(
num_vars: usize,
mut rng: &mut R,
) -> MultilinearPolynomial<Scalar> {
MultilinearPolynomial::new(
std::iter::from_fn(|| Some(Scalar::random(&mut rng)))
.take(1 << num_vars)
.collect(),
)
}
fn bind_sequence<F: PrimeField>(
poly: &MultilinearPolynomial<F>,
values: &[F],
) -> MultilinearPolynomial<F> {
assert!(poly.Z.len().is_power_of_two());
assert!(poly.Z.len() >= 1 << values.len());
let mut tmp = poly.clone();
for v in values.iter() {
tmp.bind_poly_var_top(v);
}
tmp
}
fn bind_and_evaluate_with<F: PrimeField>() {
for _ in 0..50 {
let n = 7;
let poly = random(n, &mut OsRng);
let pt: Vec<_> = std::iter::from_fn(|| Some(F::random(&mut OsRng)))
.take(n)
.collect();
assert_eq!(evaluate(&poly, &pt), bind_sequence(&poly, &pt).Z[0])
}
}
#[test]
fn test_bind_and_evaluate() {
bind_and_evaluate_with::<pallas::Scalar>();
}
}